From 098e41f003335966f9e6b7f8cf04ae88e851fb82 Mon Sep 17 00:00:00 2001 From: Lucas Machado Date: Thu, 17 Sep 2026 16:54:19 +0200 Subject: [PATCH 01/20] feat: safe reads, reliable jobs and live feedback - Read-only scans with server-side timeouts, MySQL KILL QUERY on cancel - Production flag: writes need the label typed back (web, CLI) - Panics fail one run, not the process; failures name side, phase, table - Unknown counts are never zero; partitioned tables and FKs to non-key columns seed correctly - Run strip with reattach, quiet/lost states; TUI and CLI progress - Workspace draws before counting; settings survive navigation --- Makefile | 7 + cmd/seedstorm/main.go | 49 +- e2e/support/selectors.ts | 17 + e2e/support/test.fixture.ts | 2 + e2e/tests/compare.spec.ts | 24 + e2e/tests/memory.spec.ts | 74 +++ e2e/tests/production.spec.ts | 56 ++ integration/introspect_constraints_test.go | 138 +++++ integration/loadsim_containers_test.go | 143 +++++ integration/loadsim_profiles_test.go | 55 ++ integration/loadsim_readscope_test.go | 46 ++ integration/partitions_test.go | 95 ++++ integration/production_test.go | 58 ++ integration/read_scope_test.go | 200 +++++++ integration/reliability_test.go | 136 +++++ integration/reliability_web_test.go | 105 ++++ integration/server_info_test.go | 58 ++ integration/web_jobs_helpers_test.go | 4 +- internal/cli/clone_schema.go | 12 + internal/cli/compare.go | 6 +- internal/cli/endpoints.go | 23 +- internal/cli/gaps.go | 50 +- internal/cli/generate.go | 7 +- internal/cli/helpers.go | 23 + internal/cli/introspect.go | 15 +- internal/cli/mirror.go | 7 + internal/cli/production.go | 33 ++ internal/cli/progress.go | 44 ++ internal/cli/seed.go | 27 +- internal/cli/snapshot.go | 3 +- internal/compare/compare.go | 50 +- internal/compare/compare_test.go | 84 +++ internal/compare/mirror.go | 30 +- internal/db/access.go | 28 +- internal/db/clone.go | 9 +- internal/db/copy.go | 11 +- internal/db/counts.go | 81 ++- internal/db/db.go | 25 +- internal/db/mysql.go | 52 +- internal/db/partitions.go | 147 ++++++ internal/db/partitions_test.go | 46 ++ internal/db/postgres.go | 122 +++-- internal/db/read_scope.go | 189 +++++++ internal/db/server_info.go | 96 ++++ internal/db/stats.go | 74 ++- internal/db/types.go | 23 + internal/faker/catalog.go | 5 + internal/faker/existing.go | 25 +- internal/faker/faker.go | 47 +- internal/faker/partitions.go | 224 ++++++++ internal/faker/partitions_test.go | 112 ++++ internal/faker/references.go | 78 +++ internal/faker/references_test.go | 86 +++ internal/faker/stream.go | 76 ++- internal/faultinject/faultinject_off.go | 11 + internal/faultinject/faultinject_on.go | 47 ++ internal/runerr/runerr.go | 108 ++++ internal/runerr/runerr_test.go | 47 ++ internal/safego/safego.go | 51 ++ internal/safego/safego_test.go | 54 ++ internal/schema/schema.go | 8 + internal/seeder/gaps.go | 30 ++ internal/seeder/gaps_test.go | 21 + internal/seeder/mirror.go | 38 +- internal/seeder/mirror_snapshot_test.go | 2 +- internal/seeder/panic_test.go | 79 +++ internal/seeder/seed.go | 82 ++- internal/seeder/seeder.go | 23 +- internal/seeder/writer.go | 25 +- internal/seeder/writer_test.go | 34 ++ internal/tui/clone.go | 40 +- internal/tui/execute.go | 186 +++++-- internal/tui/execute_test.go | 39 ++ internal/tui/gaps.go | 58 +- internal/tui/mirror.go | 67 ++- internal/tui/mirror_test.go | 22 + internal/tui/tui.go | 3 +- internal/tuning/host.go | 122 +++++ internal/tuning/host_test.go | 103 ++++ internal/tuning/recommend.go | 248 +++++++++ internal/tuning/recommend_test.go | 114 ++++ internal/web/counts_cache_test.go | 73 +++ internal/web/failures.go | 56 ++ internal/web/handlers_api.go | 42 +- internal/web/handlers_compare.go | 58 +- internal/web/handlers_connections.go | 70 ++- internal/web/handlers_jobs.go | 90 +++- internal/web/handlers_jobs_test.go | 73 +++ internal/web/handlers_pages.go | 14 +- internal/web/handlers_runs.go | 56 +- internal/web/jobs.go | 83 ++- internal/web/jobs_list_test.go | 127 +++++ internal/web/preflight.go | 81 +++ internal/web/production.go | 119 +++++ internal/web/production_test.go | 134 +++++ internal/web/recover.go | 23 + internal/web/recover_test.go | 74 +++ internal/web/runners.go | 63 ++- internal/web/runners_test.go | 127 ++++- internal/web/server.go | 12 +- internal/web/session.go | 132 ++++- internal/web/session_test.go | 40 ++ internal/web/static/app.js | 582 ++++++++++++++++++--- internal/web/static/compare.js | 110 +++- internal/web/static/profiles.js | 38 +- internal/web/static/style.css | 28 + internal/web/store.go | 4 + internal/web/templates/compare.html.tmpl | 1 + internal/web/templates/connect.html.tmpl | 14 + internal/web/templates/layout.html.tmpl | 3 +- internal/web/templates/workspace.html.tmpl | 2 + 111 files changed, 6642 insertions(+), 586 deletions(-) create mode 100644 e2e/tests/memory.spec.ts create mode 100644 e2e/tests/production.spec.ts create mode 100644 integration/introspect_constraints_test.go create mode 100644 integration/loadsim_containers_test.go create mode 100644 integration/loadsim_profiles_test.go create mode 100644 integration/loadsim_readscope_test.go create mode 100644 integration/partitions_test.go create mode 100644 integration/production_test.go create mode 100644 integration/read_scope_test.go create mode 100644 integration/reliability_test.go create mode 100644 integration/reliability_web_test.go create mode 100644 integration/server_info_test.go create mode 100644 internal/cli/production.go create mode 100644 internal/db/partitions.go create mode 100644 internal/db/partitions_test.go create mode 100644 internal/db/read_scope.go create mode 100644 internal/db/server_info.go create mode 100644 internal/faker/partitions.go create mode 100644 internal/faker/partitions_test.go create mode 100644 internal/faker/references.go create mode 100644 internal/faker/references_test.go create mode 100644 internal/faultinject/faultinject_off.go create mode 100644 internal/faultinject/faultinject_on.go create mode 100644 internal/runerr/runerr.go create mode 100644 internal/runerr/runerr_test.go create mode 100644 internal/safego/safego.go create mode 100644 internal/safego/safego_test.go create mode 100644 internal/seeder/gaps.go create mode 100644 internal/seeder/gaps_test.go create mode 100644 internal/seeder/panic_test.go create mode 100644 internal/tuning/host.go create mode 100644 internal/tuning/host_test.go create mode 100644 internal/tuning/recommend.go create mode 100644 internal/tuning/recommend_test.go create mode 100644 internal/web/counts_cache_test.go create mode 100644 internal/web/failures.go create mode 100644 internal/web/jobs_list_test.go create mode 100644 internal/web/preflight.go create mode 100644 internal/web/production.go create mode 100644 internal/web/production_test.go create mode 100644 internal/web/recover.go create mode 100644 internal/web/recover_test.go diff --git a/Makefile b/Makefile index 47bcc3a..9727e7f 100644 --- a/Makefile +++ b/Makefile @@ -64,6 +64,13 @@ test-integration: dev-up done cd integration && go test -race -v -tags integration -count=1 ./... -timeout 1500s +.PHONY: test-loadsim +# Resource-limited evals: throwaway database containers shaped like managed +# Cloud SQL instances (integration/loadsim_*_test.go). Needs Docker; skips when +# the machine has too little free memory. Pass ARGS to filter, e.g. ARGS=-run=TestLoadsim_Read +test-loadsim: + cd integration && go test -race -v -tags "integration loadsim" -count=1 -run TestLoadsim $(ARGS) ./... -timeout 1800s + .PHONY: test-e2e # Playwright journeys against a freshly built `seedstorm serve` and the compose # databases (make dev-up). Pass ARGS to filter, e.g. ARGS=tests/compare.spec.ts diff --git a/cmd/seedstorm/main.go b/cmd/seedstorm/main.go index 8e40cda..e2087f2 100644 --- a/cmd/seedstorm/main.go +++ b/cmd/seedstorm/main.go @@ -2,18 +2,61 @@ package main import ( "context" + "errors" "fmt" "os" + "os/signal" + "syscall" + + tea "github.com/charmbracelet/bubbletea" "github.com/AxeForging/seedstorm/internal/app" + "github.com/AxeForging/seedstorm/internal/safego" _ "github.com/go-sql-driver/mysql" _ "github.com/jackc/pgx/v5/stdlib" ) +// Exit codes: 1 a run failed, 70 an internal error (a bug: please report it), +// 130 interrupted. +const ( + exitFailed = 1 + exitInternal = 70 + exitInterrupted = 130 +) + func main() { - if err := app.New().Run(context.Background(), os.Args); err != nil { - fmt.Fprintf(os.Stderr, "error: %v\n", err) - os.Exit(1) + os.Exit(run()) +} + +func run() int { + // The first Ctrl+C cancels the run, so reads stop on the server (MySQL + // queries are killed) and the summary is printed; a second one exits now. + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + go func() { + <-ctx.Done() + stop() + again := make(chan os.Signal, 1) + signal.Notify(again, os.Interrupt, syscall.SIGTERM) + <-again + fmt.Fprintln(os.Stderr, "interrupted again: exiting now") + os.Exit(exitInterrupted) + }() + + err := safego.Run("seedstorm", func() error { return app.New().Run(ctx, os.Args) }) + switch { + case err == nil: + return 0 + case ctx.Err() != nil: + fmt.Fprintf(os.Stderr, "interrupted: %v\n", err) + return exitInterrupted + } + var p *safego.PanicError + if errors.As(err, &p) || errors.Is(err, tea.ErrProgramPanic) { + fmt.Fprintf(os.Stderr, "error: %v\nThis is a bug in seedstorm; run with --log-level debug for the stack and please report it.\n", err) + return exitInternal } + fmt.Fprintf(os.Stderr, "error: %v\n", err) + return exitFailed } diff --git a/e2e/support/selectors.ts b/e2e/support/selectors.ts index d25adc8..2e3b50f 100644 --- a/e2e/support/selectors.ts +++ b/e2e/support/selectors.ts @@ -29,6 +29,7 @@ export const sel = { actionBar: "ws-actionbar", canvas: "ws-canvas", graphLoading: "ws-graph-loading", + countsStatus: "ws-counts-status", tableCount: "ws-count-total", populatedCount: "ws-stat-populated", access: "ws-access", @@ -60,6 +61,7 @@ export const sel = { tuningToggle: "ws-tuning-toggle", tuningSummary: "ws-tuning-summary", workers: "ws-workers", + genWorkers: "ws-gen-workers", profile: "ws-profile", cloneTarget: "ws-clone-target", cloneViews: "ws-clone-views", @@ -115,4 +117,19 @@ export const sel = { ignoreHit: "pf-ignore-hit", yamlDownload: "pf-yaml-download", }, + runStrip: { + root: "run-strip", + status: "run-strip-status", + failure: "run-strip-failure", + next: "run-strip-next", + detail: "run-strip-detail", + }, + production: { + checkbox: "conn-production", + badge: "conn-production-badge", + confirmLabel: "conn-confirm-label", + dialog: "prod-confirm-dialog", + input: "prod-confirm-input", + ok: "prod-confirm-ok", + }, } as const; diff --git a/e2e/support/test.fixture.ts b/e2e/support/test.fixture.ts index f65c248..1660e9f 100644 --- a/e2e/support/test.fixture.ts +++ b/e2e/support/test.fixture.ts @@ -35,6 +35,8 @@ export async function openWorkspace(page: Page, tables: number): Promise { await page.goto("/"); await expect(page.getByTestId(sel.workspace.tableCount)).toHaveText(String(tables)); await expect(page.getByTestId(sel.workspace.graphLoading)).toBeHidden(); + // Row counts fill in after the graph is drawn. + await expect(page.getByTestId(sel.workspace.countsStatus)).toHaveAttribute("data-state", "ready"); await waitForCamera(page); } diff --git a/e2e/tests/compare.spec.ts b/e2e/tests/compare.spec.ts index 1fe655e..33b5314 100644 --- a/e2e/tests/compare.spec.ts +++ b/e2e/tests/compare.spec.ts @@ -120,3 +120,27 @@ test("compare, export and import counts, then mirror from the imported file", as await expectReport(SOURCE_ROWS); }); }); + +// A job that fails on the server must end in the page: the button comes back +// and the reason is shown. The failure event used to be named "error", which +// browsers treat as a broken connection, so the page waited forever. +test("a compare that fails on the server ends with its reason instead of spinning", async ({ page }) => { + await connectPostgres(page, DB.tgt); + await connectPostgres(page, DB.src); + await page.goto("/compare"); + await page.getByTestId(cmp.target).selectOption({ label: pgConnectionLabel(DB.tgt) }); + + // The target disappears between choosing it and running the job. + await page.route("**/api/compare", async (route) => { + const body = route.request().postDataJSON(); + body.target = { id: "gone-" + Date.now() }; + await route.continue({ postData: JSON.stringify(body) }); + }); + + const run = page.getByTestId(cmp.run); + await run.click(); + await expect(page.getByTestId(cmp.outcome)).toContainText("Compare failed", { timeout: 15_000 }); + await expect(page.getByTestId(cmp.outcome)).toContainText("target connection not found"); + await expect(run).toBeEnabled(); + await expect(run).toHaveText("Compare"); +}); diff --git a/e2e/tests/memory.spec.ts b/e2e/tests/memory.spec.ts new file mode 100644 index 0000000..c293021 --- /dev/null +++ b/e2e/tests/memory.spec.ts @@ -0,0 +1,74 @@ +// Leaving a page and coming back keeps what the user set, reattaches to the +// job they started, and never remembers a destructive toggle. +import { DB } from "../support/db.helpers"; +import { sel } from "../support/selectors"; +import { connectPostgres, expect, openWorkspace, pgConnectionLabel, test } from "../support/test.fixture"; + +const ws = sel.workspace; +const cmp = sel.compare; + +test("workspace and compare settings survive a trip to another page", async ({ page }) => { + const created = await page.request.post("/api/profiles", { + headers: { "Content-Type": "application/json" }, + data: { rules: { version: 1, name: "e2e-memory", rules: [{ column: "*name*", template: "N {{seq}}" }] } }, + }); + expect(created.ok()).toBe(true); + const profileId = (await created.json()).id as string; + try { + await connectPostgres(page, DB.tgt); + await connectPostgres(page, DB.src); + await openWorkspace(page, 2); + + await page.getByTestId(ws.rows).fill("37"); + await page.getByTestId(ws.tuningToggle).click(); + await page.getByTestId(ws.workers).fill("3"); + await page.getByTestId(ws.genWorkers).fill("2"); + await page.getByTestId(ws.profile).selectOption(profileId); + await page.locator("#cfg-truncate").check(); + + await page.goto("/compare"); + await page.getByTestId(cmp.target).selectOption({ label: pgConnectionLabel(DB.tgt) }); + const chosenTarget = await page.getByTestId(cmp.target).inputValue(); + // Mirror settings sit next to a report: compare first. + await page.getByTestId(cmp.run).click(); + await expect(page.getByTestId(cmp.results)).toBeVisible(); + await page.getByTestId(cmp.advancedToggle).click(); + await page.getByTestId(cmp.stopOnError).check(); + + await page.goto("/profiles"); + await page.goto("/compare"); + await expect(page.getByTestId(cmp.target)).toHaveValue(chosenTarget); + // The report for the same pair is restored, with the remembered settings. + await expect(page.getByTestId(cmp.results)).toBeVisible(); + await page.getByTestId(cmp.advancedToggle).click(); + await expect(page.getByTestId(cmp.stopOnError)).toBeChecked(); + + await openWorkspace(page, 2); + await expect(page.getByTestId(ws.rows)).toHaveValue("37"); + await expect(page.getByTestId(ws.workers)).toHaveValue("3"); + await expect(page.getByTestId(ws.genWorkers)).toHaveValue("2"); + await expect(page.getByTestId(ws.tuningSummary)).toHaveText("3 writers · 2 gen"); + await expect(page.getByTestId(ws.profile)).toHaveValue(profileId); + await expect(page.locator("#cfg-truncate"), "truncate is never remembered").not.toBeChecked(); + } finally { + await page.request.delete(`/api/profiles?id=${profileId}`); + } +}); + +test("a run started before leaving the workspace is shown when coming back", async ({ page }) => { + test.setTimeout(120_000); + await connectPostgres(page, DB.src); + await openWorkspace(page, 2); + // A dry run of a few million rows takes seconds and writes nothing. + await page.getByTestId(ws.rows).fill("1500000"); + await page.locator("#cfg-dryrun").check(); + await page.getByTestId(ws.run).click(); + await expect(page.getByTestId(sel.runStrip.root).first()).toHaveAttribute("data-state", "running"); + await page.goto("/compare"); + await openWorkspace(page, 2); + const strip = page.getByTestId(sel.runStrip.root).first(); + await expect(strip).toBeVisible(); + await expect(strip.getByTestId(sel.runStrip.status)).toHaveText(/Running|Still working|Done/); + await expect(strip).toHaveAttribute("data-state", "done", { timeout: 90_000 }); + await expect(strip.getByTestId(sel.runStrip.status)).toHaveText("Done"); +}); diff --git a/e2e/tests/production.spec.ts b/e2e/tests/production.spec.ts new file mode 100644 index 0000000..dea4213 --- /dev/null +++ b/e2e/tests/production.spec.ts @@ -0,0 +1,56 @@ +// A saved connection marked production is never written from the UI until its +// label is typed back: cancelling writes nothing, typing it runs the seed. +import { DB, pg, pgCounts, withPg } from "../support/db.helpers"; +import { sel } from "../support/selectors"; +import { connectPostgres, expect, openWorkspace, test } from "../support/test.fixture"; + +const LABEL = "e2e-orders-prod"; +const prod = sel.production; + +test("seeding a production connection asks for its label first", async ({ page }) => { + await withPg(DB.tgt, (c) => c.query("TRUNCATE orders, customers RESTART IDENTITY")); + const saved = await page.request.post("/api/saved-connections", { + headers: { "X-Seedstorm-Request": "1" }, + data: { label: LABEL, dbType: "postgres", host: pg.host, port: pg.port, dbName: DB.tgt, user: pg.user, password: pg.password, production: true }, + }); + expect(saved.ok()).toBe(true); + const savedId = (await saved.json()).id as string; + test.info().attach("saved connection", { body: savedId }); + + try { + await test.step("the saved list shows the production badge", async () => { + await page.goto("/connect?mode=chooser"); + await expect(page.getByTestId(prod.badge).first()).toBeVisible(); + }); + + // Connected ad hoc (a DSN, not the saved entry): it is still recognised. + await connectPostgres(page, DB.tgt); + await openWorkspace(page, 2); + await page.getByTestId(sel.workspace.rows).fill("3"); + + await test.step("cancelling the confirmation writes nothing", async () => { + await page.getByTestId(sel.workspace.run).click(); + const dialog = page.getByTestId(prod.dialog); + await expect(dialog).toBeVisible(); + await expect(dialog.getByTestId(prod.ok)).toBeDisabled(); + await dialog.getByRole("button", { name: "Cancel" }).click(); + await expect(dialog).toBeHidden(); + await expect(page.locator(".job-phase-log").last()).toContainText("Not written"); + expect(await pgCounts(DB.tgt, ["customers", "orders"])).toEqual({ customers: 0, orders: 0 }); + }); + + await test.step("typing the label runs the seed", async () => { + await page.getByTestId(sel.workspace.run).click(); + const dialog = page.getByTestId(prod.dialog); + await dialog.getByTestId(prod.input).fill("e2e-orders"); + await expect(dialog.getByTestId(prod.ok)).toBeDisabled(); + await dialog.getByTestId(prod.input).fill(LABEL); + await dialog.getByTestId(prod.ok).click(); + await expect(page.getByTestId(sel.job.status)).toHaveText("done", { timeout: 60_000 }); + expect(await pgCounts(DB.tgt, ["customers", "orders"])).toEqual({ customers: 3, orders: 3 }); + }); + } finally { + await page.request.delete(`/api/saved-connections?id=${savedId}`, { headers: { "X-Seedstorm-Request": "1" } }); + await withPg(DB.tgt, (c) => c.query("TRUNCATE orders, customers RESTART IDENTITY")); + } +}); diff --git a/integration/introspect_constraints_test.go b/integration/introspect_constraints_test.go new file mode 100644 index 0000000..49a746f --- /dev/null +++ b/integration/introspect_constraints_test.go @@ -0,0 +1,138 @@ +//go:build integration + +package integration_test + +import ( + "context" + "fmt" + "path/filepath" + "strings" + "testing" + + "github.com/AxeForging/seedstorm/internal/db" +) + +// fkTargets maps "table.column" to the "table.column" its FK references. +func fkTargets(tables []db.Table) map[string]string { + out := map[string]string{} + for _, t := range tables { + for _, c := range t.Columns { + if c.FK != nil { + out[t.Name+"."+c.Name] = c.FK.TableName + "." + c.FK.ColumnName + } + } + } + return out +} + +func pkColumns(tables []db.Table) map[string][]string { + out := map[string][]string{} + for _, t := range tables { + for _, c := range t.Columns { + if c.IsPK { + out[t.Name] = append(out[t.Name], c.Name) + } + } + } + return out +} + +// Postgres FKs were read from information_schema joined by constraint name +// only: a two-column FK paired every column with every referenced column (last +// one won), two tables with a same-named constraint mixed their columns, and a +// role with only SELECT saw no constraints at all. +func TestIntrospect_PostgresConstraintsArePairedAndVisibleToReadOnlyRoles(t *testing.T) { + e := postgresEngine() + dsn, owner := e.scratchDB(t, "ss_introspect_fk") + execSQL(t, owner, ` + CREATE TABLE regions (country CHAR(2), code VARCHAR(8), name TEXT, PRIMARY KEY (country, code)); + CREATE TABLE stores (id SERIAL PRIMARY KEY, country CHAR(2) NOT NULL, region VARCHAR(8) NOT NULL, + CONSTRAINT loc_fk FOREIGN KEY (region, country) REFERENCES regions (code, country)); + CREATE TABLE brands (id SERIAL PRIMARY KEY, name TEXT); + CREATE TABLE products (id SERIAL PRIMARY KEY, brand_id INT NOT NULL, store_id INT NOT NULL, + CONSTRAINT loc_fk FOREIGN KEY (brand_id) REFERENCES brands (id), + CONSTRAINT store_fk FOREIGN KEY (store_id) REFERENCES stores (id))`) + + want := map[string]string{ + "stores.country": "regions.country", + "stores.region": "regions.code", + "products.brand_id": "brands.id", + "products.store_id": "stores.id", + } + check := func(t *testing.T, dsn string) { + t.Helper() + tables, err := db.Introspect(e.driver, dsn) + if err != nil { + t.Fatal(err) + } + got := fkTargets(tables) + for col, target := range want { + if got[col] != target { + t.Errorf("FK %s -> %q, want %q (all: %v)", col, got[col], target, got) + } + } + if len(got) != len(want) { + t.Errorf("FKs = %v, want exactly %v", got, want) + } + pks := pkColumns(tables) + if strings.Join(pks["regions"], ",") != "country,code" || strings.Join(pks["stores"], ",") != "id" { + t.Errorf("PKs = %v", pks) + } + } + + t.Run("owner", func(t *testing.T) { check(t, dsn) }) + + t.Run("select-only role", func(t *testing.T) { + const role = "ss_introspect_reader" + ctx := context.Background() + drop := func() { + var exists bool + _ = owner.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM pg_roles WHERE rolname = $1)`, role).Scan(&exists) + if exists { + _, _ = owner.ExecContext(ctx, `DROP OWNED BY `+role) + _, _ = owner.ExecContext(ctx, `DROP ROLE `+role) + } + } + drop() + t.Cleanup(drop) + execSQL(t, owner, fmt.Sprintf(` + CREATE ROLE %[1]s LOGIN PASSWORD 'ss_reader_pw'; + GRANT CONNECT ON DATABASE ss_introspect_fk TO %[1]s; + GRANT USAGE ON SCHEMA public TO %[1]s; + GRANT SELECT ON ALL TABLES IN SCHEMA public TO %[1]s`, role)) + check(t, strings.Replace(dsn, "seedstorm:seedstorm@", role+":ss_reader_pw@", 1)) + }) +} + +// A foreign key may reference a UNIQUE column that is not the primary key, or +// a column of a composite key that is unique on its own. +// Seeding used to fill it with the parent's primary-key values, which the +// database refused as FK violations. +func TestSeed_ForeignKeysToNonKeyColumnsInsertValidReferences(t *testing.T) { + for _, e := range []engine{postgresEngine(), mysqlEngine()} { + t.Run(e.name, func(t *testing.T) { + dsn, conn := e.scratchDB(t, "ss_seed_refcols") + execSQL(t, conn, ` + CREATE TABLE accounts (id INT PRIMARY KEY, code VARCHAR(16) NOT NULL UNIQUE); + CREATE TABLE ledger (id INT PRIMARY KEY, account_code VARCHAR(16) NOT NULL, + FOREIGN KEY (account_code) REFERENCES accounts (code)); + CREATE TABLE memberships (tenant_id INT NOT NULL, member_no INT NOT NULL, PRIMARY KEY (tenant_id, member_no), UNIQUE (member_no)); + CREATE TABLE badges (id INT PRIMARY KEY, member_no INT NOT NULL, + FOREIGN KEY (member_no) REFERENCES memberships (member_no))`) + schemaPath := filepath.Join(t.TempDir(), "schema.yaml") + runBin(t, "introspect", "--db", e.name, "--dsn", dsn, "--out", schemaPath) + runBin(t, "seed", "--db", e.name, "--dsn", dsn, "--schema", schemaPath, "--rows", "40", "--workers", "1") + + var ledger, badges int + if err := conn.QueryRow(`SELECT COUNT(*) FROM ledger`).Scan(&ledger); err != nil { + t.Fatal(err) + } + if err := conn.QueryRow(`SELECT COUNT(*) FROM badges`).Scan(&badges); err != nil { + t.Fatal(err) + } + if ledger != 40 || badges != 40 { + t.Fatalf("rows: ledger=%d badges=%d, want 40 each (the database enforces every reference)", ledger, badges) + } + }) + } +} diff --git a/integration/loadsim_containers_test.go b/integration/loadsim_containers_test.go new file mode 100644 index 0000000..e57abb0 --- /dev/null +++ b/integration/loadsim_containers_test.go @@ -0,0 +1,143 @@ +//go:build integration && loadsim + +package integration_test + +import ( + "database/sql" + "fmt" + "net" + "os" + "os/exec" + "strconv" + "strings" + "testing" + "time" +) + +// loadsimDB is a throwaway database container limited like a profile. +type loadsimDB struct { + engine engine + dsn string + conn *sql.DB + name string // container name +} + +// requireHeadroom skips resource-limited tests on a machine that cannot spare +// the memory, instead of starving everything else running on it. +func requireHeadroom(t *testing.T, needMB int) { + t.Helper() + raw, err := os.ReadFile("/proc/meminfo") + if err != nil { + return + } + for _, line := range strings.Split(string(raw), "\n") { + if strings.HasPrefix(line, "MemAvailable:") { + fields := strings.Fields(line) + kb, _ := strconv.Atoi(fields[1]) + if kb/1024 < needMB { + t.Skipf("only %dMB available, need %dMB for this profile", kb/1024, needMB) + } + } + } +} + +func freePort(t *testing.T) int { + t.Helper() + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer l.Close() + return l.Addr().(*net.TCPAddr).Port +} + +// dataDevice is the block device holding Docker's data, for IOPS limits. +func dataDevice(t *testing.T) string { + t.Helper() + out, err := exec.Command("findmnt", "-no", "SOURCE", "-T", "/var/lib/docker").Output() + if err != nil { + return "" + } + src := strings.TrimSpace(string(out)) + if i := strings.Index(src, "["); i > 0 { + src = src[:i] + } + // /dev/nvme0n1p7 -> /dev/nvme0n1, /dev/sda1 -> /dev/sda + for _, prefix := range []string{"/dev/nvme", "/dev/mmcblk"} { + if strings.HasPrefix(src, prefix) { + if i := strings.LastIndex(src, "p"); i > len(prefix) { + return src[:i] + } + } + } + return strings.TrimRight(src, "0123456789") +} + +// startLoadsimDB starts a limited database container for profile and waits +// until it answers. The container is removed when the test ends. +func startLoadsimDB(t *testing.T, driver string, p loadsimProfile, limitIO bool) *loadsimDB { + t.Helper() + port := freePort(t) + name := fmt.Sprintf("ss-loadsim-%s-%s-%d", strings.ReplaceAll(p.name, "cloudsql-", ""), map[string]string{postgresDriver: "pg", mysqlDriver: "my"}[driver], port) + args := []string{"run", "-d", "--rm", "--name", name, "--cpus", p.cpus, "--memory", p.memory, "--memory-swap", p.memory, + "-p", fmt.Sprintf("127.0.0.1:%d:%d", port, map[string]int{postgresDriver: 5432, mysqlDriver: 3306}[driver])} + if dev := dataDevice(t); limitIO && p.writeIOPS > 0 && dev != "" { + args = append(args, "--device-write-iops", fmt.Sprintf("%s:%d", dev, p.writeIOPS)) + } + var e engine + var dsn string + switch driver { + case postgresDriver: + args = append(args, "-e", "POSTGRES_USER=seedstorm", "-e", "POSTGRES_PASSWORD=seedstorm", "-e", "POSTGRES_DB=loadsim", "postgres:17-alpine") + for _, setting := range p.postgres { + args = append(args, "-c", setting) + } + dsn = fmt.Sprintf("postgres://seedstorm:seedstorm@127.0.0.1:%d/loadsim?sslmode=disable", port) + e = postgresEngine() + default: + args = append(args, "-e", "MYSQL_ROOT_PASSWORD=root", "-e", "MYSQL_DATABASE=loadsim", "-e", "MYSQL_USER=seedstorm", "-e", "MYSQL_PASSWORD=seedstorm", "mysql:8.4") + args = append(args, p.mysql...) + dsn = fmt.Sprintf("seedstorm:seedstorm@tcp(127.0.0.1:%d)/loadsim?parseTime=true&multiStatements=true", port) + e = mysqlEngine() + } + if out, err := exec.Command("docker", args...).CombinedOutput(); err != nil { + t.Fatalf("docker run %s: %v\n%s", name, err, out) + } + t.Cleanup(func() { _ = exec.Command("docker", "rm", "-f", name).Run() }) + + deadline := time.Now().Add(3 * time.Minute) + for { + conn, err := sql.Open(driver, dsn) + if err == nil { + if err = conn.Ping(); err == nil { + t.Cleanup(func() { conn.Close() }) + return &loadsimDB{engine: e, dsn: dsn, conn: conn, name: name} + } + conn.Close() + } + if time.Now().After(deadline) { + logs, _ := exec.Command("docker", "logs", "--tail", "30", name).CombinedOutput() + t.Fatalf("%s did not accept connections: %v\n%s", name, err, logs) + } + time.Sleep(time.Second) + } +} + +// oomKilled reports whether the container was killed for memory, from the +// kernel's own counter (memory.peak sits at the limit on healthy runs: it +// counts page cache). +func (d *loadsimDB) oomKilled(t *testing.T) bool { + t.Helper() + out, err := exec.Command("docker", "exec", d.name, "cat", "/sys/fs/cgroup/memory.events").Output() + if err != nil { + state, _ := exec.Command("docker", "inspect", "-f", "{{.State.OOMKilled}}", d.name).Output() + return strings.TrimSpace(string(state)) == "true" + } + for _, line := range strings.Split(string(out), "\n") { + if strings.HasPrefix(line, "oom_kill ") { + n, _ := strconv.Atoi(strings.TrimSpace(strings.TrimPrefix(line, "oom_kill "))) + return n > 0 + } + } + return false +} diff --git a/integration/loadsim_profiles_test.go b/integration/loadsim_profiles_test.go new file mode 100644 index 0000000..4f47f09 --- /dev/null +++ b/integration/loadsim_profiles_test.go @@ -0,0 +1,55 @@ +//go:build integration && loadsim + +package integration_test + +// Database profiles modelled on managed Cloud SQL instances. Server settings +// that Cloud SQL derives from memory are set explicitly, never left to the +// container's defaults (docs/specs 004, D13). + +type loadsimProfile struct { + name string + cpus string // docker --cpus + memory string // docker --memory (and --memory-swap: no swap) + // writeIOPS limits writes to the data disk (0: unlimited). The Cloud SQL + // storage docs give 30 IOPS per GB of SSD. + writeIOPS int + postgres []string // postgres -c settings + mysql []string // mysqld flags + // source records where the settings come from. + source string +} + +var loadsimProfiles = map[string]loadsimProfile{ + // db-f1-micro: 1 shared vCPU, 628.74MB, 10GB SSD. + "cloudsql-micro": { + name: "cloudsql-micro", cpus: "1", memory: "629m", writeIOPS: 300, + postgres: []string{ + "max_connections=25", "shared_buffers=207MB", "effective_cache_size=251MB", + "work_mem=4MB", "maintenance_work_mem=64MB", "temp_buffers=8MB", + }, + mysql: []string{ + "--innodb-buffer-pool-size=53477376", "--innodb-flush-log-at-trx-commit=1", "--innodb-flush-method=O_DIRECT", + "--innodb-io-capacity=5000", "--innodb-io-capacity-max=10000", "--innodb-log-buffer-size=67108864", + "--innodb-redo-log-capacity=104857600", "--max-allowed-packet=33554432", "--max-connections=280", + "--performance-schema=OFF", "--table-open-cache=4000", "--thread-cache-size=10", + "--tmp-table-size=16777216", "--max-heap-table-size=16777216", "--sort-buffer-size=262144", "--join-buffer-size=262144", + }, + source: "postgres: Cloud SQL docs (flags, memory best practices); mysql: SHOW VARIABLES on a MySQL 8.4.10 Enterprise 1 vCPU / 628.74MB instance, 2026-09-17", + }, + // db-custom-2-7680: 2 vCPU, 7.5GB, 100GB SSD assumed. + "cloudsql-2vcpu": { + name: "cloudsql-2vcpu", cpus: "2", memory: "7680m", writeIOPS: 3000, + postgres: []string{ + "max_connections=400", "shared_buffers=2534MB", "effective_cache_size=3072MB", + "work_mem=4MB", "maintenance_work_mem=64MB", "temp_buffers=8MB", + }, + mysql: []string{ + "--innodb-buffer-pool-size=5793165312", "--innodb-flush-log-at-trx-commit=1", "--innodb-flush-method=O_DIRECT", + "--innodb-io-capacity=5000", "--innodb-io-capacity-max=10000", "--innodb-log-buffer-size=67108864", + "--innodb-redo-log-capacity=104857600", "--max-allowed-packet=33554432", "--max-connections=280", + "--performance-schema=OFF", "--table-open-cache=4000", "--thread-cache-size=10", + "--tmp-table-size=16777216", "--max-heap-table-size=16777216", "--sort-buffer-size=262144", "--join-buffer-size=262144", + }, + source: "postgres: Cloud SQL docs; mysql: ESTIMATED (buffer pool ~72% of memory per Cloud SQL docs, other values as the micro reference)", + }, +} diff --git a/integration/loadsim_readscope_test.go b/integration/loadsim_readscope_test.go new file mode 100644 index 0000000..24e8b99 --- /dev/null +++ b/integration/loadsim_readscope_test.go @@ -0,0 +1,46 @@ +//go:build integration && loadsim + +package integration_test + +import ( + "context" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/db" +) + +// Assertion 7: the read safeguards behave on the smallest managed instance as +// they do on an unconstrained server. +func TestLoadsim_ReadSafeguardsOnTheSmallestInstance(t *testing.T) { + p := loadsimProfiles["cloudsql-micro"] + for _, driver := range []string{postgresDriver, mysqlDriver} { + t.Run(driver, func(t *testing.T) { + requireHeadroom(t, 2048) + d := startLoadsimDB(t, driver, p, false) + execSQL(t, d.conn, `CREATE TABLE ledger (id INT PRIMARY KEY)`) + + err := db.ReadOnce(context.Background(), d.conn, driver, db.ReadLimits{}, func(ctx context.Context, q db.Querier) error { + _, err := q.ExecContext(ctx, `INSERT INTO ledger (id) VALUES (1)`) + return err + }) + if err == nil || countRows(t, d.conn, "ledger") != 0 { + t.Fatalf("write inside a read scope: err=%v rows=%d", err, countRows(t, d.conn, "ledger")) + } + + start := time.Now() + err = db.ReadOnce(context.Background(), d.conn, driver, db.ReadLimits{StatementTimeout: 500 * time.Millisecond}, func(ctx context.Context, q db.Querier) error { + return q.QueryRowContext(ctx, slowQuery(driver)).Scan(new(int64)) + }) + if got := db.ReadOutcomeOf(context.Background(), err); got != db.OutcomeTimedOut { + t.Fatalf("outcome = %s (%v), want timed out", got, err) + } + if d := time.Since(start); d > 10*time.Second { + t.Fatalf("timeout took %s on the micro profile", d) + } + if d.oomKilled(t) { + t.Fatal("the database was OOM-killed") + } + }) + } +} diff --git a/integration/partitions_test.go b/integration/partitions_test.go new file mode 100644 index 0000000..66442c8 --- /dev/null +++ b/integration/partitions_test.go @@ -0,0 +1,95 @@ +//go:build integration + +package integration_test + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/AxeForging/seedstorm/internal/compare" +) + +// Postgres partitioned tables were listed together with each of their +// partitions, so counts doubled, and seeding generated partition-key values +// outside every partition ("no partition of relation found for row"). +func TestPartitionedTables_CountedOnceAndSeededInsideTheirBounds(t *testing.T) { + e := postgresEngine() + dsn, conn := e.scratchDB(t, "ss_partitions") + execSQL(t, conn, ` + CREATE TABLE customers (id SERIAL PRIMARY KEY, name TEXT NOT NULL); + CREATE TABLE events (id BIGINT NOT NULL, customer_id INT NOT NULL REFERENCES customers (id), + created_at DATE NOT NULL, kind TEXT NOT NULL, PRIMARY KEY (id, created_at)) PARTITION BY RANGE (created_at); + CREATE TABLE events_2025 PARTITION OF events FOR VALUES FROM ('2025-01-01') TO ('2026-01-01'); + CREATE TABLE events_2026 PARTITION OF events FOR VALUES FROM ('2026-01-01') TO ('2027-01-01'); + CREATE TABLE tickets (id INT NOT NULL, region TEXT NOT NULL, PRIMARY KEY (id, region)) PARTITION BY LIST (region); + CREATE TABLE tickets_eu PARTITION OF tickets FOR VALUES IN ('eu-west', 'eu-north'); + CREATE TABLE tickets_us PARTITION OF tickets FOR VALUES IN ('us-east'); + CREATE TABLE scores (id INT NOT NULL, bucket INT NOT NULL, PRIMARY KEY (id, bucket)) PARTITION BY RANGE (bucket); + CREATE TABLE scores_low PARTITION OF scores FOR VALUES FROM (0) TO (100); + CREATE TABLE scores_rest PARTITION OF scores DEFAULT; + CREATE TABLE shards (id INT PRIMARY KEY, body TEXT) PARTITION BY HASH (id); + CREATE TABLE shards_0 PARTITION OF shards FOR VALUES WITH (modulus 2, remainder 0); + CREATE TABLE shards_1 PARTITION OF shards FOR VALUES WITH (modulus 2, remainder 1)`) + + dir := t.TempDir() + schemaPath := filepath.Join(dir, "schema.yaml") + runBin(t, "introspect", "--db", "postgres", "--dsn", dsn, "--out", schemaPath) + raw, err := os.ReadFile(schemaPath) + if err != nil { + t.Fatal(err) + } + for _, partition := range []string{"events_2025", "events_2026", "tickets_eu", "tickets_us", "scores_low", "scores_rest", "shards_0", "shards_1"} { + if strings.Contains(string(raw), "\n "+partition+":") { + t.Errorf("partition %s is introspected as its own table", partition) + } + } + + runBin(t, "seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "60", "--workers", "1") + for _, table := range []string{"events", "tickets", "scores", "shards"} { + if n := countRows(t, conn, table); n < 60 { + t.Errorf("%s has %d rows after seeding 60", table, n) + } + } + if n := scalar(t, conn, `SELECT COUNT(*) FROM events_2025`) + scalar(t, conn, `SELECT COUNT(*) FROM events_2026`); n != countRows(t, conn, "events") { + t.Errorf("events rows outside its partitions: partitions hold %d of %d", n, countRows(t, conn, "events")) + } + + snap := decodeJSON[compare.Snapshot](t, runBin(t, "snapshot", "--db", "postgres", "--dsn", dsn, "--format", "json")) + for _, partition := range []string{"events_2025", "tickets_eu", "shards_0"} { + if _, ok := snap.Tables[partition]; ok { + t.Errorf("snapshot lists partition %s: parent rows would be counted twice", partition) + } + } + if got, want := snap.Tables["events"].Rows, int64(countRows(t, conn, "events")); got != want { + t.Errorf("snapshot events rows = %d, want %d", got, want) + } + if snap.Tables["events"].Bytes <= 0 { + t.Errorf("snapshot events bytes = %d, want the size of its partitions", snap.Tables["events"].Bytes) + } +} + +// A partition key seedstorm cannot generate inside the bounds (an expression) +// is refused before anything is written, with a way out. +func TestPartitionedTables_ExpressionKeyIsRefusedBeforeWriting(t *testing.T) { + e := postgresEngine() + dsn, conn := e.scratchDB(t, "ss_partitions_expr") + execSQL(t, conn, ` + CREATE TABLE plain (id INT PRIMARY KEY); + INSERT INTO plain VALUES (1), (2), (3); + CREATE TABLE logs (id INT NOT NULL, created_at TIMESTAMP NOT NULL) PARTITION BY RANGE (date_trunc('month', created_at)); + CREATE TABLE logs_jan PARTITION OF logs FOR VALUES FROM ('2026-01-01') TO ('2026-02-01')`) + schemaPath := filepath.Join(t.TempDir(), "schema.yaml") + runBin(t, "introspect", "--db", "postgres", "--dsn", dsn, "--out", schemaPath) + _, stderr, err := runBinResult(t, "seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "5", "--workers", "1", "--truncate", "--yes") + if err == nil { + t.Fatal("seed succeeded on a table partitioned by an expression") + } + if !strings.Contains(stderr, "logs") || !strings.Contains(stderr, "partitioned by an expression") || !strings.Contains(stderr, "value rule") { + t.Fatalf("stderr does not explain the refusal:\n%s", stderr) + } + if n := countRows(t, conn, "plain"); n != 3 { + t.Fatalf("plain has %d rows, want its 3: the run truncated or wrote before refusing", n) + } +} diff --git a/integration/production_test.go b/integration/production_test.go new file mode 100644 index 0000000..e7cdc47 --- /dev/null +++ b/integration/production_test.go @@ -0,0 +1,58 @@ +//go:build integration + +package integration_test + +import ( + "path/filepath" + "strings" + "testing" +) + +// A database marked production (--production or SEEDSTORM_PRODUCTION) is never +// written without --allow-production: the command stops before truncating or +// inserting anything. Dry runs and reads still work. +func TestProduction_CLIRefusesWritesUnlessAllowed(t *testing.T) { + e := postgresEngine() + dsn, conn := e.scratchDB(t, "ss_production_cli") + execSQL(t, conn, `CREATE TABLE accounts (id INT PRIMARY KEY, name TEXT); INSERT INTO accounts VALUES (1, 'kept')`) + schemaPath := filepath.Join(t.TempDir(), "schema.yaml") + runBin(t, "introspect", "--db", "postgres", "--dsn", dsn, "--out", schemaPath) + + refused := []struct { + name string + args []string + env bool + }{ + {"seed --truncate", []string{"seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "5", "--truncate", "--yes", "--production"}, false}, + {"seed via env", []string{"seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "5"}, true}, + {"gaps --fill", []string{"gaps", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--fill", "--yes", "--production"}, false}, + {"clone-schema", []string{"clone-schema", "--source-db", "postgres", "--source-dsn", dsn, "--target-db", "postgres", "--target-dsn", dsn, "--production"}, false}, + } + for _, c := range refused { + t.Run(c.name, func(t *testing.T) { + if c.env { + t.Setenv("SEEDSTORM_PRODUCTION", "true") + } + _, stderr, err := runBinResult(t, c.args...) + if err == nil { + t.Fatal("the write to a production database ran") + } + if !strings.Contains(stderr, "production") || !strings.Contains(stderr, "--allow-production") { + t.Fatalf("stderr does not explain the refusal:\n%s", stderr) + } + if n := countRows(t, conn, "accounts"); n != 1 { + t.Fatalf("accounts has %d rows, want the 1 it had", n) + } + }) + } + + t.Run("dry run is allowed", func(t *testing.T) { + runBin(t, "seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "5", "--dry-run", "--production") + }) + t.Run("allowed explicitly", func(t *testing.T) { + runBin(t, "seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "5", "--production", "--allow-production") + if n := countRows(t, conn, "accounts"); n != 6 { + t.Fatalf("accounts has %d rows after an allowed seed, want 6", n) + } + }) +} diff --git a/integration/read_scope_test.go b/integration/read_scope_test.go new file mode 100644 index 0000000..910182c --- /dev/null +++ b/integration/read_scope_test.go @@ -0,0 +1,200 @@ +//go:build integration + +package integration_test + +import ( + "context" + "database/sql" + "strconv" + "strings" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/db" +) + +// slowQuery is a read that runs for many seconds on either engine without +// sleeping (MySQL's SLEEP swallows interrupts instead of failing). +func slowQuery(driver string) string { + if driver == postgresDriver { + return `SELECT COUNT(*) FROM generate_series(1, 400000000)` + } + return `SELECT COUNT(*) FROM information_schema.COLUMNS a, information_schema.COLUMNS b, information_schema.COLUMNS c` +} + +// Every read seedstorm runs against a database it only reads is refused by the +// server if it tries to write, whatever the code does. +func TestReadScope_ServerRefusesWrites(t *testing.T) { + for _, e := range engines() { + t.Run(e.name, func(t *testing.T) { + _, conn := e.scratchDB(t, "ss_readscope_ro") + execSQL(t, conn, `CREATE TABLE notes (id INT PRIMARY KEY)`) + err := db.ReadOnce(context.Background(), conn, e.driver, db.ReadLimits{}, func(ctx context.Context, q db.Querier) error { + _, err := q.ExecContext(ctx, `INSERT INTO notes (id) VALUES (1)`) + return err + }) + if err == nil { + t.Fatal("an INSERT inside a read scope succeeded") + } + if n := countRows(t, conn, "notes"); n != 0 { + t.Fatalf("notes has %d rows", n) + } + }) + } +} + +// A statement timeout is enforced by the server and reported as timed out; +// the connection is usable afterwards and the limit does not leak into it. +func TestReadScope_StatementTimeoutIsServerSide(t *testing.T) { + for _, e := range engines() { + t.Run(e.name, func(t *testing.T) { + _, conn := e.scratchDB(t, "ss_readscope_timeout") + conn.SetMaxOpenConns(1) // the next read gets the same server session + start := time.Now() + err := db.ReadOnce(context.Background(), conn, e.driver, db.ReadLimits{StatementTimeout: 500 * time.Millisecond}, func(ctx context.Context, q db.Querier) error { + var n int64 + return q.QueryRowContext(ctx, slowQuery(e.driver)).Scan(&n) + }) + if got := db.ReadOutcomeOf(context.Background(), err); got != db.OutcomeTimedOut { + t.Fatalf("outcome = %s (err %v), want timed out", got, err) + } + if d := time.Since(start); d > 5*time.Second { + t.Fatalf("timed-out read took %s", d) + } + // The same session runs an ordinary read without inheriting the limit. + var one int + if err := conn.QueryRowContext(context.Background(), `SELECT 1`).Scan(&one); err != nil || one != 1 { + t.Fatalf("connection unusable after a timeout: %v", err) + } + if e.driver == mysqlDriver { + var v int64 + if err := conn.QueryRowContext(context.Background(), `SELECT @@SESSION.max_execution_time`).Scan(&v); err != nil || v != 0 { + t.Fatalf("max_execution_time leaked into the pool: %d (%v)", v, err) + } + } + }) + } +} + +// A read waiting behind a table lock (an ALTER, a LOCK TABLE) gives up after +// the lock timeout instead of queueing and holding up every writer behind it. +func TestReadScope_LockTimeoutStopsWaitingBehindALock(t *testing.T) { + for _, e := range engines() { + t.Run(e.name, func(t *testing.T) { + _, conn := e.scratchDB(t, "ss_readscope_lock") + execSQL(t, conn, `CREATE TABLE ledger (id INT PRIMARY KEY)`) + holder, err := conn.Conn(context.Background()) + if err != nil { + t.Fatal(err) + } + defer holder.Close() + lock := `LOCK TABLES ledger WRITE` + unlock := `UNLOCK TABLES` + if e.driver == postgresDriver { + if _, err := holder.ExecContext(context.Background(), `BEGIN`); err != nil { + t.Fatal(err) + } + lock, unlock = `LOCK TABLE ledger IN ACCESS EXCLUSIVE MODE`, `ROLLBACK` + } + if _, err := holder.ExecContext(context.Background(), lock); err != nil { + t.Fatal(err) + } + defer holder.ExecContext(context.Background(), unlock) //nolint:errcheck + + start := time.Now() + counts, failed := db.CountTablesWithin(context.Background(), conn, e.driver, []string{"ledger"}, db.ReadLimits{LockTimeout: time.Second}, nil) + if _, ok := counts["ledger"]; ok { + t.Fatal("ledger counted while an exclusive lock was held") + } + if got := db.ReadOutcomeOf(context.Background(), failed["ledger"]); got != db.OutcomeLocked { + t.Fatalf("outcome = %s (%v), want locked", got, failed["ledger"]) + } + if d := time.Since(start); d > 5*time.Second { + t.Fatalf("locked read waited %s", d) + } + }) + } +} + +// Cancelling a MySQL read must stop the query on the server. The driver only +// closes its socket, and the server keeps running the query: this test first +// shows that, then that ReadOnce kills it. +func TestReadScope_MySQLCancelKillsTheServerQuery(t *testing.T) { + e := mysqlEngine() + _, conn := e.scratchDB(t, "ss_readscope_kill") + running := func(marker string) bool { + var n int + _ = conn.QueryRowContext(context.Background(), + `SELECT COUNT(*) FROM information_schema.PROCESSLIST WHERE INFO LIKE ? AND INFO NOT LIKE '%PROCESSLIST%'`, "%"+marker+"%").Scan(&n) + return n > 0 + } + query := func(marker string) string { + return "SELECT /* " + marker + " */ " + strings.TrimPrefix(slowQuery(mysqlDriver), "SELECT ") + } + waitRunning := func(marker string) { + deadline := time.Now().Add(5 * time.Second) + for !running(marker) { + if time.Now().After(deadline) { + t.Fatalf("query %s never started", marker) + } + time.Sleep(50 * time.Millisecond) + } + } + + t.Run("driver alone leaves the query running", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- conn.QueryRowContext(ctx, query("ss_plain_cancel")).Scan(new(int64)) }() + waitRunning("ss_plain_cancel") + cancel() + <-done + time.Sleep(500 * time.Millisecond) + if !running("ss_plain_cancel") { + t.Skip("this server stopped the query on disconnect; the KILL below is still required elsewhere") + } + killMarker(t, conn, "ss_plain_cancel") + }) + + t.Run("ReadOnce kills it", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- db.ReadOnce(ctx, conn, mysqlDriver, db.ReadLimits{}, func(ctx context.Context, q db.Querier) error { + return q.QueryRowContext(ctx, query("ss_scope_cancel")).Scan(new(int64)) + }) + }() + waitRunning("ss_scope_cancel") + cancel() + err := <-done + if got := db.ReadOutcomeOf(ctx, err); got != db.OutcomeCancelled { + t.Fatalf("outcome = %s (%v), want cancelled", got, err) + } + deadline := time.Now().Add(3 * time.Second) + for running("ss_scope_cancel") { + if time.Now().After(deadline) { + killMarker(t, conn, "ss_scope_cancel") + t.Fatal("the query still runs on the server 3s after cancelling") + } + time.Sleep(50 * time.Millisecond) + } + }) +} + +func killMarker(t *testing.T, conn *sql.DB, marker string) { + t.Helper() + rows, err := conn.QueryContext(context.Background(), `SELECT ID FROM information_schema.PROCESSLIST WHERE INFO LIKE ? AND INFO NOT LIKE '%PROCESSLIST%'`, "%"+marker+"%") + if err != nil { + return + } + var ids []int64 + for rows.Next() { + var id int64 + if rows.Scan(&id) == nil { + ids = append(ids, id) + } + } + rows.Close() + for _, id := range ids { + _, _ = conn.ExecContext(context.Background(), `KILL QUERY `+strconv.FormatInt(id, 10)) + } +} diff --git a/integration/reliability_test.go b/integration/reliability_test.go new file mode 100644 index 0000000..7cacdf4 --- /dev/null +++ b/integration/reliability_test.go @@ -0,0 +1,136 @@ +//go:build integration + +package integration_test + +import ( + "bytes" + "errors" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "syscall" + "testing" + "time" +) + +var ( + faultBinOnce sync.Once + faultBinPath string + faultBinErr error +) + +// seedstormFaultBin builds the CLI with the faultinject tag: SEEDSTORM_FAULT +// then makes a named step panic, fail or hang, so tests exercise real failure +// paths without mocking seedstorm's own code. +func seedstormFaultBin(t *testing.T) string { + t.Helper() + faultBinOnce.Do(func() { + dir, err := os.MkdirTemp("", "seedstorm-faultbin-") + if err != nil { + faultBinErr = err + return + } + faultBinPath = filepath.Join(dir, "seedstorm") + out, err := exec.Command("go", "build", "-tags", "faultinject", "-o", faultBinPath, "../cmd/seedstorm").CombinedOutput() + if err != nil { + faultBinErr = errors.New(string(out)) + } + }) + if faultBinErr != nil { + t.Fatalf("build fault-injection binary: %v", faultBinErr) + } + return faultBinPath +} + +func runFault(t *testing.T, fault string, args ...string) (stderr string, code int) { + t.Helper() + cmd := exec.Command(seedstormFaultBin(t), append([]string{"--no-color"}, args...)...) + cmd.Env = append(os.Environ(), "SEEDSTORM_FAULT="+fault) + var errBuf bytes.Buffer + cmd.Stderr = &errBuf + err := cmd.Run() + var exit *exec.ExitError + if errors.As(err, &exit) { + return errBuf.String(), exit.ExitCode() + } + if err != nil { + t.Fatal(err) + } + return errBuf.String(), 0 +} + +// A panic is an internal error (exit 70) that says where it happened; a +// refused run is exit 1; neither prints a Go stack unless asked. +func TestCLI_ExitCodesSayWhatKindOfFailureItWas(t *testing.T) { + e := postgresEngine() + dsn, conn := e.scratchDB(t, "ss_reliability_cli") + execSQL(t, conn, `CREATE TABLE users (id INT PRIMARY KEY, name TEXT); CREATE TABLE posts (id INT PRIMARY KEY, user_id INT NOT NULL REFERENCES users (id))`) + schemaPath := filepath.Join(t.TempDir(), "schema.yaml") + runBin(t, "introspect", "--db", "postgres", "--dsn", dsn, "--out", schemaPath) + seed := []string{"seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "50", "--workers", "2"} + + stderr, code := runFault(t, "write:users:panic", seed...) + if code != 70 { + t.Fatalf("panic while writing: exit %d, want 70\n%s", code, stderr) + } + for _, want := range []string{"internal error", "write", "users", "--log-level debug"} { + if !strings.Contains(stderr, want) { + t.Fatalf("stderr lacks %q:\n%s", want, stderr) + } + } + if strings.Contains(stderr, "goroutine ") { + t.Fatalf("a stack was printed without --log-level debug:\n%s", stderr) + } + if n := countRows(t, conn, "posts"); n != 0 { + t.Fatalf("posts has %d rows after its parent's writes panicked", n) + } + + stderr, code = runFault(t, "write:users:error", seed...) + if code != 1 || !strings.Contains(stderr, "write · users") { + t.Fatalf("failed write: exit %d, want 1 naming write · users\n%s", code, stderr) + } +} + +// Ctrl+C ends a run with exit 130 and a line saying so, promptly. +func TestCLI_InterruptExitsWith130(t *testing.T) { + e := postgresEngine() + dsn, conn := e.scratchDB(t, "ss_reliability_sigint") + execSQL(t, conn, `CREATE TABLE users (id INT PRIMARY KEY, name TEXT)`) + schemaPath := filepath.Join(t.TempDir(), "schema.yaml") + runBin(t, "introspect", "--db", "postgres", "--dsn", dsn, "--out", schemaPath) + + cmd := exec.Command(seedstormFaultBin(t), "--no-color", "seed", "--db", "postgres", "--dsn", dsn, "--schema", schemaPath, "--rows", "1000", "--workers", "1") + cmd.Env = append(os.Environ(), "SEEDSTORM_FAULT=write:users:hang") + var errBuf bytes.Buffer + cmd.Stderr = &errBuf + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + time.Sleep(time.Second) + start := time.Now() + _ = cmd.Process.Signal(syscall.SIGINT) + err := cmd.Wait() + var exit *exec.ExitError + if !errors.As(err, &exit) || exit.ExitCode() != 130 { + t.Fatalf("exit = %v, want 130\n%s", err, errBuf.String()) + } + if d := time.Since(start); d > 5*time.Second { + t.Fatalf("took %s to stop after Ctrl+C", d) + } + if !strings.Contains(errBuf.String(), "interrupted") { + t.Fatalf("stderr does not say it was interrupted:\n%s", errBuf.String()) + } +} + +// The normal build carries no fault-injection code. +func TestFaultInjection_AbsentFromTheDefaultBinary(t *testing.T) { + raw, err := os.ReadFile(seedstormBin(t)) + if err != nil { + t.Fatal(err) + } + if bytes.Contains(raw, []byte("SEEDSTORM_FAULT")) { + t.Fatal("the default binary contains fault-injection code") + } +} diff --git a/integration/reliability_web_test.go b/integration/reliability_web_test.go new file mode 100644 index 0000000..13dcd6e --- /dev/null +++ b/integration/reliability_web_test.go @@ -0,0 +1,105 @@ +//go:build integration + +package integration_test + +import ( + "bytes" + "net" + "net/http" + "os" + "os/exec" + "strconv" + "strings" + "sync" + "testing" + "time" +) + +// startFaultServe runs `seedstorm serve` from the fault-injection build and +// returns its base URL; the process is stopped when the test ends. +// lockedBuffer is written by the child process's output copier while the test +// reads it. +type lockedBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *lockedBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *lockedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +func startFaultServe(t *testing.T, fault string) (string, *lockedBuffer) { + t.Helper() + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + port := l.Addr().(*net.TCPAddr).Port + _ = l.Close() + cmd := exec.Command(seedstormFaultBin(t), "--no-color", "serve", "--addr", "127.0.0.1:"+strconv.Itoa(port)) + cmd.Env = append(os.Environ(), "SEEDSTORM_FAULT="+fault, "XDG_CONFIG_HOME="+t.TempDir(), "HOME="+t.TempDir()) + logs := &lockedBuffer{} + cmd.Stdout, cmd.Stderr = logs, logs + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = cmd.Process.Kill(); _, _ = cmd.Process.Wait() }) + base := "http://127.0.0.1:" + strconv.Itoa(port) + deadline := time.Now().Add(15 * time.Second) + for { + if res, err := http.Get(base + "/connect"); err == nil { + res.Body.Close() + return base, logs + } + if time.Now().After(deadline) { + t.Fatalf("serve did not start:\n%s", logs.String()) + } + time.Sleep(100 * time.Millisecond) + } +} + +// A panic inside one web job used to kill `serve`: every session and every +// running job went with it. The job fails naming where; another job running +// at the same time finishes, and the server keeps answering. +func TestServe_PanicInOneJobLeavesTheServerAndOtherJobsRunning(t *testing.T) { + e := postgresEngine() + _, broken := e.scratchDB(t, "ss_reliability_web_a") + _, healthy := e.scratchDB(t, "ss_reliability_web_b") + execSQL(t, broken, `CREATE TABLE ledger_entries (id INT PRIMARY KEY, note TEXT)`) + execSQL(t, healthy, `CREATE TABLE accounts (id INT PRIMARY KEY, name TEXT)`) + + base, logs := startFaultServe(t, "write:ledger_entries:panic") + a, b := clientFor(t, base), clientFor(t, base) + a.connectPostgres("ss_reliability_web_a") + b.connectPostgres("ss_reliability_web_b") + + slow := b.start("/api/seed", map[string]any{"rows": 20000, "batchSize": 500, "workers": 2}) + failing := a.start("/api/seed", map[string]any{"rows": 200, "workers": 2}) + + _, status, reason := a.stream(failing).finish(t, 60*time.Second) + if status != "failed" || !strings.Contains(reason, "write · ledger_entries") || !strings.Contains(reason, "internal error") { + t.Fatalf("failing job ended %q: %q", status, reason) + } + if _, status, reason := b.stream(slow).finish(t, 120*time.Second); status != "done" { + t.Fatalf("the other job ended %q (%q) after a panic elsewhere\n%s", status, reason, logs.String()) + } + if n := countRows(t, healthy, "accounts"); n != 20000 { + t.Fatalf("accounts = %d rows, want 20000", n) + } + res, err := http.Get(base + "/connect") + if err != nil || res.StatusCode != http.StatusOK { + t.Fatalf("server no longer answers after the panic: %v %v\n%s", res, err, logs.String()) + } + res.Body.Close() + if strings.Contains(logs.String(), "goroutine ") { + t.Fatalf("serve printed a Go stack at the default log level:\n%s", logs.String()) + } +} diff --git a/integration/server_info_test.go b/integration/server_info_test.go new file mode 100644 index 0000000..983ad81 --- /dev/null +++ b/integration/server_info_test.go @@ -0,0 +1,58 @@ +//go:build integration + +package integration_test + +import ( + "context" + "fmt" + "strings" + "testing" + + "github.com/AxeForging/seedstorm/internal/db" +) + +// Tuning reads what a database says about its capacity: connection limit and +// use, reserved buffers, the database's size, version, replica or not. +func TestDetectServer_ReadsCapacityOnBothEngines(t *testing.T) { + for _, e := range engines() { + t.Run(e.name, func(t *testing.T) { + _, conn := e.scratchDB(t, "ss_server_info") + execSQL(t, conn, `CREATE TABLE sized (id INT PRIMARY KEY, body VARCHAR(200))`) + for i := 0; i < 50; i++ { + execSQL(t, conn, fmt.Sprintf(`INSERT INTO sized (id, body) VALUES (%d, 'row')`, i)) + } + info, err := db.DetectServer(context.Background(), conn, e.driver) + if err != nil { + t.Fatal(err) + } + if info.MaxConnections <= 0 || info.UsedConnections <= 0 || info.UsedConnections > info.MaxConnections { + t.Fatalf("connections %d/%d", info.UsedConnections, info.MaxConnections) + } + if info.UsedBytes <= 0 || info.Version == "" || info.ServerID == "" { + t.Fatalf("info = %+v", info) + } + if info.Replica { + t.Fatal("a primary reported as a replica") + } + switch e.driver { + case postgresDriver: + if info.SharedBuffersBytes <= 0 { + t.Fatalf("shared_buffers = %d", info.SharedBuffersBytes) + } + default: + if info.BufferPoolBytes <= 0 || info.MaxAllowedPacket <= 0 { + t.Fatalf("mysql buffers = %+v", info) + } + } + // Two databases on the same server share the server id. + _, other := e.scratchDB(t, "ss_server_info_other") + otherInfo, err := db.DetectServer(context.Background(), other, e.driver) + if err != nil { + t.Fatal(err) + } + if otherInfo.ServerID != info.ServerID || strings.Contains(info.ServerID, "ss_server_info") { + t.Fatalf("server ids %q vs %q: must be equal and not name the database", info.ServerID, otherInfo.ServerID) + } + }) + } +} diff --git a/integration/web_jobs_helpers_test.go b/integration/web_jobs_helpers_test.go index dd32569..1c840df 100644 --- a/integration/web_jobs_helpers_test.go +++ b/integration/web_jobs_helpers_test.go @@ -83,7 +83,7 @@ func (c *webClient) start(path string, body any) string { // streamEvent is one SSE event of a job stream. type streamEvent struct { - Kind string // log | phase | progress | status | error | end + Kind string // log | phase | progress | status | failure | end Seq int Text string // log line, phase name, or progress label Done int @@ -175,7 +175,7 @@ func (s *jobStream) read(res *http.Response) { switch event { case "status": s.status = ev.Text - case "error": + case "failure": s.errMsg = ev.Text case "end": s.ended = true diff --git a/internal/cli/clone_schema.go b/internal/cli/clone_schema.go index 0ae372f..003e8fb 100644 --- a/internal/cli/clone_schema.go +++ b/internal/cli/clone_schema.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "strings" + "time" "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/logging" @@ -72,11 +73,18 @@ This is same-engine schema cloning for local/test databases, not a lossless migr Aliases: []string{"i"}, Usage: "Review and confirm the clone in the terminal UI", }, + productionFlags()[0], + productionFlags()[1], }, Action: func(ctx context.Context, cmd *cli.Command) error { log := logging.Log sourceType := normalizeDBType(cmd.String("source-db")) targetType := normalizeDBType(cmd.String("target-db")) + if !cmd.Bool("dry-run") { + if err := refuseProductionWrite(cmd, "clone a schema into the target"); err != nil { + return err + } + } objects, err := db.ParseCloneObjects(cmd.String("objects")) if err != nil { return err @@ -92,10 +100,14 @@ This is same-engine schema cloning for local/test databases, not a lossless migr if cmd.Bool("interactive") { return tui.RunClone(ctx, sourceType, cmd.String("source-dsn"), targetType, cmd.String("target-dsn"), opts) } + start := time.Now() + log.Info().Str("source", dsnLabel(sourceType, cmd.String("source-dsn"))).Str("target", dsnLabel(targetType, cmd.String("target-dsn"))).Bool("dry_run", opts.DryRun). + Msg("Cloning schema: reading the source, then running DDL on the target") result, err := db.CloneSchema(ctx, sourceType, cmd.String("source-dsn"), targetType, cmd.String("target-dsn"), opts) if err != nil { return err } + log.Info().Int("tables", result.Tables).Int("statements", len(result.Statements)).Dur("duration", time.Since(start).Round(time.Millisecond)).Msg("Schema cloned") for _, skipped := range result.Skipped { log.Warn().Str("kind", string(skipped.Kind)).Str("name", skipped.Name).Str("reason", skipped.Reason).Msg("Object not cloned") } diff --git a/internal/cli/compare.go b/internal/cli/compare.go index a3e77d0..3bdcc35 100644 --- a/internal/cli/compare.go +++ b/internal/cli/compare.go @@ -4,10 +4,12 @@ import ( "context" "fmt" "os" + "time" "github.com/urfave/cli/v3" "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/logging" "github.com/AxeForging/seedstorm/internal/seeder" ) @@ -30,13 +32,15 @@ instead of a live database.`, if err != nil { return err } + logging.Log.Info().Msg("Connecting to source and target") source, target, err := openEndpoints(ctx, cmd) if err != nil { return err } defer closeEndpoints(source, target) - report, err := seeder.Snapshots(ctx, source, target, mode, nil) + logging.Log.Info().Str("source", source.Label).Str("target", target.Label).Str("counts", string(mode)).Msg("Reading table volumes") + report, err := seeder.Snapshots(ctx, source, target, mode, sideStepLogger("Counting", time.Now)) if err != nil { return err } diff --git a/internal/cli/endpoints.go b/internal/cli/endpoints.go index 6ec01e6..de3cecf 100644 --- a/internal/cli/endpoints.go +++ b/internal/cli/endpoints.go @@ -8,10 +8,12 @@ import ( "os" "path/filepath" "strings" + "time" "github.com/urfave/cli/v3" "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/seeder" ) @@ -37,11 +39,11 @@ func openEndpoints(ctx context.Context, cmd *cli.Command) (source, target seeder return source, target, fmt.Errorf("use either --source-dsn or --source-snapshot, not both") case snapshotPath != "": if source, err = snapshotEndpoint(snapshotPath); err != nil { - return source, target, fmt.Errorf("source snapshot: %w", err) + return source, target, runerr.OnSide(runerr.SideSource, runerr.At(runerr.PhaseConnect, "", fmt.Errorf("snapshot file: %w", err))) } case sourceDSN != "": if source, err = openEndpoint(ctx, cmd.String("source-db"), sourceDSN); err != nil { - return source, target, fmt.Errorf("source: %w", err) + return source, target, runerr.OnSide(runerr.SideSource, runerr.At(runerr.PhaseConnect, "", err)) } default: return source, target, fmt.Errorf("a source is required: pass --source-dsn (or SEEDSTORM_SOURCE_DSN) or --source-snapshot ") @@ -49,7 +51,7 @@ func openEndpoints(ctx context.Context, cmd *cli.Command) (source, target seeder target, err = openEndpoint(ctx, cmd.String("target-db"), cmd.String("target-dsn")) if err != nil { closeEndpoints(source) - return source, target, fmt.Errorf("target: %w", err) + return source, target, runerr.OnSide(runerr.SideTarget, runerr.At(runerr.PhaseConnect, "", err)) } return source, target, nil } @@ -89,13 +91,24 @@ func openEndpoint(ctx context.Context, dbFlag, dsn string) (seeder.Endpoint, err if err != nil { return seeder.Endpoint{}, fmt.Errorf("open connection: %w", err) } - if err := conn.PingContext(ctx); err != nil { + if err := pingWithin(ctx, conn); err != nil { _ = conn.Close() - return seeder.Endpoint{}, fmt.Errorf("ping database: %w", err) + return seeder.Endpoint{}, fmt.Errorf("%s did not answer: %w", dsnLabel(driver, dsn), err) } return seeder.Endpoint{Conn: conn, DBType: driver, DSN: dsn, Label: dsnLabel(driver, dsn)}, nil } +// connectTimeout bounds how long a database may take to answer before a +// command reports it unreachable instead of waiting on the OS TCP timeout. +var connectTimeout = 10 * time.Second + +// pingWithin pings conn, giving up after connectTimeout. +func pingWithin(ctx context.Context, conn *sql.DB) error { + pctx, cancel := context.WithTimeout(ctx, connectTimeout) + defer cancel() + return conn.PingContext(pctx) +} + // dsnLabel names a connection for reports without leaking its password. func dsnLabel(driver, dsn string) string { if driver == "pgx" { diff --git a/internal/cli/gaps.go b/internal/cli/gaps.go index ddf208a..7d92fb1 100644 --- a/internal/cli/gaps.go +++ b/internal/cli/gaps.go @@ -13,6 +13,7 @@ import ( "github.com/AxeForging/seedstorm/internal/faker" "github.com/AxeForging/seedstorm/internal/graph" "github.com/AxeForging/seedstorm/internal/logging" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/schema" "github.com/AxeForging/seedstorm/internal/seeder" "github.com/AxeForging/seedstorm/internal/tui" @@ -93,6 +94,8 @@ Use --fill --dry-run to preview the SQL without executing it.`, }, workersFlag(), genWorkersFlag(), + productionFlags()[0], + productionFlags()[1], profileFlag(), }, Action: func(ctx context.Context, cmd *cli.Command) error { @@ -111,6 +114,11 @@ Use --fill --dry-run to preview the SQL without executing it.`, dryRun := cmd.Bool("dry-run") yes := cmd.Bool("yes") batchSize := cmd.Int("batch-size") + if fill && !dryRun { + if err := refuseProductionWrite(cmd, "fill its empty tables"); err != nil { + return err + } + } log.Info().Str("path", schemaPath).Msg("Loading schema") s, err := schema.Load(schemaPath) @@ -133,15 +141,20 @@ Use --fill --dry-run to preview the SQL without executing it.`, return fmt.Errorf("failed to open connection: %w", err) } defer dbConn.Close() - if err := dbConn.PingContext(ctx); err != nil { - return fmt.Errorf("failed to ping database: %w", err) + if err := pingWithin(ctx, dbConn); err != nil { + return runerr.At(runerr.PhaseConnect, "", fmt.Errorf("%s did not answer: %w", dsnLabel(dbType, dsn), err)) } // Query current row counts for all tables. log.Info().Int("tables", len(allSorted)).Msg("Scanning tables") - counts, err := db.GetTableRowCounts(ctx, dbConn, dbType, allSorted) - if err != nil { - return fmt.Errorf("row count scan failed: %w", err) + counts, failed := db.CountTables(ctx, dbConn, dbType, allSorted, stepLogger("Counting rows", time.Now)) + if err := ctx.Err(); err != nil { + return err + } + for _, t := range allSorted { + if ferr, ok := failed[t]; ok { + log.Warn().Str("table", t).Err(ferr).Msg("Row count failed: the table is left out of the gaps (unknown is not empty)") + } } profile, err := loadProfile(cmd, s) @@ -158,12 +171,7 @@ Use --fill --dry-run to preview the SQL without executing it.`, fkParents := buildFKParents(s, allSorted) // Identify gap tables in topological order, minus the profile's ignored ones. - var gapTables []string - for _, t := range allSorted { - if counts[t] == 0 { - gapTables = append(gapTables, t) - } - } + gapTables := seeder.GapTables(allSorted, counts, nil) if gapTables, err = profile.applyIgnore(ctx, dbConn, dbType, gapTables); err != nil { return err } @@ -211,19 +219,23 @@ Use --fill --dry-run to preview the SQL without executing it.`, res, err := seeder.Seed(ctx, dbConn, dbType, s, allSorted, gapTables, seeder.SeedOptions{ Rows: rows, EnumRows: enumRows, TableRows: tableRows, BatchSize: batchSize, DryRun: dryRun, Workers: cmd.Int("workers"), OnProgress: onProgress, OnTable: onTable, - GenWorkers: cmd.Int("gen-workers"), + GenWorkers: genWorkers(cmd), Generate: faker.GenerateOptions{ SelfRefDepth: selfRefDepth, Overrides: profile.overrides, OnWarning: logWarning, }, - OnRows: printDryRunSQL(dryRun, dbType), + OnRows: printDryRunSQL(dryRun, dbType), + OnNotice: func(msg string) { log.Warn().Msg(msg) }, OnTableStart: func(table string) error { log.Info().Str("table", table).Msg("Seeding table") return nil }, }) if err != nil { + if !dryRun { + logPartialRun(res, gapTables) + } return err } totalRows := res.Total @@ -285,7 +297,11 @@ func printGapReport(allSorted []string, counts map[string]int64, fkParents map[s totalGapRows := 0 for _, tableName := range allSorted { - count := counts[tableName] + count, known := counts[tableName] + if !known { + fmt.Printf(" %-*s %6s count failed → skipped\n", nameWidth, tableName, "?") + continue + } if gapSet[tableName] { deps := formatFKDeps(fkParents[tableName], gapSet, counts) fmt.Printf(" %-*s %6d EMPTY → would seed %d rows%s\n", @@ -317,7 +333,11 @@ func formatFKDeps(parents []string, gapSet map[string]bool, counts map[string]in if gapSet[p] { parts = append(parts, p+" (filling)") } else { - parts = append(parts, fmt.Sprintf("%s (%d rows)", p, counts[p])) + if n, ok := counts[p]; ok { + parts = append(parts, fmt.Sprintf("%s (%d rows)", p, n)) + } else { + parts = append(parts, p+" (count unknown)") + } } } return " [FK → " + strings.Join(parts, ", ") + "]" diff --git a/internal/cli/generate.go b/internal/cli/generate.go index 220b0dd..9702137 100644 --- a/internal/cli/generate.go +++ b/internal/cli/generate.go @@ -136,8 +136,11 @@ func generateCmd() *cli.Command { Overrides: profile.overrides, OnWarning: logWarning, }, - OnTableStart: w.Table, - OnRows: func(_ string, rows []map[string]interface{}) error { return w.Rows(rows) }, + OnTableStart: func(table string) error { + log.Info().Str("table", table).Msg("Generating table") + return w.Table(table) + }, + OnRows: func(_ string, rows []map[string]interface{}) error { return w.Rows(rows) }, }); err != nil { return fmt.Errorf("generation failed: %w", err) } diff --git a/internal/cli/helpers.go b/internal/cli/helpers.go index 1682efb..21a1649 100644 --- a/internal/cli/helpers.go +++ b/internal/cli/helpers.go @@ -6,11 +6,13 @@ import ( "fmt" "io" "os" + "strings" "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/faker" "github.com/AxeForging/seedstorm/internal/fsutil" "github.com/AxeForging/seedstorm/internal/logging" + "github.com/AxeForging/seedstorm/internal/seeder" ) // syncSequences moves Postgres sequences past the ids a run inserted so the @@ -86,3 +88,24 @@ func buildInsert(tableName string, row map[string]interface{}, dbType string) (s func buildBatchInsert(tableName string, rows []map[string]interface{}, dbType string) (string, []interface{}) { return db.BuildBatchInsert(tableName, rows, dbType) } + +// logPartialRun says what a failed run wrote before it stopped, so the user +// knows which tables hold new rows and which were never reached. +func logPartialRun(res seeder.SeedResult, order []string) { + var written, notWritten []string + for _, t := range order { + if n := res.Counts[t]; n > 0 { + written = append(written, fmt.Sprintf("%s (%d)", t, n)) + } else { + notWritten = append(notWritten, t) + } + } + ev := logging.Log.Warn().Int("rows_written", res.Total) + if len(written) > 0 { + ev = ev.Str("written", strings.Join(written, ", ")) + } + if len(notWritten) > 0 { + ev = ev.Str("not_written", strings.Join(notWritten, ", ")) + } + ev.Msg("Run stopped before finishing") +} diff --git a/internal/cli/introspect.go b/internal/cli/introspect.go index 528166d..2345204 100644 --- a/internal/cli/introspect.go +++ b/internal/cli/introspect.go @@ -2,11 +2,14 @@ package cli import ( "context" + "database/sql" "fmt" + "time" "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/faker" "github.com/AxeForging/seedstorm/internal/logging" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/schema" "github.com/urfave/cli/v3" ) @@ -47,11 +50,19 @@ Outputs a schema.yaml that can be used for seeding or AI enrichment.`, log.Info(). Str("db", cmd.String("db")). Msg("Connecting to database") - - tables, err := db.Introspect(dbType, dsn) + conn, err := sql.Open(dbType, dsn) if err != nil { return fmt.Errorf("introspection failed: %w", err) } + defer conn.Close() + if err := pingWithin(ctx, conn); err != nil { + return runerr.At(runerr.PhaseConnect, "", fmt.Errorf("%s did not answer: %w", dsnLabel(dbType, dsn), err)) + } + log.Info().Msg("Reading the catalog") + tables, err := db.IntrospectConn(ctx, conn, dbType, stepLogger("Introspecting", time.Now)) + if err != nil { + return runerr.At(runerr.PhaseIntrospect, "", fmt.Errorf("introspection failed: %w", err)) + } log.Info(). Int("tables", len(tables)). diff --git a/internal/cli/mirror.go b/internal/cli/mirror.go index 8e9beff..a27aac2 100644 --- a/internal/cli/mirror.go +++ b/internal/cli/mirror.go @@ -40,6 +40,7 @@ func mirrorCmd() *cli.Command { &cli.IntFlag{Name: "seed", Usage: "Random seed for reproducible data generation (0 = random)"}, &cli.BoolFlag{Name: "interactive", Aliases: []string{"i"}, Usage: "Review the plan, preview samples and confirm in the terminal UI"}, ) + flags = append(flags, productionFlags()...) return &cli.Command{ Name: "mirror", Usage: "Seed a target database so its table volumes follow a source database", @@ -57,6 +58,11 @@ same-database safety check cannot run then, so double-check --target-dsn.`, if err != nil { return err } + if !cmd.Bool("dry-run") { + if err := refuseProductionWrite(cmd, "mirror into the target"); err != nil { + return err + } + } counts, err := countMode(cmd) if err != nil { return err @@ -87,6 +93,7 @@ same-database safety check cannot run then, so double-check --target-dsn.`, log.Info().Str("source", source.Label).Str("target", target.Label).Msg("Comparing databases") job, err := seeder.PrepareMirror(ctx, source, target, seeder.MirrorConfig{ + OnCount: sideStepLogger("Counting", time.Now), Options: compare.MirrorOptions{ Mode: mode, Scale: cmd.Float("scale"), diff --git a/internal/cli/production.go b/internal/cli/production.go new file mode 100644 index 0000000..a7b728c --- /dev/null +++ b/internal/cli/production.go @@ -0,0 +1,33 @@ +package cli + +import ( + "fmt" + + "github.com/urfave/cli/v3" +) + +// productionFlags mark the database a command writes to as production. Writes +// are then refused unless --allow-production is also given; reads and dry runs +// are unaffected. +func productionFlags() []cli.Flag { + return []cli.Flag{ + &cli.BoolFlag{ + Name: "production", + Usage: "The database written to is production: refuse to write unless --allow-production is also set", + Sources: cli.EnvVars("SEEDSTORM_PRODUCTION"), + }, + &cli.BoolFlag{ + Name: "allow-production", + Usage: "Confirm writing to a database marked --production", + }, + } +} + +// refuseProductionWrite stops a write to a production database that was not +// explicitly allowed. action says what would have been written. +func refuseProductionWrite(cmd *cli.Command, action string) error { + if !cmd.Bool("production") || cmd.Bool("allow-production") { + return nil + } + return fmt.Errorf("the database is marked production (--production or SEEDSTORM_PRODUCTION): refusing to %s; pass --allow-production to confirm, or --dry-run to preview", action) +} diff --git a/internal/cli/progress.go b/internal/cli/progress.go index 59f3297..e265903 100644 --- a/internal/cli/progress.go +++ b/internal/cli/progress.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "strings" + "sync" "time" "github.com/urfave/cli/v3" @@ -12,6 +13,7 @@ import ( "github.com/AxeForging/seedstorm/internal/graph" "github.com/AxeForging/seedstorm/internal/logging" "github.com/AxeForging/seedstorm/internal/seeder" + "github.com/AxeForging/seedstorm/internal/tuning" ) // progressInterval is the most often a run logs a progress line. @@ -33,6 +35,18 @@ func genWorkersFlag() cli.Flag { } } +// genWorkers reads --gen-workers, lowered to the cores this process may use +// (a container CPU quota counts, not only the host's cores) with a log line +// saying so. +func genWorkers(cmd *cli.Command) int { + requested := cmd.Int("gen-workers") + n := tuning.ClampGenerators(requested) + if requested > n { + logging.Log.Info().Int("requested", requested).Int("using", n).Msg("Generators limited to the CPUs available") + } + return n +} + // progressLogger logs rows written, rate and ETA at most every progressInterval, // plus a line per finished table, so a long table is never silent. func progressLogger(now func() time.Time) (onProgress, onTable func(seeder.Progress)) { @@ -93,3 +107,33 @@ func populatedCheck(ctx context.Context, conn *sql.DB, dbType string) func(strin return counts[table] > 0, nil } } + +// stepLogger logs a table-by-table step (counting, introspecting) at most every +// progressInterval, plus its last table, so a long step is never silent. +func stepLogger(what string, now func() time.Time) func(done, total int, table string) { + var last time.Time + return func(done, total int, table string) { + t := now() + if done < total && t.Sub(last) < progressInterval { + return + } + last = t + logging.Log.Info().Int("done", done).Int("total", total).Str("table", table).Msg(what) + } +} + +// sideStepLogger is stepLogger for a two-sided step, one throttle per side. +func sideStepLogger(what string, now func() time.Time) func(side string, done, total int, table string) { + loggers := map[string]func(int, int, string){} + var mu sync.Mutex + return func(side string, done, total int, table string) { + mu.Lock() + l, ok := loggers[side] + if !ok { + l = stepLogger(what+" ("+side+")", now) + loggers[side] = l + } + mu.Unlock() + l(done, total, table) + } +} diff --git a/internal/cli/seed.go b/internal/cli/seed.go index 9174334..c7da893 100644 --- a/internal/cli/seed.go +++ b/internal/cli/seed.go @@ -13,6 +13,7 @@ import ( "github.com/AxeForging/seedstorm/internal/faker" "github.com/AxeForging/seedstorm/internal/graph" "github.com/AxeForging/seedstorm/internal/logging" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/schema" "github.com/AxeForging/seedstorm/internal/seeder" "github.com/AxeForging/seedstorm/internal/tui" @@ -101,6 +102,8 @@ Use --dry-run to print SQL statements without executing them.`, workersFlag(), genWorkersFlag(), profileFlag(), + productionFlags()[0], + productionFlags()[1], }, Action: func(ctx context.Context, cmd *cli.Command) error { log := logging.Log @@ -120,6 +123,11 @@ Use --dry-run to print SQL statements without executing them.`, yes := cmd.Bool("yes") batchSize := cmd.Int("batch-size") seed := cmd.Int("seed") + if !dryRun { + if err := refuseProductionWrite(cmd, "seed it"); err != nil { + return err + } + } if seed != 0 { faker.SeedRandom(int64(seed)) @@ -169,8 +177,8 @@ Use --dry-run to print SQL statements without executing them.`, } defer dbConn.Close() - if err := dbConn.PingContext(ctx); err != nil { - return fmt.Errorf("failed to ping database: %w", err) + if err := pingWithin(ctx, dbConn); err != nil { + return runerr.At(runerr.PhaseConnect, "", fmt.Errorf("%s did not answer: %w", dsnLabel(dbType, dsn), err)) } // Every table stays in the preload so FKs can reference rows of ignored @@ -186,6 +194,11 @@ Use --dry-run to print SQL statements without executing them.`, fmt.Println("--- SQL ---") } + // Refuse tables that cannot be generated before anything is truncated. + if err := faker.CheckSeedable(s, sortedTables, profile.overrides); err != nil { + return err + } + // Truncate tables before seeding if truncate && !dryRun { if !yes { @@ -198,7 +211,7 @@ Use --dry-run to print SQL statements without executing them.`, } log.Info().Int("tables", len(sortedTables)).Msg("Truncating tables") if err := db.TruncateConcurrently(ctx, dbConn, dbType, sortedTables, cmd.Int("workers"), nil); err != nil { - return fmt.Errorf("truncate failed: %w", err) + return runerr.At(runerr.PhaseTruncate, "", fmt.Errorf("truncate failed: %w", err)) } log.Info().Msg("Truncate complete") } @@ -214,19 +227,23 @@ Use --dry-run to print SQL statements without executing them.`, res, err := seeder.Seed(ctx, dbConn, dbType, s, allTables, sortedTables, seeder.SeedOptions{ Rows: rows, EnumRows: enumRows, TableRows: tableRows, BatchSize: batchSize, DryRun: dryRun, Workers: cmd.Int("workers"), OnProgress: onProgress, OnTable: onTable, - GenWorkers: cmd.Int("gen-workers"), Reproducible: cmd.Int("seed") != 0, + GenWorkers: genWorkers(cmd), Reproducible: cmd.Int("seed") != 0, Generate: faker.GenerateOptions{ SelfRefDepth: selfRefDepth, Overrides: profile.overrides, OnWarning: logWarning, }, - OnRows: printDryRunSQL(dryRun, dbType), + OnRows: printDryRunSQL(dryRun, dbType), + OnNotice: func(msg string) { log.Warn().Msg(msg) }, OnTableStart: func(table string) error { log.Info().Str("table", table).Msg("Seeding table") return nil }, }) if err != nil { + if !dryRun { + logPartialRun(res, sortedTables) + } return err } diff --git a/internal/cli/snapshot.go b/internal/cli/snapshot.go index ed1b1d6..c7bd512 100644 --- a/internal/cli/snapshot.go +++ b/internal/cli/snapshot.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "strings" + "time" "github.com/urfave/cli/v3" @@ -51,7 +52,7 @@ A hand-written file with only row counts also works: defer closeEndpoints(ep) log.Info().Str("database", ep.Label).Str("counts", string(mode)).Msg("Reading table counts") - snap, err := compare.Take(ctx, ep.Conn, ep.DBType, ep.Label, mode, nil) + snap, err := compare.Take(ctx, ep.Conn, ep.DBType, ep.Label, mode, stepLogger("Counting", time.Now)) if err != nil { return err } diff --git a/internal/compare/compare.go b/internal/compare/compare.go index 8b88828..ec008d7 100644 --- a/internal/compare/compare.go +++ b/internal/compare/compare.go @@ -86,14 +86,26 @@ func Take(ctx context.Context, conn *sql.DB, dbType, label string, mode CountMod // A missing estimate is unknown, and a zero one may just be stale (a // table filled since statistics were last gathered). Counting those // exactly is cheap when the table really is empty and correct when not. - if n, ok := estimates[name]; ok && n > 0 { + n, hasEstimate := estimates[name] + size, hasSize := sizes[name] + if !hasSize { + size = db.UnknownCount + } + switch { + case hasEstimate && n > 0: counts[name], estimated[name] = n, true - } else { - one, err := db.GetTableRowCounts(ctx, conn, dbType, []string{name}) - if err != nil { + case mode == CountEstimate && !exactFallback(n, hasEstimate, size): + // Too large to scan in estimate mode: the count stays unknown. + default: + if err := ctx.Err(); err != nil { return snap, err } - counts[name] = one[name] + // A table that cannot be counted stays unknown instead of failing + // the whole snapshot. + one, _ := db.CountTables(ctx, conn, dbType, []string{name}, nil) + if n, ok := one[name]; ok { + counts[name] = n + } } if progress != nil { progress(i+1, len(names), name) @@ -121,6 +133,9 @@ const ( StatusDiffers Status = "differs" StatusSourceOnly Status = "source_only" StatusTargetOnly Status = "target_only" + // StatusUnknown is a matched table whose row count one side could not + // report: it is neither the same nor different. + StatusUnknown Status = "unknown" ) // Row compares one table across both databases. @@ -148,6 +163,8 @@ type Totals struct { Differs int `json:"differs"` SourceOnly int `json:"sourceOnly"` TargetOnly int `json:"targetOnly"` + // Unknown counts matched tables whose row count failed on either side. + Unknown int `json:"unknown"` // ColumnDrift counts matched tables whose column names differ. ColumnDrift int `json:"columnDrift"` } @@ -201,9 +218,12 @@ func Diff(source, target Snapshot) Report { if tgtName != name { row.TargetTable = tgtName } - row.Delta = knownRows(tgt.Rows) - knownRows(src.Rows) - if src.Rows != tgt.Rows { + switch { + case src.Rows < 0 || tgt.Rows < 0: + row.Status = StatusUnknown + case src.Rows != tgt.Rows: row.Status = StatusDiffers + row.Delta = tgt.Rows - src.Rows } row.MissingColumns, row.ExtraColumns = columnDiff(src.Columns, tgt.Columns) r.Rows = append(r.Rows, row) @@ -229,6 +249,8 @@ func Diff(source, target Snapshot) Report { r.Totals.SourceOnly++ case StatusTargetOnly: r.Totals.TargetOnly++ + case StatusUnknown: + r.Totals.Unknown++ } if len(row.MissingColumns) > 0 || len(row.ExtraColumns) > 0 { r.Totals.ColumnDrift++ @@ -281,3 +303,17 @@ func sortedNames(m map[string]TableStat) []string { sort.Strings(out) return out } + +// exactFallbackMaxBytes is the largest table estimate mode still counts +// exactly when its statistics are missing or zero. +const exactFallbackMaxBytes = 256 << 20 + +// exactFallback reports whether estimate mode counts a table exactly: its +// estimate is missing or zero, and the table is small enough (or of unknown +// size) that COUNT(*) is cheap. Large tables stay unknown instead of scanned. +func exactFallback(estimate int64, hasEstimate bool, bytes int64) bool { + if hasEstimate && estimate > 0 { + return false + } + return bytes <= exactFallbackMaxBytes +} diff --git a/internal/compare/compare_test.go b/internal/compare/compare_test.go index 23a719e..1eb467c 100644 --- a/internal/compare/compare_test.go +++ b/internal/compare/compare_test.go @@ -365,3 +365,87 @@ func TestRenderReport_MarksEstimatedCounts(t *testing.T) { t.Fatalf("legend missing:\n%s", out) } } + +// A count the database could not report is unknown, not zero: a table whose +// count failed on either side is neither "same" nor "differs", and its delta +// is not invented from a zero. +func TestDiff_UnknownCountIsItsOwnStatus(t *testing.T) { + src := snap("a", map[string]TableStat{"users": stat(100), "orders": {Rows: db.UnknownCount, Bytes: db.UnknownCount}, "tags": stat(3)}) + tgt := snap("b", map[string]TableStat{"users": {Rows: db.UnknownCount, Bytes: db.UnknownCount}, "orders": stat(7), "tags": stat(3)}) + r := Diff(src, tgt) + for _, table := range []string{"users", "orders"} { + row := rowFor(t, r, table) + if row.Status != StatusUnknown || row.Delta != 0 { + t.Errorf("%s: status=%s delta=%d, want unknown with no delta", table, row.Status, row.Delta) + } + } + if row := rowFor(t, r, "tags"); row.Status != StatusSame { + t.Errorf("tags status = %s", row.Status) + } + if r.Totals.Unknown != 2 || r.Totals.Differs != 0 || r.Totals.Same != 1 { + t.Fatalf("totals = %+v", r.Totals) + } +} + +// Mirror planned an unknown target count as an empty table and inserted the +// full source volume into it. A target it cannot count is skipped, and it is +// never picked as an empty parent to fill either. +func TestPlanMirror_UnknownTargetCountIsNeverFilled(t *testing.T) { + src := snap("src", map[string]TableStat{"users": stat(50), "orders": stat(80)}) + tgt := snap("tgt", map[string]TableStat{ + "users": {Rows: db.UnknownCount, Bytes: db.UnknownCount}, + "orders": stat(0), + "order_items": stat(0), + "audit_logs": stat(0), + }) + plan, err := PlanMirror(Diff(src, tgt), shopTarget(), MirrorOptions{Tables: []string{"orders"}, ParentRows: 10}) + if err != nil { + t.Fatal(err) + } + for _, e := range plan.Entries { + if e.Table == "users" { + t.Fatalf("users has an unknown target count but is planned: %+v", e) + } + } + if e := entry(t, plan, "orders"); e.Insert != 80 { + t.Fatalf("orders = %+v", e) + } + + plan, err = PlanMirror(Diff(src, tgt), shopTarget(), MirrorOptions{}) + if err != nil { + t.Fatal(err) + } + reasons := map[string]string{} + for _, s := range plan.Skipped { + reasons[s.Table] = s.Reason + } + if reasons["users"] != ReasonTargetUnknown { + t.Fatalf("users skip = %q, all skips %+v", reasons["users"], plan.Skipped) + } +} + +// Estimate mode exists to avoid COUNT(*) on huge tables. A zero or missing +// estimate is still counted exactly when the table is small (stale statistics +// on a fresh table), but a large table stays unknown instead of being scanned. +func TestExactFallback_OnlyForTablesCheapToCount(t *testing.T) { + cases := []struct { + name string + estimate int64 + hasEstimate bool + bytes int64 + want bool + }{ + {"positive estimate is used", 5000, true, 1 << 30, false}, + {"zero estimate on a small table", 0, true, 64 << 10, true}, + {"missing estimate on a small table", 0, false, 64 << 10, true}, + {"missing estimate, size unknown", 0, false, db.UnknownCount, true}, + {"zero estimate on a large table", 0, true, 4 << 30, false}, + {"missing estimate at the limit", 0, false, exactFallbackMaxBytes, true}, + {"missing estimate just over the limit", 0, false, exactFallbackMaxBytes + 1, false}, + } + for _, c := range cases { + if got := exactFallback(c.estimate, c.hasEstimate, c.bytes); got != c.want { + t.Errorf("%s: exactFallback = %v, want %v", c.name, got, c.want) + } + } +} diff --git a/internal/compare/mirror.go b/internal/compare/mirror.go index 0e5e29c..a8fc2c2 100644 --- a/internal/compare/mirror.go +++ b/internal/compare/mirror.go @@ -53,14 +53,17 @@ type MirrorOptions struct { // Reasons attached to plan entries and skips. const ( - ReasonMatch = "match source" - ReasonCapped = "capped by max rows" - ReasonParent = "required parent is empty" - ReasonDependent = "truncated with its parent" - ReasonNotTarget = "not in target" - ReasonUnknown = "source row count unknown" - ReasonSatisfied = "target already has enough rows" - ReasonNotInGraph = "not introspected on target" + ReasonMatch = "match source" + ReasonCapped = "capped by max rows" + ReasonParent = "required parent is empty" + ReasonDependent = "truncated with its parent" + ReasonNotTarget = "not in target" + ReasonUnknown = "source row count unknown" + // ReasonTargetUnknown: the target's count failed, so how many rows it + // already has is unknown; filling it could double its volume. + ReasonTargetUnknown = "target row count unknown" + ReasonSatisfied = "target already has enough rows" + ReasonNotInGraph = "not introspected on target" // ReasonIgnored marks a table the profile ignores. ReasonIgnored = "ignored by profile" // ReasonIgnoredParent marks a table whose required FK parent is ignored and empty. @@ -160,10 +163,14 @@ func PlanMirror(report Report, target *schema.Schema, opts MirrorOptions) (Mirro targetRows := make(map[string]int64) candidates := make(map[string]Row) byTarget := make(map[string]Row) + // unknownTarget holds tables whose target count failed: never filled, and + // never assumed empty when a child needs a parent with rows. + unknownTarget := make(map[string]bool) for _, row := range report.Rows { name := targetName(row) if row.Target != nil { targetRows[name] = knownRows(row.Target.Rows) + unknownTarget[name] = row.Target.Rows < 0 byTarget[name] = row } if !selected(row) { @@ -181,6 +188,9 @@ func PlanMirror(report Report, target *schema.Schema, opts MirrorOptions) (Mirro case row.Source.Rows < 0: staticSkips = append(staticSkips, PlanSkip{Table: name, Reason: ReasonUnknown}) continue + case unknownTarget[name]: + staticSkips = append(staticSkips, PlanSkip{Table: name, Reason: ReasonTargetUnknown}) + continue } if _, ok := target.Tables[name]; !ok { staticSkips = append(staticSkips, PlanSkip{Table: name, Reason: ReasonNotInGraph}) @@ -223,6 +233,10 @@ func PlanMirror(report Report, target *schema.Schema, opts MirrorOptions) (Mirro entries, truncated, skips = planPass(candidates, excluded, byTarget, target, opts, isIgnored) available = func(name string) int64 { have := targetRows[name] + if unknownTarget[name] && !truncated[name] { + // Unknown is not empty: do not refill it as a missing parent. + have = max(have, 1) + } if truncated[name] { have = 0 } diff --git a/internal/db/access.go b/internal/db/access.go index 5d2f296..b482bb4 100644 --- a/internal/db/access.go +++ b/internal/db/access.go @@ -41,18 +41,24 @@ type Access struct { // InspectAccess reports what the user behind conn is allowed to do. dbType is // "pgx" (PostgreSQL) or "mysql". -func InspectAccess(ctx context.Context, conn *sql.DB, dbType string) (Access, error) { +func InspectAccess(ctx context.Context, conn *sql.DB, dbType string) (acc Access, err error) { + var inspect func(context.Context, Querier) (Access, error) switch dbType { case "pgx", "postgres", "postgresql": - return inspectPostgresAccess(ctx, conn) + inspect = inspectPostgresAccess case "mysql": - return inspectMySQLAccess(ctx, conn) + inspect = inspectMySQLAccess default: return Access{}, fmt.Errorf("inspect access: unsupported database type %q", dbType) } + err = ReadOnce(ctx, conn, dbType, ReadLimits{}, func(ctx context.Context, q Querier) error { + acc, err = inspect(ctx, q) + return err + }) + return acc, err } -func inspectPostgresAccess(ctx context.Context, conn *sql.DB) (Access, error) { +func inspectPostgresAccess(ctx context.Context, conn Querier) (Access, error) { acc := Access{Tables: map[string]TableAccess{}, Notes: []string{}} var hasSchema, usage bool err := conn.QueryRowContext(ctx, ` @@ -121,15 +127,9 @@ func inspectPostgresAccess(ctx context.Context, conn *sql.DB) (Access, error) { return acc, nil } -func inspectMySQLAccess(ctx context.Context, conn *sql.DB) (Access, error) { - // Pin one session so CURRENT_ROLE() and SHOW GRANTS describe the same - // connection. - c, err := conn.Conn(ctx) - if err != nil { - return Access{}, fmt.Errorf("inspect access: %w", err) - } - defer func() { _ = c.Close() }() - +// inspectMySQLAccess runs on one session (ReadOnce pins it), so CURRENT_ROLE() +// and SHOW GRANTS describe the same connection. +func inspectMySQLAccess(ctx context.Context, c Querier) (Access, error) { var user string var database sql.NullString if err := c.QueryRowContext(ctx, `SELECT CURRENT_USER(), DATABASE()`).Scan(&user, &database); err != nil { @@ -199,7 +199,7 @@ func inspectMySQLAccess(ctx context.Context, conn *sql.DB) (Access, error) { return acc, nil } -func mysqlShowGrants(ctx context.Context, c *sql.Conn, query string) ([]string, error) { +func mysqlShowGrants(ctx context.Context, c Querier, query string) ([]string, error) { rows, err := c.QueryContext(ctx, query) if err != nil { return nil, fmt.Errorf("%s: %w", query, err) diff --git a/internal/db/clone.go b/internal/db/clone.go index b0085b1..d32baed 100644 --- a/internal/db/clone.go +++ b/internal/db/clone.go @@ -256,14 +256,7 @@ func ddlProgressLabel(stmt string) string { } func introspectWithConn(conn *sql.DB, dbType string) ([]Table, error) { - switch dbType { - case "pgx": - return introspectPostgres(conn) - case "mysql": - return introspectMySQL(conn) - default: - return nil, fmt.Errorf("unsupported database type %q", dbType) - } + return IntrospectConn(context.Background(), conn, dbType, nil) } func buildCreateTable(table Table, dbType string) (string, error) { diff --git a/internal/db/copy.go b/internal/db/copy.go index 2208544..fadf361 100644 --- a/internal/db/copy.go +++ b/internal/db/copy.go @@ -13,6 +13,9 @@ import ( "time" "github.com/jackc/pgx/v5/stdlib" + + "github.com/AxeForging/seedstorm/internal/faultinject" + "github.com/AxeForging/seedstorm/internal/safego" ) // ErrCopyUnsupported reports an engine or connection without COPY support. @@ -39,7 +42,13 @@ func CopyRows(ctx context.Context, conn *sql.DB, dbType, table string, rows []ma } reader, writer := io.Pipe() go func() { - writer.CloseWithError(writeCSVRows(writer, cols, rows)) + // A panic while encoding rows fails this COPY instead of the process. + writer.CloseWithError(safego.Run("copy "+table, func() error { + if err := faultinject.Hit(ctx, "copy", table); err != nil { + return err + } + return writeCSVRows(writer, cols, rows) + })) }() _, err := pc.Conn().PgConn().CopyFrom(ctx, reader, stmt) _ = reader.Close() diff --git a/internal/db/counts.go b/internal/db/counts.go index 32db7d0..ab31cbc 100644 --- a/internal/db/counts.go +++ b/internal/db/counts.go @@ -4,21 +4,88 @@ import ( "context" "database/sql" "fmt" + "sync" ) // GetTableRowCounts queries SELECT COUNT(*) for each table and returns a -// map of table name → row count. Tables that cannot be queried return 0 and -// the error is surfaced immediately. +// map of table name → row count. The first table that cannot be counted ends +// the scan with its error; use CountTables to keep counting past failures. func GetTableRowCounts(ctx context.Context, conn *sql.DB, dbType string, tableNames []string) (map[string]int64, error) { counts := make(map[string]int64, len(tableNames)) for _, tableName := range tableNames { - var n int64 - //nolint:gosec - row := conn.QueryRowContext(ctx, fmt.Sprintf("SELECT COUNT(*) FROM %s", QuoteIdent(tableName, dbType))) - if err := row.Scan(&n); err != nil { - return nil, fmt.Errorf("count rows in %s: %w", tableName, err) + n, err := countRows(ctx, conn, dbType, tableName) + if err != nil { + return nil, err } counts[tableName] = n } return counts, nil } + +// CountTables counts every table it can within DefaultCountLimits. Counted +// tables are in counts; a table whose count failed is left out of counts (its +// row count is unknown, never 0) and its error is in failed. onTable, if set, +// is called after each table. +func CountTables(ctx context.Context, conn *sql.DB, dbType string, tableNames []string, onTable func(done, total int, table string)) (counts map[string]int64, failed map[string]error) { + return CountTablesWithin(ctx, conn, dbType, tableNames, DefaultCountLimits, onTable) +} + +// CountTablesWithin is CountTables with explicit limits: each count runs in its +// own read-only transaction (ReadOnce), lim.Concurrency at a time. After ctx +// ends the remaining tables fail with ctx's error. +func CountTablesWithin(ctx context.Context, conn *sql.DB, dbType string, tableNames []string, lim ReadLimits, onTable func(done, total int, table string)) (counts map[string]int64, failed map[string]error) { + counts = make(map[string]int64, len(tableNames)) + failed = map[string]error{} + var mu sync.Mutex + done := 0 + record := func(table string, n int64, err error) { + mu.Lock() + defer mu.Unlock() + if err != nil { + failed[table] = err + } else { + counts[table] = n + } + done++ + if onTable != nil { + onTable(done, len(tableNames), table) + } + } + work := make(chan string) + var wg sync.WaitGroup + for i := 0; i < max(1, lim.Concurrency); i++ { + wg.Add(1) + go func() { + defer wg.Done() + for table := range work { + if err := ctx.Err(); err != nil { + record(table, 0, err) + continue + } + var n int64 + err := ReadOnce(ctx, conn, dbType, lim, func(ctx context.Context, q Querier) error { + var err error + n, err = countRows(ctx, q, dbType, table) + return err + }) + record(table, n, err) + } + }() + } + for _, table := range tableNames { + work <- table + } + close(work) + wg.Wait() + return counts, failed +} + +func countRows(ctx context.Context, conn Querier, dbType, tableName string) (int64, error) { + var n int64 + //nolint:gosec + row := conn.QueryRowContext(ctx, fmt.Sprintf("SELECT COUNT(*) FROM %s", QuoteIdent(tableName, dbType))) + if err := row.Scan(&n); err != nil { + return 0, fmt.Errorf("count rows in %s: %w", tableName, err) + } + return n, nil +} diff --git a/internal/db/db.go b/internal/db/db.go index 4378f04..4437713 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -1,10 +1,18 @@ package db import ( + "context" "database/sql" "fmt" + "time" + + "github.com/AxeForging/seedstorm/internal/faultinject" ) +// introspectConnectTimeout bounds how long Introspect waits for the database +// to answer before reporting it unreachable. +var introspectConnectTimeout = 10 * time.Second + // Introspect connects to the database and returns all discovered tables. func Introspect(dbType, dsn string) ([]Table, error) { db, err := sql.Open(dbType, dsn) @@ -13,15 +21,26 @@ func Introspect(dbType, dsn string) ([]Table, error) { } defer db.Close() - if err := db.Ping(); err != nil { + ctx, cancel := context.WithTimeout(context.Background(), introspectConnectTimeout) + defer cancel() + if err := db.PingContext(ctx); err != nil { return nil, fmt.Errorf("failed to ping database: %w", err) } + return IntrospectConn(context.Background(), db, dbType, nil) +} +// IntrospectConn reads every table of the connection's database. Its reads +// stop when ctx ends; onTable, if set, is called after each table (catalog +// queries for the whole database run before the first one). +func IntrospectConn(ctx context.Context, conn *sql.DB, dbType string, onTable func(done, total int, table string)) ([]Table, error) { + if err := faultinject.Hit(ctx, "introspect", ""); err != nil { + return nil, err + } switch dbType { case "pgx": - return introspectPostgres(db) + return introspectPostgres(ctx, conn, onTable) case "mysql": - return introspectMySQL(db) + return introspectMySQL(ctx, conn, onTable) default: return nil, fmt.Errorf("unsupported database type %q (use mysql or postgres)", dbType) } diff --git a/internal/db/mysql.go b/internal/db/mysql.go index 0352ffb..3f8eebe 100644 --- a/internal/db/mysql.go +++ b/internal/db/mysql.go @@ -1,6 +1,7 @@ package db import ( + "context" "database/sql" "fmt" "regexp" @@ -8,15 +9,15 @@ import ( "strings" ) -func introspectMySQL(db *sql.DB) ([]Table, error) { +func introspectMySQL(ctx context.Context, db *sql.DB, onTable func(done, total int, table string)) ([]Table, error) { // Get current database name var dbName string - if err := db.QueryRow("SELECT DATABASE()").Scan(&dbName); err != nil { + if err := db.QueryRowContext(ctx, "SELECT DATABASE()").Scan(&dbName); err != nil { return nil, fmt.Errorf("failed to get current database: %w", err) } // List all tables - tableRows, err := db.Query(` + tableRows, err := db.QueryContext(ctx, ` SELECT TABLE_NAME FROM information_schema.TABLES WHERE TABLE_SCHEMA = ? @@ -37,7 +38,7 @@ func introspectMySQL(db *sql.DB) ([]Table, error) { } // Fetch FK relationships for the whole database - fkMap, err := mysqlFKMap(db, dbName) + fkMap, err := mysqlFKMap(ctx, db, dbName) if err != nil { return nil, err } @@ -48,38 +49,41 @@ func introspectMySQL(db *sql.DB) ([]Table, error) { // or stored in metadata on those versions. var checkMap map[string]map[string][]string var rangeMap map[string]map[string]rangeConstraint - supportsCheck, err := mysqlSupportsCheckConstraints(db) + supportsCheck, err := mysqlSupportsCheckConstraints(ctx, db) if err != nil { return nil, err } if supportsCheck { - checkMap, err = mysqlCheckMap(db, dbName) + checkMap, err = mysqlCheckMap(ctx, db, dbName) if err != nil { return nil, err } - rangeMap, err = mysqlRangeMap(db, dbName) + rangeMap, err = mysqlRangeMap(ctx, db, dbName) if err != nil { return nil, err } } - indexMap, err := mysqlIndexMap(db, dbName) + indexMap, err := mysqlIndexMap(ctx, db, dbName) if err != nil { return nil, err } - tableComments, err := mysqlTableCommentMap(db, dbName) + tableComments, err := mysqlTableCommentMap(ctx, db, dbName) if err != nil { return nil, err } var tables []Table for _, tableName := range tableNames { - cols, err := mysqlColumns(db, dbName, tableName, fkMap, checkMap, rangeMap) + cols, err := mysqlColumns(ctx, db, dbName, tableName, fkMap, checkMap, rangeMap) if err != nil { return nil, fmt.Errorf("failed to introspect table %s: %w", tableName, err) } tables = append(tables, Table{Name: tableName, Columns: cols, Indexes: indexMap[tableName], Comment: tableComments[tableName]}) + if onTable != nil { + onTable(len(tables), len(tableNames), tableName) + } } return tables, nil @@ -87,9 +91,9 @@ func introspectMySQL(db *sql.DB) ([]Table, error) { // mysqlSupportsCheckConstraints reports whether the server exposes // information_schema.CHECK_CONSTRAINTS. Introduced in MySQL 8.0.16. -func mysqlSupportsCheckConstraints(db *sql.DB) (bool, error) { +func mysqlSupportsCheckConstraints(ctx context.Context, db *sql.DB) (bool, error) { var one int - err := db.QueryRow(` + err := db.QueryRowContext(ctx, ` SELECT 1 FROM information_schema.TABLES WHERE TABLE_SCHEMA = 'information_schema' @@ -104,8 +108,8 @@ func mysqlSupportsCheckConstraints(db *sql.DB) (bool, error) { return true, nil } -func mysqlFKMap(db *sql.DB, dbName string) (map[string]map[string]*ForeignKey, error) { - rows, err := db.Query(` +func mysqlFKMap(ctx context.Context, db *sql.DB, dbName string) (map[string]map[string]*ForeignKey, error) { + rows, err := db.QueryContext(ctx, ` SELECT kcu.TABLE_NAME, kcu.COLUMN_NAME, @@ -136,8 +140,8 @@ func mysqlFKMap(db *sql.DB, dbName string) (map[string]map[string]*ForeignKey, e return fkMap, nil } -func mysqlColumns(db *sql.DB, dbName, tableName string, fkMap map[string]map[string]*ForeignKey, checkMap map[string]map[string][]string, rangeMap map[string]map[string]rangeConstraint) ([]Column, error) { - rows, err := db.Query(` +func mysqlColumns(ctx context.Context, db *sql.DB, dbName, tableName string, fkMap map[string]map[string]*ForeignKey, checkMap map[string]map[string][]string, rangeMap map[string]map[string]rangeConstraint) ([]Column, error) { + rows, err := db.QueryContext(ctx, ` SELECT COLUMN_NAME, DATA_TYPE, @@ -217,8 +221,8 @@ func mysqlColumns(db *sql.DB, dbName, tableName string, fkMap map[string]map[str } // mysqlCheckMap returns map[table][column]=[]values for CHECK IN constraints (MySQL 8.0.16+). -func mysqlCheckMap(db *sql.DB, dbName string) (map[string]map[string][]string, error) { - rows, err := db.Query(` +func mysqlCheckMap(ctx context.Context, db *sql.DB, dbName string) (map[string]map[string][]string, error) { + rows, err := db.QueryContext(ctx, ` SELECT tc.TABLE_NAME, cc.CHECK_CLAUSE FROM information_schema.TABLE_CONSTRAINTS tc JOIN information_schema.CHECK_CONSTRAINTS cc @@ -280,8 +284,8 @@ var ( ) // mysqlRangeMap returns map[table][column]=rangeConstraint for CHECK (col >= N AND col <= M). -func mysqlRangeMap(db *sql.DB, dbName string) (map[string]map[string]rangeConstraint, error) { - rows, err := db.Query(` +func mysqlRangeMap(ctx context.Context, db *sql.DB, dbName string) (map[string]map[string]rangeConstraint, error) { + rows, err := db.QueryContext(ctx, ` SELECT tc.TABLE_NAME, cc.CHECK_CLAUSE FROM information_schema.TABLE_CONSTRAINTS tc JOIN information_schema.CHECK_CONSTRAINTS cc @@ -351,8 +355,8 @@ func parseEnumValues(columnType string) []string { return values } -func mysqlIndexMap(db *sql.DB, dbName string) (map[string][]Index, error) { - rows, err := db.Query(` +func mysqlIndexMap(ctx context.Context, db *sql.DB, dbName string) (map[string][]Index, error) { + rows, err := db.QueryContext(ctx, ` SELECT TABLE_NAME, INDEX_NAME, @@ -403,8 +407,8 @@ func parseIndexPrefixes(raw string) []int { return out } -func mysqlTableCommentMap(db *sql.DB, dbName string) (map[string]string, error) { - rows, err := db.Query(` +func mysqlTableCommentMap(ctx context.Context, db *sql.DB, dbName string) (map[string]string, error) { + rows, err := db.QueryContext(ctx, ` SELECT TABLE_NAME, TABLE_COMMENT FROM information_schema.TABLES WHERE TABLE_SCHEMA = ? diff --git a/internal/db/partitions.go b/internal/db/partitions.go new file mode 100644 index 0000000..88f96c1 --- /dev/null +++ b/internal/db/partitions.go @@ -0,0 +1,147 @@ +package db + +import ( + "context" + "database/sql" + "fmt" + "regexp" + "strings" +) + +// postgresPartitionNames selects the names of partitions in schema public: they +// hold a partitioned table's rows and are never listed as tables themselves. +const postgresPartitionNames = ` + SELECT pc.relname FROM pg_class pc + JOIN pg_namespace pn ON pn.oid = pc.relnamespace + WHERE pn.nspname = 'public' AND pc.relispartition AND pc.relkind IN ('r', 'p')` + +// postgresPartitioning reads every partitioned table's key and the bounds of +// its direct partitions. +func postgresPartitioning(ctx context.Context, db *sql.DB) (map[string]*Partitioning, error) { + rows, err := db.QueryContext(ctx, ` + SELECT c.relname, pt.partstrat, + ARRAY(SELECT COALESCE(a.attname, '') + FROM unnest(pt.partattrs::int2[]) WITH ORDINALITY AS k(attnum, ord) + LEFT JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum = k.attnum + ORDER BY k.ord)::text[] + FROM pg_partitioned_table pt + JOIN pg_class c ON c.oid = pt.partrelid + JOIN pg_namespace n ON n.oid = c.relnamespace + WHERE n.nspname = 'public'`) + if err != nil { + return nil, fmt.Errorf("failed to query partitioned tables: %w", err) + } + out := map[string]*Partitioning{} + for rows.Next() { + var name, strategy, cols string + if err := rows.Scan(&name, &strategy, &cols); err != nil { + rows.Close() + return nil, err + } + p := &Partitioning{Strategy: map[string]string{"r": "range", "l": "list", "h": "hash"}[strategy]} + p.Columns = parseTextArray(cols) + out[name] = p + } + rows.Close() + if err := rows.Err(); err != nil { + return nil, err + } + + bounds, err := db.QueryContext(ctx, ` + SELECT parent.relname, pg_get_expr(child.relpartbound, child.oid) + FROM pg_inherits i + JOIN pg_class child ON child.oid = i.inhrelid + JOIN pg_class parent ON parent.oid = i.inhparent + JOIN pg_namespace n ON n.oid = parent.relnamespace + WHERE n.nspname = 'public' AND child.relispartition AND child.relkind IN ('r', 'p') + ORDER BY parent.relname, child.relname`) + if err != nil { + return nil, fmt.Errorf("failed to query partition bounds: %w", err) + } + defer bounds.Close() + for bounds.Next() { + var parent, bound string + if err := bounds.Scan(&parent, &bound); err != nil { + return nil, err + } + if p := out[parent]; p != nil { + applyPartitionBound(p, bound) + } + } + return out, bounds.Err() +} + +var ( + reRangeBound = regexp.MustCompile(`(?i)^FOR VALUES FROM \((.*)\) TO \((.*)\)$`) + reListBound = regexp.MustCompile(`(?i)^FOR VALUES IN \((.*)\)$`) +) + +// applyPartitionBound records one partition's bound, as pg_get_expr prints it. +func applyPartitionBound(p *Partitioning, bound string) { + bound = strings.TrimSpace(bound) + switch { + case strings.EqualFold(bound, "DEFAULT"): + p.Default = true + case reRangeBound.MatchString(bound): + m := reRangeBound.FindStringSubmatch(bound) + p.Ranges = append(p.Ranges, PartitionRange{From: strings.TrimSpace(m[1]), To: strings.TrimSpace(m[2])}) + case reListBound.MatchString(bound): + for _, v := range splitSQLList(reListBound.FindStringSubmatch(bound)[1]) { + if !strings.EqualFold(v, "NULL") { + p.Values = append(p.Values, unquoteSQL(v)) + } + } + } +} + +// splitSQLList splits a comma-separated list of SQL literals, keeping commas +// inside quotes. +func splitSQLList(s string) []string { + var out []string + var cur strings.Builder + inQuote := false + for i := 0; i < len(s); i++ { + ch := s[i] + switch { + case ch == '\'': + inQuote = !inQuote + cur.WriteByte(ch) + case ch == ',' && !inQuote: + out = append(out, strings.TrimSpace(cur.String())) + cur.Reset() + default: + cur.WriteByte(ch) + } + } + if strings.TrimSpace(cur.String()) != "" { + out = append(out, strings.TrimSpace(cur.String())) + } + return out +} + +// unquoteSQL turns 'it”s' into it's and strips a ::type cast; unquoted +// literals are returned trimmed. +func unquoteSQL(v string) string { + v = strings.TrimSpace(v) + if i := strings.LastIndex(v, "::"); i > 0 && strings.HasSuffix(v[:i], "'") { + v = v[:i] + } + if len(v) >= 2 && v[0] == '\'' && v[len(v)-1] == '\'' { + return strings.ReplaceAll(v[1:len(v)-1], "''", "'") + } + return v +} + +// parseTextArray parses a Postgres text[] literal like {a,"",b}. +func parseTextArray(s string) []string { + s = strings.TrimSpace(s) + s = strings.TrimPrefix(strings.TrimSuffix(s, "}"), "{") + if s == "" { + return nil + } + var out []string + for _, part := range strings.Split(s, ",") { + out = append(out, strings.Trim(part, `"`)) + } + return out +} diff --git a/internal/db/partitions_test.go b/internal/db/partitions_test.go new file mode 100644 index 0000000..a1f6f49 --- /dev/null +++ b/internal/db/partitions_test.go @@ -0,0 +1,46 @@ +package db + +import ( + "reflect" + "testing" +) + +func TestApplyPartitionBound_ReadsWhatPostgresPrints(t *testing.T) { + p := &Partitioning{Strategy: "range"} + for _, b := range []string{ + "FOR VALUES FROM ('2025-01-01') TO ('2026-01-01')", + "FOR VALUES FROM (MINVALUE) TO (0)", + "DEFAULT", + } { + applyPartitionBound(p, b) + } + want := []PartitionRange{{From: "'2025-01-01'", To: "'2026-01-01'"}, {From: "MINVALUE", To: "0"}} + if !reflect.DeepEqual(p.Ranges, want) || !p.Default { + t.Fatalf("range partitioning = %+v", p) + } + + l := &Partitioning{Strategy: "list"} + applyPartitionBound(l, "FOR VALUES IN ('eu-west', 'it''s, fine', NULL)") + applyPartitionBound(l, "FOR VALUES IN (7)") + if want := []string{"eu-west", "it's, fine", "7"}; !reflect.DeepEqual(l.Values, want) { + t.Fatalf("list values = %q, want %q", l.Values, want) + } + + h := &Partitioning{Strategy: "hash"} + applyPartitionBound(h, "FOR VALUES WITH (modulus 2, remainder 0)") + if h.Default || len(h.Ranges) != 0 || len(h.Values) != 0 { + t.Fatalf("hash bounds recorded as ranges or values: %+v", h) + } +} + +func TestParseTextArray_KeepsExpressionSlots(t *testing.T) { + if got := parseTextArray(`{created_at}`); !reflect.DeepEqual(got, []string{"created_at"}) { + t.Fatalf("got %q", got) + } + if got := parseTextArray(`{"",region}`); !reflect.DeepEqual(got, []string{"", "region"}) { + t.Fatalf("got %q", got) + } + if got := parseTextArray(`{}`); got != nil { + t.Fatalf("got %q", got) + } +} diff --git a/internal/db/postgres.go b/internal/db/postgres.go index f304075..df38cab 100644 --- a/internal/db/postgres.go +++ b/internal/db/postgres.go @@ -1,6 +1,7 @@ package db import ( + "context" "database/sql" "fmt" "regexp" @@ -8,12 +9,13 @@ import ( "strings" ) -func introspectPostgres(db *sql.DB) ([]Table, error) { - tableRows, err := db.Query(` +func introspectPostgres(ctx context.Context, db *sql.DB, onTable func(done, total int, table string)) ([]Table, error) { + tableRows, err := db.QueryContext(ctx, ` SELECT table_name FROM information_schema.tables WHERE table_schema = 'public' AND table_type = 'BASE TABLE' + AND table_name NOT IN (`+postgresPartitionNames+`) ORDER BY table_name`) if err != nil { return nil, fmt.Errorf("failed to list tables: %w", err) @@ -29,69 +31,79 @@ func introspectPostgres(db *sql.DB) ([]Table, error) { tableNames = append(tableNames, name) } - fkMap, err := postgresFKMap(db) + fkMap, err := postgresFKMap(ctx, db) if err != nil { return nil, err } - pkMap, err := postgresPKMap(db) + pkMap, err := postgresPKMap(ctx, db) if err != nil { return nil, err } - uniqueMap, err := postgresUniqueMap(db) + uniqueMap, err := postgresUniqueMap(ctx, db) if err != nil { return nil, err } - checkMap, err := postgresCheckMap(db) + checkMap, err := postgresCheckMap(ctx, db) if err != nil { return nil, err } - rangeMap, err := postgresRangeMap(db) + rangeMap, err := postgresRangeMap(ctx, db) if err != nil { return nil, err } - indexMap, err := postgresIndexMap(db, uniqueMap) + indexMap, err := postgresIndexMap(ctx, db, uniqueMap) if err != nil { return nil, err } - tableComments, columnComments, err := postgresCommentMaps(db) + tableComments, columnComments, err := postgresCommentMaps(ctx, db) + if err != nil { + return nil, err + } + + partitions, err := postgresPartitioning(ctx, db) if err != nil { return nil, err } var tables []Table for _, tableName := range tableNames { - cols, err := postgresColumns(db, tableName, fkMap, pkMap, uniqueMap, checkMap, rangeMap, columnComments) + cols, err := postgresColumns(ctx, db, tableName, fkMap, pkMap, uniqueMap, checkMap, rangeMap, columnComments) if err != nil { return nil, fmt.Errorf("failed to introspect table %s: %w", tableName, err) } - tables = append(tables, Table{Name: tableName, Columns: cols, Indexes: indexMap[tableName], Comment: tableComments[tableName]}) + tables = append(tables, Table{Name: tableName, Columns: cols, Indexes: indexMap[tableName], Comment: tableComments[tableName], Partition: partitions[tableName]}) + if onTable != nil { + onTable(len(tables), len(tableNames), tableName) + } } return tables, nil } -func postgresFKMap(db *sql.DB) (map[string]map[string]*ForeignKey, error) { - rows, err := db.Query(` - SELECT - kcu.table_name, - kcu.column_name, - ccu.table_name AS foreign_table, - ccu.column_name AS foreign_column - FROM information_schema.table_constraints AS tc - JOIN information_schema.key_column_usage AS kcu - ON tc.constraint_name = kcu.constraint_name - AND tc.table_schema = kcu.table_schema - JOIN information_schema.constraint_column_usage AS ccu - ON ccu.constraint_name = tc.constraint_name - AND ccu.table_schema = tc.table_schema - WHERE tc.constraint_type = 'FOREIGN KEY' - AND tc.table_schema = 'public'`) +// postgresFKMap reads foreign keys from pg_constraint. conkey and confkey are +// parallel arrays, so unnesting them together pairs each column with the one it +// references, also for multi-column keys. information_schema is not used: it +// joins constraints by name (same-named constraints on two tables mix) and +// hides constraints from roles that only have SELECT. +func postgresFKMap(ctx context.Context, db *sql.DB) (map[string]map[string]*ForeignKey, error) { + rows, err := db.QueryContext(ctx, ` + SELECT t.relname, a.attname, ft.relname, fa.attname + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + JOIN pg_namespace n ON n.oid = t.relnamespace + JOIN pg_class ft ON ft.oid = c.confrelid + CROSS JOIN LATERAL unnest(c.conkey, c.confkey) AS k(attnum, fattnum) + JOIN pg_attribute a ON a.attrelid = c.conrelid AND a.attnum = k.attnum + JOIN pg_attribute fa ON fa.attrelid = c.confrelid AND fa.attnum = k.fattnum + WHERE c.contype = 'f' + AND n.nspname = 'public' + ORDER BY t.relname, c.conname, a.attname`) if err != nil { return nil, fmt.Errorf("failed to query FK constraints: %w", err) } @@ -108,18 +120,20 @@ func postgresFKMap(db *sql.DB) (map[string]map[string]*ForeignKey, error) { } fkMap[table][column] = &ForeignKey{TableName: refTable, ColumnName: refColumn} } - return fkMap, nil + return fkMap, rows.Err() } -func postgresPKMap(db *sql.DB) (map[string]map[string]bool, error) { - rows, err := db.Query(` - SELECT kcu.table_name, kcu.column_name - FROM information_schema.table_constraints AS tc - JOIN information_schema.key_column_usage AS kcu - ON tc.constraint_name = kcu.constraint_name - AND tc.table_schema = kcu.table_schema - WHERE tc.constraint_type = 'PRIMARY KEY' - AND tc.table_schema = 'public'`) +// postgresPKMap reads primary keys from pg_constraint (visible to any role that +// can read the table, unlike information_schema.table_constraints). +func postgresPKMap(ctx context.Context, db *sql.DB) (map[string]map[string]bool, error) { + rows, err := db.QueryContext(ctx, ` + SELECT t.relname, a.attname + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + JOIN pg_namespace n ON n.oid = t.relnamespace + JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = ANY(c.conkey) + WHERE c.contype = 'p' + AND n.nspname = 'public'`) if err != nil { return nil, fmt.Errorf("failed to query PK constraints: %w", err) } @@ -136,13 +150,13 @@ func postgresPKMap(db *sql.DB) (map[string]map[string]bool, error) { } pkMap[table][column] = true } - return pkMap, nil + return pkMap, rows.Err() } type rangeConstraint struct{ Min, Max int64 } -func postgresColumns(db *sql.DB, tableName string, fkMap map[string]map[string]*ForeignKey, pkMap map[string]map[string]bool, uniqueMap map[string]map[string]bool, checkMap map[string]map[string][]string, rangeMap map[string]map[string]rangeConstraint, columnComments map[string]map[string]string) ([]Column, error) { - rows, err := db.Query(` +func postgresColumns(ctx context.Context, db *sql.DB, tableName string, fkMap map[string]map[string]*ForeignKey, pkMap map[string]map[string]bool, uniqueMap map[string]map[string]bool, checkMap map[string]map[string][]string, rangeMap map[string]map[string]rangeConstraint, columnComments map[string]map[string]string) ([]Column, error) { + rows, err := db.QueryContext(ctx, ` SELECT c.column_name, c.data_type, @@ -201,7 +215,7 @@ func postgresColumns(db *sql.DB, tableName string, fkMap map[string]map[string]* // Resolve enum values for user-defined enum types if dataType == "USER-DEFINED" { - col.EnumValues, _ = postgresEnumValues(db, udtName) + col.EnumValues, _ = postgresEnumValues(ctx, db, udtName) } if fkMap[tableName] != nil { @@ -229,8 +243,8 @@ func postgresColumns(db *sql.DB, tableName string, fkMap map[string]map[string]* } // postgresUniqueMap returns map[table][column]=true for single-column UNIQUE constraints. -func postgresUniqueMap(db *sql.DB) (map[string]map[string]bool, error) { - rows, err := db.Query(` +func postgresUniqueMap(ctx context.Context, db *sql.DB) (map[string]map[string]bool, error) { + rows, err := db.QueryContext(ctx, ` SELECT t.relname, a.attname FROM pg_constraint c JOIN pg_class t ON c.conrelid = t.oid @@ -259,8 +273,8 @@ func postgresUniqueMap(db *sql.DB) (map[string]map[string]bool, error) { } // postgresCheckMap returns map[table][column]=[]values for single-column CHECK IN constraints. -func postgresCheckMap(db *sql.DB) (map[string]map[string][]string, error) { - rows, err := db.Query(` +func postgresCheckMap(ctx context.Context, db *sql.DB) (map[string]map[string][]string, error) { + rows, err := db.QueryContext(ctx, ` SELECT t.relname, a.attname, pg_get_constraintdef(c.oid) FROM pg_constraint c JOIN pg_class t ON c.conrelid = t.oid @@ -291,8 +305,8 @@ func postgresCheckMap(db *sql.DB) (map[string]map[string][]string, error) { } // postgresRangeMap returns map[table][column]=rangeConstraint for CHECK (col >= N AND col <= M). -func postgresRangeMap(db *sql.DB) (map[string]map[string]rangeConstraint, error) { - rows, err := db.Query(` +func postgresRangeMap(ctx context.Context, db *sql.DB) (map[string]map[string]rangeConstraint, error) { + rows, err := db.QueryContext(ctx, ` SELECT t.relname, a.attname, pg_get_constraintdef(c.oid) FROM pg_constraint c JOIN pg_class t ON c.conrelid = t.oid @@ -368,8 +382,8 @@ func parsePostgresCheckValues(clause string) []string { return values } -func postgresEnumValues(db *sql.DB, typeName string) ([]string, error) { - rows, err := db.Query(` +func postgresEnumValues(ctx context.Context, db *sql.DB, typeName string) ([]string, error) { + rows, err := db.QueryContext(ctx, ` SELECT e.enumlabel FROM pg_type t JOIN pg_enum e ON t.oid = e.enumtypid @@ -395,8 +409,8 @@ func isPostgresSerialDefault(value string) bool { return strings.HasPrefix(value, "nextval(") } -func postgresIndexMap(db *sql.DB, uniqueMap map[string]map[string]bool) (map[string][]Index, error) { - rows, err := db.Query(postgresIndexQuery()) +func postgresIndexMap(ctx context.Context, db *sql.DB, uniqueMap map[string]map[string]bool) (map[string][]Index, error) { + rows, err := db.QueryContext(ctx, postgresIndexQuery()) if err != nil { return nil, fmt.Errorf("failed to query indexes: %w", err) } @@ -443,8 +457,8 @@ func postgresIndexQuery() string { GROUP BY t.relname, i.relname, ix.indisunique` } -func postgresCommentMaps(db *sql.DB) (map[string]string, map[string]map[string]string, error) { - tableRows, err := db.Query(` +func postgresCommentMaps(ctx context.Context, db *sql.DB) (map[string]string, map[string]map[string]string, error) { + tableRows, err := db.QueryContext(ctx, ` SELECT c.relname, obj_description(c.oid, 'pg_class') FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace @@ -465,7 +479,7 @@ func postgresCommentMaps(db *sql.DB) (map[string]string, map[string]map[string]s tableComments[table] = comment } - columnRows, err := db.Query(` + columnRows, err := db.QueryContext(ctx, ` SELECT c.relname, a.attname, col_description(c.oid, a.attnum) FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace diff --git a/internal/db/read_scope.go b/internal/db/read_scope.go new file mode 100644 index 0000000..9eb376c --- /dev/null +++ b/internal/db/read_scope.go @@ -0,0 +1,189 @@ +package db + +import ( + "context" + "database/sql" + "database/sql/driver" + "errors" + "fmt" + "strings" + "sync" + "time" + + "github.com/go-sql-driver/mysql" + "github.com/jackc/pgx/v5/pgconn" +) + +// ReadLimits bound a read seedstorm runs against a database it must not +// disturb: the server enforces both timeouts. +type ReadLimits struct { + // StatementTimeout ends a statement that runs longer (0: no limit). + StatementTimeout time.Duration + // LockTimeout ends a statement waiting this long for a lock, so a read + // never queues behind DDL and holds up every writer queued behind it + // (0: wait as the server does). + LockTimeout time.Duration + // Concurrency is how many statements run at once (<= 1: one). + Concurrency int +} + +// DefaultCountLimits are the limits of row counts: no statement timeout (a +// long exact count is what the user asked for) but never wait behind a lock. +var DefaultCountLimits = ReadLimits{LockTimeout: 2 * time.Second} + +// Querier is what a read runs statements on. +type Querier interface { + QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) +} + +// ReadOnce runs fn in its own short read-only transaction: the server refuses +// any write, the transaction ends as soon as fn returns (a long-lived snapshot +// would hold back vacuum or purge), and the limits apply on the server. On +// MySQL a cancelled ctx also kills the query on the server: the driver only +// closes its connection, and the server would keep running it. +func ReadOnce(ctx context.Context, conn *sql.DB, dbType string, lim ReadLimits, fn func(ctx context.Context, q Querier) error) error { + if dbType == "mysql" { + return readOnceMySQL(ctx, conn, lim, fn) + } + tx, err := conn.BeginTx(ctx, &sql.TxOptions{ReadOnly: true}) + if err != nil { + return err + } + defer tx.Rollback() //nolint:errcheck // after Commit this is a no-op + if lim.StatementTimeout > 0 { + if _, err := tx.ExecContext(ctx, fmt.Sprintf("SET LOCAL statement_timeout = %d", lim.StatementTimeout.Milliseconds())); err != nil { + return err + } + } + if lim.LockTimeout > 0 { + if _, err := tx.ExecContext(ctx, fmt.Sprintf("SET LOCAL lock_timeout = %d", lim.LockTimeout.Milliseconds())); err != nil { + return err + } + } + if err := fn(ctx, tx); err != nil { + return err + } + return tx.Commit() +} + +func readOnceMySQL(ctx context.Context, conn *sql.DB, lim ReadLimits, fn func(ctx context.Context, q Querier) error) error { + c, err := conn.Conn(ctx) + if err != nil { + return err + } + var id int64 + if err := c.QueryRowContext(ctx, "SELECT CONNECTION_ID()").Scan(&id); err != nil { + _ = c.Close() + return err + } + // Session limits are reset before the connection returns to the pool, or + // the connection is discarded: later statements (seeding's own scans) must + // never inherit them. + limited := lim.StatementTimeout > 0 || lim.LockTimeout > 0 + defer func() { + if limited { + reset, cancel := context.WithTimeout(context.Background(), 3*time.Second) + _, rerr := c.ExecContext(reset, "SET SESSION max_execution_time = DEFAULT, lock_wait_timeout = DEFAULT") + cancel() + if rerr != nil { + _ = c.Raw(func(any) error { return driver.ErrBadConn }) + } + } + _ = c.Close() + }() + if lim.StatementTimeout > 0 { + if _, err := c.ExecContext(ctx, fmt.Sprintf("SET SESSION max_execution_time = %d", lim.StatementTimeout.Milliseconds())); err != nil { + return err + } + } + if lim.LockTimeout > 0 { + secs := max(1, int64((lim.LockTimeout+time.Second-1)/time.Second)) + if _, err := c.ExecContext(ctx, fmt.Sprintf("SET SESSION lock_wait_timeout = %d, innodb_lock_wait_timeout = %d", secs, secs)); err != nil { + return err + } + } + + finished := make(chan struct{}) + var killer sync.WaitGroup + killer.Add(1) + go func() { + defer killer.Done() + select { + case <-finished: + case <-ctx.Done(): + kill, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + _, _ = conn.ExecContext(kill, fmt.Sprintf("KILL QUERY %d", id)) + } + }() + defer func() { + close(finished) + killer.Wait() + }() + + tx, err := c.BeginTx(ctx, &sql.TxOptions{ReadOnly: true}) + if err != nil { + return err + } + defer tx.Rollback() //nolint:errcheck // after Commit this is a no-op + if err := fn(ctx, tx); err != nil { + return err + } + return tx.Commit() +} + +// ReadOutcome is how a bounded read ended. +type ReadOutcome string + +const ( + OutcomeOK ReadOutcome = "ok" + OutcomeTimedOut ReadOutcome = "timed out" + OutcomeLocked ReadOutcome = "locked" + OutcomeCancelled ReadOutcome = "cancelled" + OutcomeFailed ReadOutcome = "failed" +) + +// ReadOutcomeOf classifies a read's error. ctx is the read's context: a +// Postgres statement cancel is a timeout unless ctx was cancelled. +func ReadOutcomeOf(ctx context.Context, err error) ReadOutcome { + if err == nil { + return OutcomeOK + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return OutcomeCancelled + } + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + switch pgErr.Code { + case "57014": // query_canceled: statement_timeout, or a user cancel + if ctx.Err() != nil { + return OutcomeCancelled + } + return OutcomeTimedOut + case "55P03": // lock_not_available: lock_timeout + return OutcomeLocked + case "40001": + // A hot standby cancels reads that conflict with replication. + if strings.Contains(pgErr.Message, "conflict with recovery") { + return OutcomeCancelled + } + } + } + var myErr *mysql.MySQLError + if errors.As(err, &myErr) { + switch myErr.Number { + case 3024: // query execution was interrupted, maximum statement execution time exceeded + return OutcomeTimedOut + case 1317: // query execution was interrupted (KILL QUERY) + return OutcomeCancelled + case 1205: // lock wait timeout exceeded (metadata or row locks) + return OutcomeLocked + } + } + if ctx.Err() != nil { + return OutcomeCancelled + } + return OutcomeFailed +} diff --git a/internal/db/server_info.go b/internal/db/server_info.go new file mode 100644 index 0000000..23bd1c1 --- /dev/null +++ b/internal/db/server_info.go @@ -0,0 +1,96 @@ +package db + +import ( + "context" + "database/sql" + "fmt" + "strings" +) + +// ServerInfo is what a database server says about its capacity, read-only. +type ServerInfo struct { + Engine string `json:"engine"` + Version string `json:"version"` + MaxConnections int `json:"maxConnections"` + UsedConnections int `json:"usedConnections"` + // Memory the server reserves for its buffers. + SharedBuffersBytes int64 `json:"sharedBuffersBytes,omitempty"` + BufferPoolBytes int64 `json:"bufferPoolBytes,omitempty"` + LogBufferBytes int64 `json:"logBufferBytes,omitempty"` + MaxAllowedPacket int64 `json:"maxAllowedPacket,omitempty"` + // UsedBytes is the current database's size on disk. + UsedBytes int64 `json:"usedBytes"` + // Replica is a read-only standby (Postgres recovery, MySQL read_only). + Replica bool `json:"replica"` + // ServerID is equal for two databases on the same server and never names + // a database. + ServerID string `json:"serverId"` +} + +// DetectServer reads the server's capacity in a read-only transaction. +func DetectServer(ctx context.Context, conn *sql.DB, dbType string) (ServerInfo, error) { + var info ServerInfo + err := ReadOnce(ctx, conn, dbType, ReadLimits{}, func(ctx context.Context, q Querier) error { + if dbType == "mysql" { + return detectMySQL(ctx, q, &info) + } + return detectPostgres(ctx, q, &info) + }) + if err != nil { + return ServerInfo{}, fmt.Errorf("read server capacity: %w", err) + } + return info, nil +} + +func detectPostgres(ctx context.Context, q Querier, info *ServerInfo) error { + info.Engine = "postgres" + return q.QueryRowContext(ctx, ` + SELECT current_setting('server_version'), + current_setting('max_connections')::int, + (SELECT count(*) FROM pg_stat_activity)::int, + pg_size_bytes(current_setting('shared_buffers')), + pg_database_size(current_database()), + pg_is_in_recovery(), + pg_postmaster_start_time()::text || '/' || COALESCE(inet_server_port()::text, '')`). + Scan(&info.Version, &info.MaxConnections, &info.UsedConnections, &info.SharedBuffersBytes, &info.UsedBytes, &info.Replica, &info.ServerID) +} + +func detectMySQL(ctx context.Context, q Querier, info *ServerInfo) error { + info.Engine = "mysql" + var readOnly int + if err := q.QueryRowContext(ctx, ` + SELECT @@version, @@max_connections, @@innodb_buffer_pool_size, @@innodb_log_buffer_size, + @@max_allowed_packet, @@read_only, @@server_uuid`). + Scan(&info.Version, &info.MaxConnections, &info.BufferPoolBytes, &info.LogBufferBytes, &info.MaxAllowedPacket, &readOnly, &info.ServerID); err != nil { + return err + } + info.Replica = readOnly == 1 + var name string + if err := q.QueryRowContext(ctx, `SHOW STATUS LIKE 'Threads_connected'`).Scan(&name, &info.UsedConnections); err != nil { + return err + } + var used sql.NullInt64 + if err := q.QueryRowContext(ctx, ` + SELECT SUM(COALESCE(DATA_LENGTH, 0) + COALESCE(INDEX_LENGTH, 0)) + FROM information_schema.TABLES WHERE TABLE_SCHEMA = DATABASE()`).Scan(&used); err != nil { + return err + } + info.UsedBytes = used.Int64 + info.Version = strings.TrimSpace(info.Version) + return nil +} + +// ConnectionUsage returns the server's connection limit and how many are open. +func ConnectionUsage(ctx context.Context, conn *sql.DB, dbType string) (maxConns, used int, err error) { + err = ReadOnce(ctx, conn, dbType, DefaultCountLimits, func(ctx context.Context, q Querier) error { + if dbType == "mysql" { + if err := q.QueryRowContext(ctx, `SELECT @@max_connections`).Scan(&maxConns); err != nil { + return err + } + var name string + return q.QueryRowContext(ctx, `SHOW STATUS LIKE 'Threads_connected'`).Scan(&name, &used) + } + return q.QueryRowContext(ctx, `SELECT current_setting('max_connections')::int, (SELECT count(*) FROM pg_stat_activity)::int`).Scan(&maxConns, &used) + }) + return maxConns, used, err +} diff --git a/internal/db/stats.go b/internal/db/stats.go index e98909c..7a706e5 100644 --- a/internal/db/stats.go +++ b/internal/db/stats.go @@ -13,13 +13,22 @@ const UnknownCount int64 = -1 // ListTableColumns returns every base table in the connection's schema with its // column names in ordinal order. It is a light alternative to Introspect for // callers that only need names. -func ListTableColumns(ctx context.Context, conn *sql.DB, dbType string) (map[string][]string, error) { +func ListTableColumns(ctx context.Context, conn *sql.DB, dbType string) (out map[string][]string, err error) { + err = ReadOnce(ctx, conn, dbType, ReadLimits{}, func(ctx context.Context, q Querier) error { + out, err = listTableColumns(ctx, q, dbType) + return err + }) + return out, err +} + +func listTableColumns(ctx context.Context, conn Querier, dbType string) (map[string][]string, error) { query := ` SELECT c.table_name, c.column_name FROM information_schema.columns c JOIN information_schema.tables t ON t.table_schema = c.table_schema AND t.table_name = c.table_name WHERE c.table_schema = 'public' AND t.table_type = 'BASE TABLE' + AND c.table_name NOT IN (` + postgresPartitionNames + `) ORDER BY c.table_name, c.ordinal_position` if dbType == "mysql" { query = ` @@ -48,12 +57,25 @@ func ListTableColumns(ctx context.Context, conn *sql.DB, dbType string) (map[str // GetTableSizes returns on-disk bytes (data plus indexes) per base table. // MySQL reports cached statistics, so treat its numbers as approximate. -func GetTableSizes(ctx context.Context, conn *sql.DB, dbType string) (map[string]int64, error) { +func GetTableSizes(ctx context.Context, conn *sql.DB, dbType string) (out map[string]int64, err error) { + err = ReadOnce(ctx, conn, dbType, ReadLimits{}, func(ctx context.Context, q Querier) error { + out, err = getTableSizes(ctx, q, dbType) + return err + }) + return out, err +} + +func getTableSizes(ctx context.Context, conn Querier, dbType string) (map[string]int64, error) { + // A partitioned table's size is the sum of its partitions (its own + // relation is empty); partitions are not listed separately. query := ` - SELECT c.relname, pg_total_relation_size(c.oid) + SELECT c.relname, + CASE WHEN c.relkind = 'p' + THEN COALESCE((SELECT SUM(pg_total_relation_size(pt.relid))::bigint FROM pg_partition_tree(c.oid) pt WHERE pt.isleaf), 0) + ELSE pg_total_relation_size(c.oid) END FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace - WHERE n.nspname = 'public' AND c.relkind IN ('r', 'p')` + WHERE n.nspname = 'public' AND c.relkind IN ('r', 'p') AND NOT c.relispartition` if dbType == "mysql" { query = ` SELECT TABLE_NAME, COALESCE(DATA_LENGTH, 0) + COALESCE(INDEX_LENGTH, 0) @@ -71,13 +93,16 @@ func GetTableSizes(ctx context.Context, conn *sql.DB, dbType string) (map[string // without ANALYZE, over planner statistics (reltuples, 0 or -1 before the first // ANALYZE depending on version). MySQL 8 caches information_schema statistics // for up to a day, so the session asks for fresh ones. -func GetEstimatedRowCounts(ctx context.Context, conn *sql.DB, dbType string) (map[string]int64, error) { +func GetEstimatedRowCounts(ctx context.Context, conn *sql.DB, dbType string) (out map[string]int64, err error) { + err = ReadOnce(ctx, conn, dbType, ReadLimits{}, func(ctx context.Context, q Querier) error { + out, err = getEstimatedRowCounts(ctx, q, dbType) + return err + }) + return out, err +} + +func getEstimatedRowCounts(ctx context.Context, c Querier, dbType string) (map[string]int64, error) { if dbType == "mysql" { - c, err := conn.Conn(ctx) - if err != nil { - return nil, fmt.Errorf("read estimated row counts: %w", err) - } - defer func() { _ = c.Close() }() // MySQL 5.7 has no such variable and never caches: ignore the error. _, _ = c.ExecContext(ctx, "SET SESSION information_schema_stats_expiry = 0") rows, err := c.QueryContext(ctx, ` @@ -89,18 +114,28 @@ func GetEstimatedRowCounts(ctx context.Context, conn *sql.DB, dbType string) (ma } return collectNameInt64(rows, "estimated row counts") } - return scanNameInt64(ctx, conn, ` + // A partitioned table's estimate is the sum over its leaf partitions, + // unknown when any of them has none. + return scanNameInt64(ctx, c, ` + WITH leaf AS ( + SELECT c.oid, + CASE WHEN COALESCE(s.n_live_tup, 0) > 0 THEN s.n_live_tup + WHEN c.reltuples > 0 THEN c.reltuples::bigint + ELSE -1 END AS est + FROM pg_class c + LEFT JOIN pg_stat_user_tables s ON s.relid = c.oid + ) SELECT c.relname, - CASE WHEN COALESCE(s.n_live_tup, 0) > 0 THEN s.n_live_tup - WHEN c.reltuples > 0 THEN c.reltuples::bigint - ELSE -1 END + CASE WHEN c.relkind = 'p' THEN + COALESCE((SELECT CASE WHEN MIN(l.est) < 0 THEN -1 ELSE SUM(l.est)::bigint END + FROM pg_partition_tree(c.oid) pt JOIN leaf l ON l.oid = pt.relid WHERE pt.isleaf), -1) + ELSE (SELECT l.est FROM leaf l WHERE l.oid = c.oid) END FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace - LEFT JOIN pg_stat_user_tables s ON s.relid = c.oid - WHERE n.nspname = 'public' AND c.relkind IN ('r', 'p')`, "estimated row counts") + WHERE n.nspname = 'public' AND c.relkind IN ('r', 'p') AND NOT c.relispartition`, "estimated row counts") } -func scanNameInt64(ctx context.Context, conn *sql.DB, query, what string) (map[string]int64, error) { +func scanNameInt64(ctx context.Context, conn Querier, query, what string) (map[string]int64, error) { rows, err := conn.QueryContext(ctx, query) if err != nil { return nil, fmt.Errorf("read %s: %w", what, err) @@ -132,7 +167,10 @@ func Identity(ctx context.Context, conn *sql.DB, dbType string) (string, error) query = `SELECT CONCAT(COALESCE(DATABASE(), ''), '@', @@server_uuid)` } var id string - if err := conn.QueryRowContext(ctx, query).Scan(&id); err != nil { + err := ReadOnce(ctx, conn, dbType, ReadLimits{}, func(ctx context.Context, q Querier) error { + return q.QueryRowContext(ctx, query).Scan(&id) + }) + if err != nil { return "", fmt.Errorf("read database identity: %w", err) } return dbType + ":" + id, nil diff --git a/internal/db/types.go b/internal/db/types.go index 454a630..4585ddc 100644 --- a/internal/db/types.go +++ b/internal/db/types.go @@ -6,6 +6,29 @@ type Table struct { Columns []Column Indexes []Index Comment string + // Partition describes a Postgres partitioned table; its partitions are + // not listed as tables of their own. + Partition *Partitioning +} + +// Partitioning is a partitioned table's key and the bounds of its partitions. +type Partitioning struct { + Strategy string // range | list | hash + // Columns are the key columns; empty entries are expressions. + Columns []string + // Ranges are the FROM/TO bounds of range partitions, as Postgres prints + // them (quoted literals, MINVALUE, MAXVALUE). + Ranges []PartitionRange + // Values are the listed values of list partitions. + Values []string + // Default reports a DEFAULT partition, which accepts any key. + Default bool +} + +// PartitionRange is one range partition's bounds for a single-column key. +type PartitionRange struct { + From string + To string } // Column represents a column in a database table. diff --git a/internal/faker/catalog.go b/internal/faker/catalog.go index 0e17fc5..abd0512 100644 --- a/internal/faker/catalog.go +++ b/internal/faker/catalog.go @@ -57,6 +57,8 @@ var catalog = []Generator{ {Name: "date", Expr: "date", Category: "Dates", Description: "Date (YYYY-MM-DD)"}, {Name: "time", Expr: "time", Category: "Dates", Description: "Time of day (HH:MM:SS)"}, {Name: "datetime", Expr: "datetime", Category: "Dates", Description: "Timestamp"}, + {Name: "daterange", Expr: "daterange(2025-01-01,2026-01-01)", Category: "Dates", Description: "Date from (inclusive) to (exclusive)", Params: []string{"from", "to"}}, + {Name: "datetimerange", Expr: "datetimerange(2025-01-01 00:00:00,2026-01-01 00:00:00)", Category: "Dates", Description: "Timestamp from (inclusive) to (exclusive)", Params: []string{"from", "to"}}, {Name: "uuid", Expr: "uuid", Category: "Identifiers", Description: "Random UUID"}, {Name: "json", Expr: "json", Category: "Identifiers", Description: "Small JSON object"}, } @@ -111,6 +113,9 @@ func BuildSchema(dbType string, tables []db.Table) *schema.Schema { st.Unique = append(st.Unique, append([]string(nil), idx.Columns...)) } } + if p := t.Partition; p != nil { + applyPartitioning(&st, t, p) + } out.Tables[t.Name] = st } return out diff --git a/internal/faker/existing.go b/internal/faker/existing.go index 849bd48..027c4f8 100644 --- a/internal/faker/existing.go +++ b/internal/faker/existing.go @@ -1,6 +1,7 @@ package faker import ( + "context" "database/sql" "fmt" "sort" @@ -84,7 +85,7 @@ func (e existingState) advanceSequences(table schema.Table, tableName string, n // loadExistingState reads existing primary keys, UNIQUE tuples and sequence // maxima for the tables about to be generated. preloaded names tables whose PK // pools were read (their maximum id is known); overrides are the value rules. -func loadExistingState(conn *sql.DB, targetTables []string, tables map[string]schema.Table, dbType string, preloaded map[string]bool, overrides Overrides) (existingState, error) { +func loadExistingState(ctx context.Context, conn *sql.DB, targetTables []string, tables map[string]schema.Table, dbType string, preloaded map[string]bool, overrides Overrides) (existingState, error) { state := existingState{ keys: make(map[string]*keySet), seqStart: make(map[string]map[string]int), @@ -96,7 +97,7 @@ func loadExistingState(conn *sql.DB, targetTables []string, tables map[string]sc continue } if !preloaded[tableName] || !keysCannotCollide(table) { - keys, err := loadExistingKeys(conn, tableName, table, dbType) + keys, err := loadExistingKeys(ctx, conn, tableName, table, dbType) if err != nil { return state, err } @@ -104,12 +105,12 @@ func loadExistingState(conn *sql.DB, targetTables []string, tables map[string]sc state.keys[tableName] = keys } } - tuples, err := loadUniqueTuples(conn, tableName, table, dbType, overrides[tableName]) + tuples, err := loadUniqueTuples(ctx, conn, tableName, table, dbType, overrides[tableName]) if err != nil { return state, err } state.unique[tableName] = tuples - starts, err := loadSequenceStarts(conn, tableName, table, dbType) + starts, err := loadSequenceStarts(ctx, conn, tableName, table, dbType) if err != nil { return state, err } @@ -147,13 +148,13 @@ func isIntegerType(t string) bool { return false } -func loadExistingKeys(conn *sql.DB, tableName string, table schema.Table, dbType string) (*keySet, error) { +func loadExistingKeys(ctx context.Context, conn *sql.DB, tableName string, table schema.Table, dbType string) (*keySet, error) { pkCols := sortedPKColumns(table) if len(pkCols) == 0 { return nil, nil } keys := newKeySet(0) - err := scanStoredRows(conn, tableName, pkCols, dbType, func(row map[string]interface{}) { + err := scanStoredRows(ctx, conn, tableName, pkCols, dbType, func(row map[string]interface{}) { keys.Add(compositePKKey(row, table)) }) return keys, err @@ -161,7 +162,7 @@ func loadExistingKeys(conn *sql.DB, tableName string, table schema.Table, dbType // loadUniqueTuples reads stored tuples of every UNIQUE group. A lone sequence // column no rule rewrites is skipped: new values continue past its maximum. -func loadUniqueTuples(conn *sql.DB, tableName string, table schema.Table, dbType string, overrides map[string]ColumnOverride) (map[string]*keySet, error) { +func loadUniqueTuples(ctx context.Context, conn *sql.DB, tableName string, table schema.Table, dbType string, overrides map[string]ColumnOverride) (map[string]*keySet, error) { out := make(map[string]*keySet) for _, group := range validGroups(table) { if len(group) == 1 && table.Columns[group[0]].Faker == uniqueSequenceFaker { @@ -170,7 +171,7 @@ func loadUniqueTuples(conn *sql.DB, tableName string, table schema.Table, dbType } } tuples := newKeySet(0) - err := scanStoredRows(conn, tableName, group, dbType, func(row map[string]interface{}) { + err := scanStoredRows(ctx, conn, tableName, group, dbType, func(row map[string]interface{}) { if !tupleHasNull(group, row) { tuples.Add(uniqueTupleKey(table, group, row)) } @@ -185,7 +186,7 @@ func loadUniqueTuples(conn *sql.DB, tableName string, table schema.Table, dbType // scanStoredRows reads the given columns of every stored row, normalising // driver values, and hands each row to fn. A nil conn has no stored rows. -func scanStoredRows(conn *sql.DB, tableName string, cols []string, dbType string, fn func(map[string]interface{})) error { +func scanStoredRows(ctx context.Context, conn *sql.DB, tableName string, cols []string, dbType string, fn func(map[string]interface{})) error { if conn == nil { return nil } @@ -194,7 +195,7 @@ func scanStoredRows(conn *sql.DB, tableName string, cols []string, dbType string quoted[i] = db.QuoteIdent(c, dbType) } query := fmt.Sprintf("SELECT %s FROM %s", strings.Join(quoted, ", "), db.QuoteIdent(tableName, dbType)) //nolint:gosec - rows, err := conn.Query(query) + rows, err := conn.QueryContext(ctx, query) if err != nil { return fmt.Errorf("read existing rows of %s: %w", tableName, err) } @@ -217,7 +218,7 @@ func scanStoredRows(conn *sql.DB, tableName string, cols []string, dbType string return rows.Err() } -func loadSequenceStarts(conn *sql.DB, tableName string, table schema.Table, dbType string) (map[string]int, error) { +func loadSequenceStarts(ctx context.Context, conn *sql.DB, tableName string, table schema.Table, dbType string) (map[string]int, error) { starts := make(map[string]int) for _, colName := range sortedColumnNames(table) { col := table.Columns[colName] @@ -230,7 +231,7 @@ func loadSequenceStarts(conn *sql.DB, tableName string, table schema.Table, dbTy } var maxVal interface{} query := fmt.Sprintf("SELECT MAX(%s) FROM %s", db.QuoteIdent(colName, dbType), db.QuoteIdent(tableName, dbType)) //nolint:gosec - if err := conn.QueryRow(query).Scan(&maxVal); err != nil { + if err := conn.QueryRowContext(ctx, query).Scan(&maxVal); err != nil { return nil, fmt.Errorf("read max %s.%s: %w", tableName, colName, err) } starts[colName] = sequenceStartAfter(col.Type, normalizeScanned(maxVal)) diff --git a/internal/faker/faker.go b/internal/faker/faker.go index b0acfbf..69102a5 100644 --- a/internal/faker/faker.go +++ b/internal/faker/faker.go @@ -1,6 +1,7 @@ package faker import ( + "context" "database/sql" "fmt" "regexp" @@ -131,11 +132,11 @@ func uniqueSequenceValue(colType string, i int) interface{} { // queryExistingPKs reads the PK pools of sortedTables and records in sampled // which tables were too large to keep whole. -func (gen generator) queryExistingPKs(conn *sql.DB, sortedTables []string, tables map[string]schema.Table, generatedPKs map[string][]interface{}, dbType string, sampled map[string]bool) error { +func (gen generator) queryExistingPKs(ctx context.Context, conn *sql.DB, sortedTables []string, tables map[string]schema.Table, generatedPKs map[string][]interface{}, dbType string, sampled map[string]bool) error { for _, tableName := range sortedTables { table := tables[tableName] for _, colName := range sortedPKColumns(table) { - wasSampled, err := gen.scanPKs(conn, tableName, colName, generatedPKs, dbType) + wasSampled, err := gen.scanPKs(ctx, conn, tableName, colName, generatedPKs, dbType) if err != nil { return err } @@ -147,14 +148,19 @@ func (gen generator) queryExistingPKs(conn *sql.DB, sortedTables []string, table return nil } -func (gen generator) scanPKs(conn *sql.DB, tableName, colName string, generatedPKs map[string][]interface{}, dbType string) (bool, error) { - rows, err := conn.Query(fmt.Sprintf("SELECT %s FROM %s", db.QuoteIdent(colName, dbType), db.QuoteIdent(tableName, dbType))) //nolint:gosec +func (gen generator) scanPKs(ctx context.Context, conn *sql.DB, tableName, colName string, generatedPKs map[string][]interface{}, dbType string) (bool, error) { + return gen.scanPool(ctx, conn, tableName, colName, tableName, generatedPKs, dbType) +} + +// scanPool reads a column's stored values into the pool under key. +func (gen generator) scanPool(ctx context.Context, conn *sql.DB, tableName, colName, key string, generatedPKs map[string][]interface{}, dbType string) (bool, error) { + rows, err := conn.QueryContext(ctx, fmt.Sprintf("SELECT %s FROM %s", db.QuoteIdent(colName, dbType), db.QuoteIdent(tableName, dbType))) //nolint:gosec if err != nil { return false, fmt.Errorf("failed to query PKs for %s.%s: %w", tableName, colName, err) } defer rows.Close() - pool := newPoolSampler(generatedPKs[tableName], poolLimit, gen.rnd) + pool := newPoolSampler(generatedPKs[key], poolLimit, gen.rnd) for rows.Next() { var pk interface{} if err := rows.Scan(&pk); err != nil { @@ -165,7 +171,7 @@ func (gen generator) scanPKs(conn *sql.DB, tableName, colName string, generatedP if err := rows.Err(); err != nil { return false, err } - generatedPKs[tableName] = pool.values() + generatedPKs[key] = pool.values() return pool.seen > pool.limit, nil } @@ -393,7 +399,7 @@ func (gen generator) enumerateCompositeFKPKRows(data map[string][]map[string]int if fkTable == "" { return 0, start, false, nil } - pool := generatedPKs[fkTable] + pool := poolFor(generatedPKs, col.FK) if len(pool) == 0 { if col.Nullable { return 0, start, false, nil @@ -523,7 +529,7 @@ func (gen generator) generateValue(col schema.Column, colName, tableName string, parts := strings.SplitN(col.FK, ".", 2) if len(parts) == 2 { fkTable := parts[0] - pks := generatedPKs[fkTable] + pks := poolFor(generatedPKs, col.FK) if len(pks) == 0 { if fkTable == tableName { // Self-referential FKs are resolved after all rows for the @@ -541,6 +547,11 @@ func (gen generator) generateValue(col schema.Column, colName, tableName string, return pks[gen.rnd.Number(0, len(pks)-1)], nil } } + if col.PK && col.PartitionKey && col.Faker != "" { + // A key column that is also the partition key must stay inside the + // partitions; the other key columns keep the row unique. + return gen.generate(col.Faker) + } if col.PK { pk, err := gen.generatePK(col.Type, nextSequentialPK(generatedPKs[tableName])) if err != nil { @@ -794,7 +805,7 @@ var knownFakers = map[string]bool{ // knownParamFakers is the set of valid faker functions that take arguments. var knownParamFakers = map[string]bool{ - "number": true, "price": true, "randomstring": true, + "number": true, "price": true, "randomstring": true, "daterange": true, "datetimerange": true, "paragraph": true, "float64": true, "lexify": true, "numerify": true, } @@ -905,6 +916,24 @@ func (gen generator) generate(fakerStr string) (interface{}, error) { return gen.rnd.Paragraph(spec.count, 3, 8, " "), nil case "float64": return gen.rnd.Float64(), nil + case "daterange", "datetimerange": + layout := dateLayout + if spec.name == "datetimerange" { + layout = timeLayout + } + from, to, err := rangeBounds(spec.args, layout) + if err != nil { + return nil, fmt.Errorf("%s: %w", spec.name, err) + } + // DateRange may return its end; keep the upper bound exclusive. + v := gen.rnd.DateRange(from, to) + if !v.Before(to) { + v = from + } + if spec.name == "daterange" { + return v.Format(dateLayout), nil + } + return v, nil case "lexify": // The pattern is raw text, not a comma-separated list. return gen.rnd.Lexify(spec.raw), nil diff --git a/internal/faker/partitions.go b/internal/faker/partitions.go new file mode 100644 index 0000000..6360d49 --- /dev/null +++ b/internal/faker/partitions.go @@ -0,0 +1,224 @@ +package faker + +import ( + "fmt" + "sort" + "strconv" + "strings" + "time" + + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/schema" +) + +// partitionKeyFaker returns the faker that keeps a partitioned table's key +// column inside its partitions' bounds, or why the key cannot be generated. +// An empty faker with no refusal means any value fits (hash partitions, or a +// DEFAULT partition). +func partitionKeyFaker(colType string, p *db.Partitioning) (faker, refuse string) { + if p == nil || p.Strategy == "hash" || p.Default { + return "", "" + } + if len(p.Columns) != 1 { + return "", "partitioned by several columns" + } + if p.Columns[0] == "" { + return "", "partitioned by an expression" + } + t := strings.ToLower(colType) + switch p.Strategy { + case "list": + if len(p.Values) == 0 { + return "", "list-partitioned with no listed values" + } + return "randomstring(" + strings.Join(p.Values, ",") + ")", "" + case "range": + switch { + case isIntegerType(t): + return intRangeFaker(p.Ranges) + case strings.Contains(t, "timestamp") || strings.Contains(t, "datetime"): + return timeRangeFaker(p.Ranges, "datetimerange", timeLayout) + case strings.Contains(t, "date"): + return timeRangeFaker(p.Ranges, "daterange", dateLayout) + } + return "", "range-partitioned on a " + rangeTypeLabel(t) + " column" + } + return "", "partitioned by " + p.Strategy +} + +const ( + dateLayout = "2006-01-02" + timeLayout = "2006-01-02 15:04:05" +) + +func rangeTypeLabel(t string) string { + if strings.Contains(t, "char") || strings.Contains(t, "text") { + return "text" + } + return t +} + +// span is one range partition as numbers (unix seconds for times). +type span struct{ from, to int64 } + +// widest merges touching ranges and returns the widest merged one: generating +// across a gap between partitions would produce keys no partition accepts. +func widest(spans []span) span { + sort.Slice(spans, func(i, j int) bool { return spans[i].from < spans[j].from }) + best, cur := spans[0], spans[0] + for _, s := range spans[1:] { + if s.from <= cur.to { + cur.to = max(cur.to, s.to) + } else { + cur = s + } + if cur.to-cur.from > best.to-best.from { + best = cur + } + } + return best +} + +const openRangeWidth = 1000 + +func intRangeFaker(ranges []db.PartitionRange) (string, string) { + var spans []span + for _, r := range ranges { + from, fok := parseIntBound(r.From) + to, tok := parseIntBound(r.To) + switch { + case !fok && !tok: + continue + case !fok: + from = to - openRangeWidth + case !tok: + to = from + openRangeWidth + } + if from < 0 && r.From == "MINVALUE" && to > 0 { + from = 0 + } + if to > from { + spans = append(spans, span{from, to}) + } + } + if len(spans) == 0 { + return "", "range-partitioned with bounds seedstorm cannot read" + } + w := widest(spans) + return fmt.Sprintf("number(%d,%d)", w.from, w.to-1), "" +} + +func parseIntBound(s string) (int64, bool) { + n, err := strconv.ParseInt(unquoteBound(s), 10, 64) + return n, err == nil +} + +func timeRangeFaker(ranges []db.PartitionRange, name, layout string) (string, string) { + const openYears = 1 + var spans []span + for _, r := range ranges { + from, fok := parseTimeBound(r.From) + to, tok := parseTimeBound(r.To) + switch { + case !fok && !tok: + continue + case !fok: + from = to.AddDate(-openYears, 0, 0) + case !tok: + to = from.AddDate(openYears, 0, 0) + } + if to.After(from) { + spans = append(spans, span{from.Unix(), to.Unix()}) + } + } + if len(spans) == 0 { + return "", "range-partitioned with bounds seedstorm cannot read" + } + w := widest(spans) + return fmt.Sprintf("%s(%s,%s)", name, time.Unix(w.from, 0).UTC().Format(layout), time.Unix(w.to, 0).UTC().Format(layout)), "" +} + +func parseTimeBound(s string) (time.Time, bool) { + v := unquoteBound(s) + for _, layout := range []string{timeLayout, dateLayout, "2006-01-02 15:04:05-07", "2006-01-02 15:04:05.999999"} { + if t, err := time.Parse(layout, v); err == nil { + return t.UTC(), true + } + } + return time.Time{}, false +} + +// unquoteBound strips quotes and a ::type cast from a printed bound literal. +func unquoteBound(s string) string { + s = strings.TrimSpace(s) + if i := strings.LastIndex(s, "::"); i > 0 { + s = s[:i] + } + return strings.Trim(s, "'") +} + +// rangeBounds parses a range faker's two bounds (from inclusive, to exclusive). +func rangeBounds(args []string, layout string) (time.Time, time.Time, error) { + if len(args) != 2 { + return time.Time{}, time.Time{}, fmt.Errorf("want 2 arguments (from, to), got %d", len(args)) + } + from, err := time.Parse(layout, strings.TrimSpace(args[0])) + if err != nil { + return time.Time{}, time.Time{}, fmt.Errorf("bad from: %w", err) + } + to, err := time.Parse(layout, strings.TrimSpace(args[1])) + if err != nil { + return time.Time{}, time.Time{}, fmt.Errorf("bad to: %w", err) + } + if !to.After(from) { + return time.Time{}, time.Time{}, fmt.Errorf("empty range %s to %s", args[0], args[1]) + } + return from, to, nil +} + +// applyPartitioning describes a partitioned table in the schema and points its +// key column at a faker that stays inside the partitions, or marks the table +// unseedable when no faker can. +func applyPartitioning(st *schema.Table, t db.Table, p *db.Partitioning) { + cols := make([]string, len(p.Columns)) + for i, c := range p.Columns { + cols[i] = c + if c == "" { + cols[i] = "expression" + } + } + st.PartitionedBy = p.Strategy + " (" + strings.Join(cols, ", ") + ")" + colType := "" + if len(p.Columns) == 1 { + for _, c := range t.Columns { + if c.Name == p.Columns[0] { + colType = c.Type + } + } + } + faker, refuse := partitionKeyFaker(colType, p) + if refuse != "" { + st.Unseedable = refuse + return + } + if faker != "" { + col := st.Columns[p.Columns[0]] + col.Faker = faker + col.PartitionKey = true + st.Columns[p.Columns[0]] = col + } +} + +// CheckSeedable refuses tables whose rows cannot be generated (see +// schema.Table.Unseedable) unless value rules set some of their columns. Runs +// call it before writing anything, truncation included. +func CheckSeedable(sc *schema.Schema, tables []string, overrides Overrides) error { + for _, name := range tables { + t, ok := sc.Tables[name] + if !ok || t.Unseedable == "" || len(overrides[name]) > 0 { + continue + } + return fmt.Errorf("table %s is %s, which seedstorm cannot generate inside its partitions: add a value rule for its partition key columns in a profile, or ignore the table", name, t.Unseedable) + } + return nil +} diff --git a/internal/faker/partitions_test.go b/internal/faker/partitions_test.go new file mode 100644 index 0000000..33dac27 --- /dev/null +++ b/internal/faker/partitions_test.go @@ -0,0 +1,112 @@ +package faker + +import ( + "strings" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/db" +) + +func TestPartitionKeyFaker_StaysInsideTheBounds(t *testing.T) { + cases := []struct { + name string + colType string + part db.Partitioning + wantFaker string + wantRefuse string + }{ + { + "contiguous date ranges", "date", + db.Partitioning{ + Strategy: "range", Columns: []string{"created_at"}, + Ranges: []db.PartitionRange{{From: "'2026-01-01'", To: "'2027-01-01'"}, {From: "'2025-01-01'", To: "'2026-01-01'"}}, + }, + "daterange(2025-01-01,2027-01-01)", "", + }, + { + "timestamp range", "timestamp without time zone", + db.Partitioning{ + Strategy: "range", Columns: []string{"at"}, + Ranges: []db.PartitionRange{{From: "'2025-01-01 00:00:00'", To: "'2025-02-01 00:00:00'"}}, + }, + "datetimerange(2025-01-01 00:00:00,2025-02-01 00:00:00)", "", + }, + { + "gap between ranges uses the widest one", "integer", + db.Partitioning{ + Strategy: "range", Columns: []string{"n"}, + Ranges: []db.PartitionRange{{From: "0", To: "10"}, {From: "100", To: "1000"}}, + }, + "number(100,999)", "", + }, + { + "open lower bound", "bigint", + db.Partitioning{ + Strategy: "range", Columns: []string{"n"}, + Ranges: []db.PartitionRange{{From: "MINVALUE", To: "500"}}, + }, + "number(0,499)", "", + }, + { + "list", "text", + db.Partitioning{Strategy: "list", Columns: []string{"region"}, Values: []string{"eu-west", "us-east"}}, + "randomstring(eu-west,us-east)", "", + }, + {"default partition accepts anything", "date", db.Partitioning{ + Strategy: "range", Columns: []string{"d"}, + Ranges: []db.PartitionRange{{From: "'2025-01-01'", To: "'2026-01-01'"}}, Default: true, + }, "", ""}, + {"hash accepts anything", "integer", db.Partitioning{Strategy: "hash", Columns: []string{"id"}}, "", ""}, + {"expression key", "timestamp", db.Partitioning{ + Strategy: "range", Columns: []string{""}, + Ranges: []db.PartitionRange{{From: "'2026-01-01'", To: "'2026-02-01'"}}, + }, "", "partitioned by an expression"}, + {"two-column key", "integer", db.Partitioning{ + Strategy: "range", Columns: []string{"a", "b"}, + Ranges: []db.PartitionRange{{From: "1, 1", To: "2, 2"}}, + }, "", "partitioned by several columns"}, + {"range on text", "text", db.Partitioning{ + Strategy: "range", Columns: []string{"s"}, + Ranges: []db.PartitionRange{{From: "'a'", To: "'m'"}}, + }, "", "range-partitioned on a text column"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + faker, refuse := partitionKeyFaker(c.colType, &c.part) + if faker != c.wantFaker { + t.Errorf("faker = %q, want %q", faker, c.wantFaker) + } + if (c.wantRefuse == "") != (refuse == "") || !strings.Contains(refuse, c.wantRefuse) { + t.Errorf("refusal = %q, want containing %q", refuse, c.wantRefuse) + } + }) + } +} + +func TestGenerate_DateRangesStayInsideTheirBounds(t *testing.T) { + from, to := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC), time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC) + for i := 0; i < 500; i++ { + v, err := defaultGen.generate("daterange(2025-01-01,2025-02-01)") + if err != nil { + t.Fatal(err) + } + d, err := time.Parse("2006-01-02", v.(string)) + if err != nil || d.Before(from) || !d.Before(to) { + t.Fatalf("daterange produced %v (%v)", v, err) + } + ts, err := defaultGen.generate("datetimerange(2025-01-01 00:00:00,2025-02-01 00:00:00)") + if err != nil { + t.Fatal(err) + } + if tt := ts.(time.Time); tt.Before(from) || !tt.Before(to) { + t.Fatalf("datetimerange produced %v", tt) + } + } + if !ValidFaker("daterange(2025-01-01,2025-02-01)") || !ValidFaker("datetimerange(2025-01-01 00:00:00,2025-02-01 00:00:00)") { + t.Fatal("range fakers are not recognised") + } + if _, err := defaultGen.generate("daterange(2025-02-01,2025-01-01)"); err == nil { + t.Fatal("an empty range was accepted") + } +} diff --git a/internal/faker/references.go b/internal/faker/references.go new file mode 100644 index 0000000..ebb440a --- /dev/null +++ b/internal/faker/references.go @@ -0,0 +1,78 @@ +package faker + +import ( + "sort" + + "github.com/AxeForging/seedstorm/internal/schema" +) + +// Foreign keys usually reference their parent's single primary key, whose +// values the parent's key pool (keyed by table name) already holds. An FK can +// also reference a UNIQUE column, or one column of a composite key; those +// values live in a pool of their own, keyed by the FK target "table.column". + +// referencedColumns lists, per parent table, the columns some FK references +// that are not the parent's single primary key, sorted. +func referencedColumns(sc *schema.Schema) map[string][]string { + seen := map[string]map[string]bool{} + for _, table := range sc.Tables { + for _, col := range table.Columns { + parent, refCol := splitFK(col.FK) + if parent == "" || refCol == "" { + continue + } + parentTable, ok := sc.Tables[parent] + if !ok || isSolePK(parentTable, refCol) { + continue + } + if _, ok := parentTable.Columns[refCol]; !ok { + continue + } + if seen[parent] == nil { + seen[parent] = map[string]bool{} + } + seen[parent][refCol] = true + } + } + out := make(map[string][]string, len(seen)) + for parent, cols := range seen { + for c := range cols { + out[parent] = append(out[parent], c) + } + sort.Strings(out[parent]) + } + return out +} + +func isSolePK(table schema.Table, colName string) bool { + pks := sortedPKColumns(table) + return len(pks) == 1 && pks[0] == colName +} + +// referencePoolKey is the pool key of a referenced non-key column. +func referencePoolKey(table, col string) string { return table + "." + col } + +// poolFor returns the values an FK may take: the referenced column's own pool +// when it has one, else the parent's key pool. +func poolFor(generatedPKs map[string][]interface{}, fk string) []interface{} { + if pool, ok := generatedPKs[fk]; ok { + return pool + } + parent, _ := splitFK(fk) + return generatedPKs[parent] +} + +// recordReferencedValues adds final rows' values of referenced columns to their +// pools, keeping each pool bounded. +func (g *Stream) recordReferencedValues(tableName string, rows []map[string]interface{}) { + for _, col := range g.refCols[tableName] { + key := referencePoolKey(tableName, col) + pool := g.pks[key] + for _, row := range rows { + if v := row[col]; v != nil { + pool = append(pool, v) + } + } + g.pks[key] = capPool(pool, poolLimit, g.gen.rnd) + } +} diff --git a/internal/faker/references_test.go b/internal/faker/references_test.go new file mode 100644 index 0000000..49002b6 --- /dev/null +++ b/internal/faker/references_test.go @@ -0,0 +1,86 @@ +package faker + +import ( + "fmt" + "testing" + + "github.com/AxeForging/seedstorm/internal/schema" +) + +// referenceSchema has foreign keys whose target is not the parent's single +// primary key: a UNIQUE code column, and one column of a composite key. +func referenceSchema() *schema.Schema { + return &schema.Schema{Tables: map[string]schema.Table{ + "accounts": {Columns: map[string]schema.Column{ + "id": {Type: "integer", PK: true}, + "code": {Type: "varchar(12)", Faker: "uuid", Unique: true}, + }}, + "ledger": {Columns: map[string]schema.Column{ + "id": {Type: "integer", PK: true}, + "account_code": {Type: "varchar(12)", FK: "accounts.code"}, + }}, + "memberships": {Columns: map[string]schema.Column{ + "tenant_id": {Type: "integer", PK: true, Faker: "number(1000,1999)"}, + "member_no": {Type: "integer", PK: true, Faker: "number(500000,599999)"}, + }}, + "badges": {Columns: map[string]schema.Column{ + "id": {Type: "integer", PK: true}, + "member_no": {Type: "integer", FK: "memberships.member_no"}, + }}, + }} +} + +func columnValues(rows []map[string]interface{}, col string) map[string]bool { + out := make(map[string]bool, len(rows)) + for _, r := range rows { + out[fmt.Sprint(r[col])] = true + } + return out +} + +// An FK must hold values of the column it references. It used to draw from the +// parent's primary-key pool, so ledger.account_code got account ids and +// badges.member_no got a mix of tenant ids and member numbers. +func TestGenerate_ForeignKeysUseTheReferencedColumn(t *testing.T) { + order := []string{"accounts", "memberships", "ledger", "badges"} + for _, fork := range []bool{false, true} { + t.Run(fmt.Sprintf("fork=%v", fork), func(t *testing.T) { + g, err := NewStream(referenceSchema(), order, order, nil, "pgx", nil) + if err != nil { + t.Fatal(err) + } + out := map[string][]map[string]interface{}{} + for _, table := range order { + gen := g + if fork { + gen = g.ForkTable(table) + } + err := gen.GenerateChunks(table, 60, 0, false, 25, DefaultGenerateOptions(), func(rows []map[string]interface{}) error { + out[table] = append(out[table], rows...) + return nil + }) + if err != nil { + t.Fatal(err) + } + if fork { + g.MergeTable(gen, table) + } + } + checks := []struct{ child, col, parent, parentCol string }{ + {"ledger", "account_code", "accounts", "code"}, + {"badges", "member_no", "memberships", "member_no"}, + } + for _, c := range checks { + allowed := columnValues(out[c.parent], c.parentCol) + if len(out[c.child]) == 0 { + t.Fatalf("%s generated nothing", c.child) + } + for i, row := range out[c.child] { + if v := fmt.Sprint(row[c.col]); !allowed[v] { + t.Fatalf("%s row %d: %s=%s is not a %s.%s value", c.child, i, c.col, v, c.parent, c.parentCol) + } + } + } + }) + } +} diff --git a/internal/faker/stream.go b/internal/faker/stream.go index ebcd242..f853ede 100644 --- a/internal/faker/stream.go +++ b/internal/faker/stream.go @@ -1,6 +1,7 @@ package faker import ( + "context" "database/sql" "fmt" "math" @@ -24,6 +25,8 @@ type Stream struct { conn *sql.DB dbType string + // ctx stops database reads (preloads, pool redraws) when a run is cancelled. + ctx context.Context // sampled marks preloaded tables whose pool is a sample of the stored keys. sampled map[string]bool // sinceDraw counts rows generated per table since its parents' samples @@ -31,6 +34,9 @@ type Stream struct { sinceDraw map[string]int // offset is how many rows value rules have already numbered, per table. offset map[string]int + // refCols lists, per table, columns FKs reference that are not its single + // primary key; their values are pooled under "table.column". + refCols map[string][]string // gen draws random values; a fork may switch to its own (UseOwnRandom). gen generator // mu guards the maps above while forks are taken and merged (ForkTable). @@ -41,14 +47,30 @@ type Stream struct { // targetTables. With a nil conn nothing is read and the stream starts empty. overrides are the value rules // that will be applied, which decide what stored UNIQUE values must be read. func NewStream(s *schema.Schema, allTables, targetTables []string, conn *sql.DB, dbType string, overrides Overrides) (*Stream, error) { + return NewStreamContext(context.Background(), s, allTables, targetTables, conn, dbType, overrides) +} + +// NewStreamContext is NewStream whose database reads stop when ctx ends. +func NewStreamContext(ctx context.Context, s *schema.Schema, allTables, targetTables []string, conn *sql.DB, dbType string, overrides Overrides) (*Stream, error) { + if err := CheckSeedable(s, targetTables, overrides); err != nil { + return nil, err + } g := &Stream{ sc: s, pks: make(map[string][]interface{}), cursor: make(map[string]int), - conn: conn, dbType: dbType, sampled: make(map[string]bool), sinceDraw: make(map[string]int), offset: make(map[string]int), - gen: defaultGen, + conn: conn, dbType: dbType, ctx: ctx, sampled: make(map[string]bool), sinceDraw: make(map[string]int), offset: make(map[string]int), + gen: defaultGen, refCols: referencedColumns(s), + } + for table, cols := range g.refCols { + for _, col := range cols { + g.pks[referencePoolKey(table, col)] = []interface{}{} + } } preloaded := make(map[string]bool, len(targetTables)) if conn != nil { - if err := g.gen.queryExistingPKs(conn, allTables, s.Tables, g.pks, dbType, g.sampled); err != nil { + if err := g.gen.queryExistingPKs(ctx, conn, allTables, s.Tables, g.pks, dbType, g.sampled); err != nil { + return nil, err + } + if err := g.queryReferencedValues(allTables); err != nil { return nil, err } for _, t := range allTables { @@ -61,7 +83,7 @@ func NewStream(s *schema.Schema, allTables, targetTables []string, conn *sql.DB, } } var err error - if g.existing, err = loadExistingState(conn, targetTables, s.Tables, dbType, preloaded, overrides); err != nil { + if g.existing, err = loadExistingState(ctx, conn, targetTables, s.Tables, dbType, preloaded, overrides); err != nil { return nil, err } return g, nil @@ -278,6 +300,7 @@ func (g *Stream) finalize(tableName string, table schema.Table, rows []map[strin return nil, fmt.Errorf("table %s self-reference backfill: %w", tableName, err) } existing.record(table, tableName, rows) + g.recordReferencedValues(tableName, rows) g.sinceDraw[tableName] += len(rows) g.pks[tableName] = capPool(g.pks[tableName], poolLimit, g.gen.rnd) return rows, nil @@ -300,9 +323,15 @@ func (g *Stream) redrawParents(tableName string, table schema.Table) error { } drawn[parent] = true g.pks[parent] = nil - if err := g.gen.queryExistingPKs(g.conn, []string{parent}, g.sc.Tables, g.pks, g.dbType, nil); err != nil { + if err := g.gen.queryExistingPKs(g.ctx, g.conn, []string{parent}, g.sc.Tables, g.pks, g.dbType, nil); err != nil { return fmt.Errorf("table %s: redraw %s keys: %w", tableName, parent, err) } + for _, col := range g.refCols[parent] { + g.pks[referencePoolKey(parent, col)] = []interface{}{} + } + if err := g.queryReferencedValues([]string{parent}); err != nil { + return fmt.Errorf("table %s: redraw %s values: %w", tableName, parent, err) + } } return nil } @@ -348,20 +377,32 @@ func (g *Stream) ForkTable(tableName string) *Stream { g.mu.Lock() defer g.mu.Unlock() child := &Stream{ - sc: g.sc, existing: g.existing, conn: g.conn, dbType: g.dbType, sampled: g.sampled, gen: g.gen, + sc: g.sc, existing: g.existing, conn: g.conn, dbType: g.dbType, ctx: g.ctx, sampled: g.sampled, gen: g.gen, pks: make(map[string][]interface{}), cursor: map[string]int{tableName: g.cursor[tableName]}, sinceDraw: map[string]int{tableName: g.sinceDraw[tableName]}, offset: map[string]int{tableName: g.offset[tableName]}, + refCols: g.refCols, } if pool, ok := g.pks[tableName]; ok { child.pks[tableName] = pool } + for _, col := range g.refCols[tableName] { + key := referencePoolKey(tableName, col) + child.pks[key] = g.pks[key] + } for _, colName := range sortedColumnNames(g.sc.Tables[tableName]) { - parent, _ := splitFK(g.sc.Tables[tableName].Columns[colName].FK) - if pool, ok := g.pks[parent]; ok && parent != "" { + fk := g.sc.Tables[tableName].Columns[colName].FK + parent, _ := splitFK(fk) + if parent == "" { + continue + } + if pool, ok := g.pks[parent]; ok { child.pks[parent] = pool } + if pool, ok := g.pks[fk]; ok { + child.pks[fk] = pool + } } return child } @@ -381,6 +422,10 @@ func (g *Stream) MergeTable(child *Stream, tableName string) { if pool, ok := child.pks[tableName]; ok { g.pks[tableName] = pool } + for _, col := range g.refCols[tableName] { + key := referencePoolKey(tableName, col) + g.pks[key] = child.pks[key] + } g.cursor[tableName] = child.cursor[tableName] g.sinceDraw[tableName] = child.sinceDraw[tableName] g.offset[tableName] = child.offset[tableName] @@ -394,6 +439,9 @@ func (g *Stream) ReleaseTable(tableName string) { g.mu.Lock() defer g.mu.Unlock() delete(g.pks, tableName) + for _, col := range g.refCols[tableName] { + delete(g.pks, referencePoolKey(tableName, col)) + } } // Schema returns the schema the stream generates for. @@ -406,3 +454,15 @@ func (g *Stream) KeyPoolLen(tableName string) int { defer g.mu.Unlock() return len(g.pks[tableName]) } + +// queryReferencedValues reads the stored values of referenced non-key columns. +func (g *Stream) queryReferencedValues(tables []string) error { + for _, table := range tables { + for _, col := range g.refCols[table] { + if _, err := g.gen.scanPool(g.ctx, g.conn, table, col, referencePoolKey(table, col), g.pks, g.dbType); err != nil { + return err + } + } + } + return nil +} diff --git a/internal/faultinject/faultinject_off.go b/internal/faultinject/faultinject_off.go new file mode 100644 index 0000000..cfebd5c --- /dev/null +++ b/internal/faultinject/faultinject_off.go @@ -0,0 +1,11 @@ +//go:build !faultinject + +// Package faultinject makes named steps of a run fail on purpose, for tests +// only. The default build compiles Hit to nothing; build with +// -tags faultinject to enable it (see faultinject_on.go). +package faultinject + +import "context" + +// Hit does nothing in the default build. +func Hit(context.Context, string, string) error { return nil } diff --git a/internal/faultinject/faultinject_on.go b/internal/faultinject/faultinject_on.go new file mode 100644 index 0000000..8803ccd --- /dev/null +++ b/internal/faultinject/faultinject_on.go @@ -0,0 +1,47 @@ +//go:build faultinject + +package faultinject + +import ( + "context" + "fmt" + "os" + "strings" + "sync" +) + +type fault struct{ point, table, mode string } + +var ( + once sync.Once + faults []fault +) + +// Hit fails the named step when SEEDSTORM_FAULT asks for it. The variable +// holds comma-separated point:table:mode entries (table * matches any); +// mode is panic, error, or hang (block until ctx ends). +func Hit(ctx context.Context, point, table string) error { + once.Do(func() { + for _, entry := range strings.Split(os.Getenv("SEEDSTORM_FAULT"), ",") { + parts := strings.Split(strings.TrimSpace(entry), ":") + if len(parts) == 3 { + faults = append(faults, fault{parts[0], parts[1], parts[2]}) + } + } + }) + for _, f := range faults { + if f.point != point || (f.table != "*" && f.table != table) { + continue + } + switch f.mode { + case "panic": + panic(fmt.Sprintf("injected panic at %s %s", point, table)) + case "error": + return fmt.Errorf("injected failure at %s %s", point, table) + case "hang": + <-ctx.Done() + return ctx.Err() + } + } + return nil +} diff --git a/internal/runerr/runerr.go b/internal/runerr/runerr.go new file mode 100644 index 0000000..3f73ea4 --- /dev/null +++ b/internal/runerr/runerr.go @@ -0,0 +1,108 @@ +// Package runerr gives a run's failure a location every surface (web, CLI, +// TUI) renders the same way: which side, which phase, which table. +package runerr + +import ( + "errors" + "strings" +) + +// Side is the database of a two-sided run (compare, mirror). +const ( + SideSource = "source" + SideTarget = "target" +) + +// Phase is the step of a run that failed. +type Phase string + +const ( + PhaseConnect Phase = "connect" + PhaseIntrospect Phase = "introspect" + PhaseCount Phase = "count" + PhasePlan Phase = "plan" + PhaseTruncate Phase = "truncate" + PhaseGenerate Phase = "generate" + PhaseWrite Phase = "write" + PhaseSync Phase = "sync" +) + +// Error is a failure with its location. Empty fields are unknown. +type Error struct { + Side string `json:"side,omitempty"` + Phase Phase `json:"phase,omitempty"` + Table string `json:"table,omitempty"` + Err error `json:"-"` +} + +func (e *Error) Error() string { + var where []string + for _, part := range []string{e.Side, string(e.Phase), e.Table} { + if part != "" { + where = append(where, part) + } + } + msg := e.Err.Error() + // A located error deeper in the chain prints its own location; this one + // already carries it, so print only that error's cause. + if inner, ok := As(e.Err); ok { + msg = strings.Replace(msg, inner.Error(), inner.Err.Error(), 1) + } + if len(where) == 0 { + return msg + } + return strings.Join(where, " · ") + ": " + msg +} + +func (e *Error) Unwrap() error { return e.Err } + +// As returns the located error in err's chain. +func As(err error) (*Error, bool) { + var e *Error + ok := errors.As(err, &e) + return e, ok +} + +// At locates err in a phase and table. A location already present (set +// closer to the failure) is kept; only missing parts are filled. +func At(phase Phase, table string, err error) error { + if err == nil { + return nil + } + if e, ok := As(err); ok { + c := *e + if c.Phase == "" { + c.Phase = phase + } + if c.Table == "" { + c.Table = table + } + return replace(err, e, &c) + } + return &Error{Phase: phase, Table: table, Err: err} +} + +// OnSide records which database of a two-sided run failed. +func OnSide(side string, err error) error { + if err == nil { + return nil + } + if e, ok := As(err); ok { + if e.Side != "" { + return err + } + c := *e + c.Side = side + return replace(err, e, &c) + } + return &Error{Side: side, Err: err} +} + +// replace swaps the located error for its updated copy when it is the +// outermost error; deeper in a chain, a new located error wraps the chain. +func replace(err error, old, updated *Error) error { + if err == error(old) { + return updated + } + return &Error{Side: updated.Side, Phase: updated.Phase, Table: updated.Table, Err: err} +} diff --git a/internal/runerr/runerr_test.go b/internal/runerr/runerr_test.go new file mode 100644 index 0000000..aaa3b4e --- /dev/null +++ b/internal/runerr/runerr_test.go @@ -0,0 +1,47 @@ +package runerr + +import ( + "errors" + "fmt" + "testing" +) + +func TestError_SaysWhereARunFailed(t *testing.T) { + cause := errors.New(`duplicate key value violates unique constraint "orders_pkey"`) + err := OnSide(SideTarget, At(PhaseWrite, "orders", cause)) + if got, want := err.Error(), `target · write · orders: duplicate key value violates unique constraint "orders_pkey"`; got != want { + t.Fatalf("Error() = %q\nwant %q", got, want) + } + if !errors.Is(err, cause) { + t.Fatal("the cause is not reachable with errors.Is") + } + e, ok := As(err) + if !ok || e.Side != SideTarget || e.Phase != PhaseWrite || e.Table != "orders" { + t.Fatalf("As = %+v, %v", e, ok) + } +} + +// The innermost layer knows best: wrapping twice keeps the first phase and +// table, and only fills what is missing. +func TestAt_DoesNotOverwriteAMoreSpecificLocation(t *testing.T) { + err := At(PhaseGenerate, "", At(PhaseWrite, "users", errors.New("refused"))) + e, _ := As(err) + if e.Phase != PhaseWrite || e.Table != "users" { + t.Fatalf("outer wrap replaced the location: %+v", e) + } + if At(PhaseWrite, "users", nil) != nil || OnSide(SideSource, nil) != nil { + t.Fatal("wrapping nil must stay nil") + } + if got := At(PhaseConnect, "", errors.New("dial tcp: refused")).Error(); got != "connect: dial tcp: refused" { + t.Fatalf("no table: %q", got) + } +} + +// Wrapped by fmt.Errorf between two locations, the location is printed once. +func TestError_LocationPrintedOnceThroughWrapping(t *testing.T) { + inner := At(PhaseWrite, "users", errors.New("refused")) + err := OnSide(SideTarget, fmt.Errorf("mirror: %w", inner)) + if got, want := err.Error(), "target · write · users: mirror: refused"; got != want { + t.Fatalf("Error() = %q, want %q", got, want) + } +} diff --git a/internal/safego/safego.go b/internal/safego/safego.go new file mode 100644 index 0000000..e5b2b3f --- /dev/null +++ b/internal/safego/safego.go @@ -0,0 +1,51 @@ +// Package safego keeps a panic in one goroutine from killing the process: a +// run that panics fails with an error naming where it happened, and the stack +// goes to the log under a short id the error carries. +package safego + +import ( + "crypto/rand" + "encoding/hex" + "fmt" + "runtime/debug" + + "github.com/AxeForging/seedstorm/internal/logging" +) + +// PanicError is a recovered panic. +type PanicError struct { + Where string + Value any + Stack []byte + // ID matches the log line that holds the stack. + ID string +} + +func (p *PanicError) Error() string { + return fmt.Sprintf("internal error in %s (id %s): %v", p.Where, p.ID, p.Value) +} + +// Run calls fn and returns its error, or a *PanicError if it panicked. +func Run(where string, fn func() error) (err error) { + defer Recover(where, &err) + return fn() +} + +// Recover, deferred directly, turns a panic of the surrounding function into a +// *PanicError stored in *errp. +func Recover(where string, errp *error) { + r := recover() + if r == nil { + return + } + p := &PanicError{Where: where, Value: r, Stack: debug.Stack(), ID: newID()} + logging.Log.Error().Str("id", p.ID).Str("where", where).Interface("panic", r).Msg("Recovered from a panic") + logging.Log.Debug().Str("id", p.ID).Str("stack", string(p.Stack)).Msg("Panic stack") + *errp = p +} + +func newID() string { + var b [4]byte + _, _ = rand.Read(b[:]) + return hex.EncodeToString(b[:]) +} diff --git a/internal/safego/safego_test.go b/internal/safego/safego_test.go new file mode 100644 index 0000000..3df0b33 --- /dev/null +++ b/internal/safego/safego_test.go @@ -0,0 +1,54 @@ +package safego + +import ( + "errors" + "regexp" + "strings" + "testing" +) + +func TestRun_TurnsAPanicIntoAnError(t *testing.T) { + err := Run("write users", func() error { + var list []int + i := 3 + _ = list[i] // a real runtime panic, not a hand-made one + return nil + }) + var p *PanicError + if !errors.As(err, &p) { + t.Fatalf("err = %v (%T), want *PanicError", err, err) + } + if p.Where != "write users" || !strings.Contains(err.Error(), "index out of range") { + t.Fatalf("panic error = %+v / %q", p, err) + } + if !regexp.MustCompile(`^[0-9a-f]{8}$`).MatchString(p.ID) { + t.Fatalf("id = %q, want 8 hex chars to find it in the server log", p.ID) + } + if !strings.Contains(string(p.Stack), "safego_test.go") { + t.Fatalf("stack does not point at the panic:\n%s", p.Stack) + } + if strings.Contains(err.Error(), "goroutine") { + t.Fatalf("the message carries the stack; it belongs in the log only: %q", err) + } +} + +func TestRun_PassesErrorsAndSuccessThrough(t *testing.T) { + want := errors.New("plain failure") + if err := Run("x", func() error { return want }); err != want { + t.Fatalf("err = %v, want the function's own error", err) + } + if err := Run("x", func() error { return nil }); err != nil { + t.Fatalf("err = %v, want nil", err) + } +} + +func TestRecover_InADeferKeepsAnExistingError(t *testing.T) { + f := func() (err error) { + defer Recover("deferred", &err) + panic("kaput") + } + var p *PanicError + if err := f(); !errors.As(err, &p) || p.Value != "kaput" { + t.Fatalf("err = %v", err) + } +} diff --git a/internal/schema/schema.go b/internal/schema/schema.go index a4ee86f..f29ab3e 100644 --- a/internal/schema/schema.go +++ b/internal/schema/schema.go @@ -18,6 +18,11 @@ type Table struct { // Unique lists multi-column UNIQUE constraints (single-column ones are // flagged on the column). Unique [][]string `yaml:"unique,omitempty"` + // PartitionedBy describes a partitioned table's key, e.g. "range (created_at)". + PartitionedBy string `yaml:"partitioned_by,omitempty"` + // Unseedable says why rows cannot be generated for the table unless a + // value rule sets them (e.g. "partitioned by an expression"). + Unseedable string `yaml:"unseedable,omitempty"` } // Column holds metadata and faker mapping for a single column. @@ -30,6 +35,9 @@ type Column struct { Nullable bool `yaml:"nullable,omitempty"` Unique bool `yaml:"unique,omitempty"` Generated bool `yaml:"generated,omitempty"` + // PartitionKey marks the key column of a partitioned table: its faker + // keeps values inside the partitions, also when the column is in the PK. + PartitionKey bool `yaml:"partition_key,omitempty"` } // Load reads a schema YAML file from disk. diff --git a/internal/seeder/gaps.go b/internal/seeder/gaps.go new file mode 100644 index 0000000..6d99c1d --- /dev/null +++ b/internal/seeder/gaps.go @@ -0,0 +1,30 @@ +package seeder + +// KnownEmpty reports whether counts says a table has no rows. A table missing +// from counts (its count failed) is unknown, never empty. +func KnownEmpty(counts map[string]int64, table string) bool { + n, ok := counts[table] + return ok && n == 0 +} + +// GapTables returns the tables of order known to be empty, in order. With only +// set, just those tables are considered. +func GapTables(order []string, counts map[string]int64, only []string) []string { + var allowed map[string]bool + if len(only) > 0 { + allowed = make(map[string]bool, len(only)) + for _, t := range only { + allowed[t] = true + } + } + var gaps []string + for _, t := range order { + if allowed != nil && !allowed[t] { + continue + } + if KnownEmpty(counts, t) { + gaps = append(gaps, t) + } + } + return gaps +} diff --git a/internal/seeder/gaps_test.go b/internal/seeder/gaps_test.go new file mode 100644 index 0000000..0cf05d0 --- /dev/null +++ b/internal/seeder/gaps_test.go @@ -0,0 +1,21 @@ +package seeder + +import ( + "reflect" + "testing" +) + +// Gaps are tables known to be empty. A table whose count failed is unknown and +// must never be treated as empty: seeding it could double a populated table. +func TestGapTables_OnlyTablesKnownToBeEmpty(t *testing.T) { + order := []string{"users", "orders", "audit", "tags"} + counts := map[string]int64{"users": 0, "orders": 12, "tags": 0} // audit's count failed + + if got := GapTables(order, counts, nil); !reflect.DeepEqual(got, []string{"users", "tags"}) { + t.Fatalf("GapTables = %v", got) + } + // Scoped to a selection: only selected tables that are known to be empty. + if got := GapTables(order, counts, []string{"audit", "tags", "orders"}); !reflect.DeepEqual(got, []string{"tags"}) { + t.Fatalf("scoped GapTables = %v", got) + } +} diff --git a/internal/seeder/mirror.go b/internal/seeder/mirror.go index 2e7814c..5478052 100644 --- a/internal/seeder/mirror.go +++ b/internal/seeder/mirror.go @@ -6,11 +6,14 @@ import ( "errors" "fmt" "sort" + "sync" "github.com/AxeForging/seedstorm/internal/compare" "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/faker" "github.com/AxeForging/seedstorm/internal/rules" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" "github.com/AxeForging/seedstorm/internal/schema" ) @@ -87,13 +90,25 @@ func Snapshots(ctx context.Context, source, target Endpoint, mode compare.CountM } return func(done, total int, table string) { onCount(side, done, total, table) } } - src, err := source.take(ctx, mode, progress("source")) - if err != nil { - return compare.Report{}, fmt.Errorf("source: %w", err) - } - tgt, err := target.take(ctx, mode, progress("target")) - if err != nil { - return compare.Report{}, fmt.Errorf("target: %w", err) + // The two databases are independent: read them at the same time. + var src, tgt compare.Snapshot + var srcErr, tgtErr error + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + srcErr = safego.Run("count source", func() (err error) { src, err = source.take(ctx, mode, progress("source")); return err }) + }() + go func() { + defer wg.Done() + tgtErr = safego.Run("count target", func() (err error) { tgt, err = target.take(ctx, mode, progress("target")); return err }) + }() + wg.Wait() + if srcErr != nil { + return compare.Report{}, runerr.OnSide(runerr.SideSource, runerr.At(runerr.PhaseCount, "", srcErr)) + } + if tgtErr != nil { + return compare.Report{}, runerr.OnSide(runerr.SideTarget, runerr.At(runerr.PhaseCount, "", tgtErr)) } return compare.Diff(src, tgt), nil } @@ -174,12 +189,17 @@ func (j *MirrorJob) Run(ctx context.Context, opts Options, onTruncate func(done, gen := j.generateOptions(opts.Generate.SelfRefDepth) gen.OnWarning = opts.Generate.OnWarning opts.Generate = gen + // Refuse tables that cannot be generated before anything is truncated. + if err := faker.CheckSeedable(j.Schema, j.Plan.Order, j.Overrides); err != nil { + return Result{}, err + } if j.Plan.Mode == compare.ModeReset && len(j.Plan.Truncate) > 0 { if err := db.TruncateConcurrently(ctx, j.target.Conn, j.target.DBType, j.Plan.Truncate, max(opts.Workers, 1), onTruncate); err != nil { - return Result{}, fmt.Errorf("truncate target: %w", err) + return Result{}, runerr.OnSide(runerr.SideTarget, runerr.At(runerr.PhaseTruncate, "", err)) } } - return Fill(ctx, j.target.Conn, j.target.DBType, j.Schema, j.Plan.Order, j.Plan.Counts(), opts) + res, err := Fill(ctx, j.target.Conn, j.target.DBType, j.Schema, j.Plan.Order, j.Plan.Counts(), opts) + return res, runerr.OnSide(runerr.SideTarget, err) } func (j *MirrorJob) generateOptions(selfRefDepth int) faker.GenerateOptions { diff --git a/internal/seeder/mirror_snapshot_test.go b/internal/seeder/mirror_snapshot_test.go index 217ce1e..1e15eeb 100644 --- a/internal/seeder/mirror_snapshot_test.go +++ b/internal/seeder/mirror_snapshot_test.go @@ -96,7 +96,7 @@ func TestPrepareMirror_LiveEndpointsStillRunTheSameDatabaseCheck(t *testing.T) { func TestSnapshots_SnapshotSideIsNotReadAndLiveSideNeedsAConnection(t *testing.T) { source := Endpoint{Snapshot: parsedSnapshot(t, "tables: {users: 1}")} _, err := Snapshots(context.Background(), source, Endpoint{DBType: "pgx"}, compare.CountExact, nil) - if err == nil || !strings.HasPrefix(err.Error(), "target: no database connection or snapshot") { + if err == nil || !strings.HasPrefix(err.Error(), "target · count: no database connection or snapshot") { t.Fatalf("err = %v", err) } diff --git a/internal/seeder/panic_test.go b/internal/seeder/panic_test.go new file mode 100644 index 0000000..38d3201 --- /dev/null +++ b/internal/seeder/panic_test.go @@ -0,0 +1,79 @@ +package seeder + +import ( + "errors" + "strings" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" +) + +// A panic while writing (here the database driver panics, the boundary the +// writer calls) used to take the whole process down: every job of `serve` with +// it. It must end this run with an error that says where. +func TestSeed_PanicInAWriterFailsTheRunInsteadOfTheProcess(t *testing.T) { + conn, rec := openRecording(t, nil, func(table string, n int) error { + if table == "users" && n == 2 { + panic("driver exploded") + } + return nil + }) + order := []string{"users", "reviewers", "audit", "posts"} + _, err := Seed(withDeadline(t, 20*time.Second), conn, "mysql", usersPostsAudit(), order, order, SeedOptions{ + Rows: 3000, BatchSize: 500, ChunkRows: 1000, Workers: 4, + }) + var p *safego.PanicError + if !errors.As(err, &p) || !strings.Contains(err.Error(), "driver exploded") { + t.Fatalf("err = %v, want a recovered panic", err) + } + e, ok := runerr.As(err) + if !ok || e.Phase != runerr.PhaseWrite || e.Table != "users" { + t.Fatalf("location = %+v (%v), want write · users", e, ok) + } + if got := rec.inserts("posts"); len(got) != 0 { + t.Fatalf("posts wrote %d batches after its parent failed", len(got)) + } +} + +// A panic on a parallel generator (here a value rule panics) fails the run the +// same way. +func TestSeed_PanicInAParallelGeneratorFailsTheRun(t *testing.T) { + conn, _ := openRecording(t, nil, nil) + order := []string{"users", "reviewers", "audit", "posts"} + overrides := faker.Overrides{"audit": {"note": func(row int, _ interface{}) (interface{}, error) { + if row == 7 { + var boom []int + _ = boom[3] + } + return "ok", nil + }}} + opts := SeedOptions{Rows: 200, BatchSize: 50, ChunkRows: 100, Workers: 4, GenWorkers: 3} + opts.Generate = faker.DefaultGenerateOptions() + opts.Generate.Overrides = overrides + _, err := Seed(withDeadline(t, 20*time.Second), conn, "mysql", usersPostsAudit(), order, order, opts) + var p *safego.PanicError + if !errors.As(err, &p) || !strings.Contains(err.Error(), "index out of range") { + t.Fatalf("err = %v, want a recovered panic", err) + } + if e, ok := runerr.As(err); !ok || e.Phase != runerr.PhaseGenerate || e.Table != "audit" { + t.Fatalf("location = %+v (%v), want generate · audit", e, ok) + } +} + +// Fill (mirror) splits chunks across workers too. +func TestFill_PanicInAPieceFailsTheTable(t *testing.T) { + conn, _ := openRecording(t, nil, func(table string, n int) error { + if table == "users" { + panic("fill driver exploded") + } + return nil + }) + res, err := Fill(withDeadline(t, 20*time.Second), conn, "mysql", usersPostsAudit(), []string{"users"}, map[string]int{"users": 2000}, Options{BatchSize: 200, Workers: 4, StopOnError: true}) + var p *safego.PanicError + if !errors.As(err, &p) { + t.Fatalf("err = %v (result %+v), want a recovered panic", err, res) + } +} diff --git a/internal/seeder/seed.go b/internal/seeder/seed.go index c1f597b..736f6ac 100644 --- a/internal/seeder/seed.go +++ b/internal/seeder/seed.go @@ -8,8 +8,13 @@ import ( "strings" "sync" + "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/faultinject" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" "github.com/AxeForging/seedstorm/internal/schema" + "github.com/AxeForging/seedstorm/internal/tuning" ) // SeedOptions describes a plain seed run: seed and gaps --fill in the CLI, TUI @@ -31,6 +36,9 @@ type SeedOptions struct { // Workers > 1, and is ignored for dry runs, when OnRows is set (callers // print chunks in order) and when Reproducible is set. See generateTables. GenWorkers int + // OnNotice receives decisions the run made that the user should know, such + // as fewer writers because the server has few free connections. + OnNotice func(msg string) // Reproducible keeps generation on one goroutine so a seeded run repeats // exactly: concurrent tables would draw from the random source in any order. Reproducible bool @@ -64,7 +72,7 @@ func Seed(ctx context.Context, conn *sql.DB, dbType string, sc *schema.Schema, p opts.BatchSize = DefaultBatchSize } opts.Generate.ChunkBytes = chunkBytesOrDefault(opts.Generate.ChunkBytes) - stream, err := faker.NewStream(sc, preload, tables, conn, dbType, opts.Generate.Overrides) + stream, err := faker.NewStreamContext(ctx, sc, preload, tables, conn, dbType, opts.Generate.Overrides) if err != nil { return SeedResult{Counts: map[string]int{}}, fmt.Errorf("data generation failed: %w", err) } @@ -84,25 +92,36 @@ func seedSequential(ctx context.Context, conn *sql.DB, dbType string, stream *fa } } count, overridden := faker.TableRowCount(tableName, opts.Rows, opts.TableRows) - err := stream.GenerateChunks(tableName, count, opts.EnumRows, overridden, chunkSize(opts.ChunkRows), opts.Generate, func(rows []map[string]interface{}) error { - if err := ctx.Err(); err != nil { + err := safego.Run("generate "+tableName, func() error { + if err := faultinject.Hit(ctx, "generate", tableName); err != nil { return err } - if opts.OnRows != nil { - if err := opts.OnRows(tableName, rows); err != nil { + return stream.GenerateChunks(tableName, count, opts.EnumRows, overridden, chunkSize(opts.ChunkRows), opts.Generate, func(rows []map[string]interface{}) error { + if err := ctx.Err(); err != nil { return err } - } - if !opts.DryRun { - if err := insertStrict(ctx, conn, dbType, tableName, rows, opts.BatchSize); err != nil { - return err + if opts.OnRows != nil { + if err := opts.OnRows(tableName, rows); err != nil { + return err + } } - } - tally.written(tableName, len(rows)) - return nil + if !opts.DryRun { + err := safego.Run("write "+tableName, func() error { + if err := faultinject.Hit(ctx, "write", tableName); err != nil { + return err + } + return insertStrict(ctx, conn, dbType, tableName, rows, opts.BatchSize) + }) + if err != nil { + return runerr.At(runerr.PhaseWrite, tableName, err) + } + } + tally.written(tableName, len(rows)) + return nil + }) }) if err != nil { - return tally.result(), err + return tally.result(), runerr.At(runerr.PhaseGenerate, tableName, err) } releases.generated(stream, tableName) tally.done(tableName) @@ -114,7 +133,9 @@ func seedConcurrent(ctx context.Context, conn *sql.DB, dbType string, sc *schema chunk := chunkSize(opts.ChunkRows) generators := 1 if opts.GenWorkers > 1 && opts.OnRows == nil && !opts.Reproducible { - generators = min(opts.GenWorkers, len(tables)) + // Never more generators than the CPU quota allows (a container's --cpus, + // not the host's cores). + generators = min(tuning.ClampGenerators(opts.GenWorkers), len(tables)) // Every generator holds a chunk: they share one chunk's worth of memory. opts.Generate.ChunkBytes = max(opts.Generate.ChunkBytes/generators, minGeneratorChunkBytes) opts.Generate.OnWarning = serialized(opts.Generate.OnWarning) @@ -126,6 +147,10 @@ func seedConcurrent(ctx context.Context, conn *sql.DB, dbType string, sc *schema if queue <= 0 { queue = chunk * approxRowMemory } + opts.Workers = clampToServer(ctx, conn, dbType, opts.Workers, generators, opts.OnNotice) + // The run never holds more connections than it uses: writers, generators + // reading parent keys, and one for sequences. + conn.SetMaxOpenConns(opts.Workers + generators + 1) w := newWriter(ctx, conn, dbType, opts.BatchSize, opts.Workers, queue) // A queued row costs at least half an average row of a full chunk, so // narrow rows queue at most about two chunks of rows: 300k narrow rows with @@ -159,6 +184,9 @@ func seedConcurrent(ctx context.Context, conn *sql.DB, dbType string, sc *schema generate := func(gen *faker.Stream, tableName string) error { tw := writers[tableName] defer tw.close() + if err := faultinject.Hit(w.ctx, "generate", tableName); err != nil { + return err + } if opts.OnTableStart != nil { if err := opts.OnTableStart(tableName); err != nil { return err @@ -199,8 +227,8 @@ const minGeneratorChunkBytes = 4 << 20 func generateTables(ctx context.Context, stream *faker.Stream, sc *schema.Schema, tables []string, generators int, generate func(*faker.Stream, string) error, releases *poolReleases) error { if generators <= 1 { for _, tableName := range tables { - if err := generate(stream, tableName); err != nil { - return err + if err := safego.Run("generate "+tableName, func() error { return generate(stream, tableName) }); err != nil { + return runerr.At(runerr.PhaseGenerate, tableName, err) } releases.generated(stream, tableName) } @@ -240,8 +268,8 @@ func generateTables(ctx context.Context, stream *faker.Stream, sc *schema.Schema // Its own random source: forks sharing the global one spent their // time waiting on its lock (4 generators ran slower than 1). fork.UseOwnRandom() - if err := generate(fork, tableName); err != nil { - cancel(err) + if err := safego.Run("generate "+tableName, func() error { return generate(fork, tableName) }); err != nil { + cancel(runerr.At(runerr.PhaseGenerate, tableName, err)) return } stream.MergeTable(fork, tableName) @@ -415,3 +443,21 @@ func (t *seedTally) result() SeedResult { } return res } + +// connectionUsage reads the server's connection limit and use; a variable so +// tests can stand in for a busy server. +var connectionUsage = db.ConnectionUsage + +// clampToServer lowers writers when the server lacks free connections for the +// run, and says so. Any error reading the limit leaves writers unchanged. +func clampToServer(ctx context.Context, conn *sql.DB, dbType string, writers, generators int, notice func(string)) int { + maxConns, used, err := connectionUsage(ctx, conn, dbType) + if err != nil { + return writers + } + clamped := tuning.ClampWriters(writers, generators, maxConns, used) + if clamped < writers && notice != nil { + notice(fmt.Sprintf("Using %d writers instead of %d: the server has %d of %d connections in use", clamped, writers, used, maxConns)) + } + return clamped +} diff --git a/internal/seeder/seeder.go b/internal/seeder/seeder.go index 94b59d4..a79310b 100644 --- a/internal/seeder/seeder.go +++ b/internal/seeder/seeder.go @@ -17,6 +17,8 @@ import ( "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" "github.com/AxeForging/seedstorm/internal/schema" ) @@ -60,6 +62,9 @@ type Options struct { Generate faker.GenerateOptions // OnProgress reports inserted/requested rows for the current table. OnProgress func(p Progress) + // OnNotice receives decisions the run made, such as fewer writers because + // the server has few free connections. + OnNotice func(msg string) } // Progress is one progress tick: the table that just advanced, and the run. @@ -120,6 +125,9 @@ func Fill(ctx context.Context, conn *sql.DB, dbType string, sc *schema.Schema, o if opts.MaxRowFailures <= 0 { opts.MaxRowFailures = DefaultMaxRowFailures } + if opts.Workers > 1 { + opts.Workers = clampToServer(ctx, conn, dbType, opts.Workers, 1, opts.OnNotice) + } // Explicit ids leave Postgres sequences behind; move them forward for every // table that received rows, including when the run stops early. defer func() { @@ -178,8 +186,12 @@ func fillTable(ctx context.Context, conn *sql.DB, dbType string, sc *schema.Sche _, selfRef := referencedTables(sc, tableName, nil) // Parents are complete by now (tables run in FK order), so their pools and // this table's stored keys are read once and kept for every chunk. - stream, err := faker.NewStream(sc, preload, []string{tableName}, conn, dbType, opts.Generate.Overrides) + stream, err := faker.NewStreamContext(ctx, sc, preload, []string{tableName}, conn, dbType, opts.Generate.Overrides) if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + tr.Missing, tr.Status, tr.Error = tr.Requested, StatusFailed, ctxErr.Error() + return tr, ctxErr + } return finishStuck(tr, "read existing rows: "+err.Error(), opts) } genOpts := opts.Generate @@ -329,7 +341,14 @@ func insertRowsConcurrently(ctx context.Context, conn *sql.DB, dbType, tableName go func() { defer wg.Done() o := &results[i] - o.inserted, o.rejected, o.lastErr, o.err = insertRows(ctx, conn, dbType, tableName, piece, opts) + o.err = safego.Run("write "+tableName, func() error { + var err error + o.inserted, o.rejected, o.lastErr, err = insertRows(ctx, conn, dbType, tableName, piece, opts) + return err + }) + if o.err != nil { + o.err = runerr.At(runerr.PhaseWrite, tableName, o.err) + } }() } wg.Wait() diff --git a/internal/seeder/writer.go b/internal/seeder/writer.go index 4d0574a..0260d62 100644 --- a/internal/seeder/writer.go +++ b/internal/seeder/writer.go @@ -10,6 +10,9 @@ import ( "time" "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/faultinject" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" "github.com/AxeForging/seedstorm/internal/schema" ) @@ -73,7 +76,11 @@ func newWriter(ctx context.Context, conn *sql.DB, dbType string, batch, workers, go func() { defer w.running.Done() for task := range w.work { - task() + // A panic in a task fails the run, not the process; the + // task's own defers still release its bookkeeping. + if err := safego.Run("writer", func() error { task(); return nil }); err != nil { + w.abort(err) + } } }() } @@ -105,7 +112,11 @@ type tableWriter struct { func (w *writer) open(name string, parents []*tableWriter, ordered bool) *tableWriter { tw := &tableWriter{w: w, name: name, parents: parents, ordered: ordered, wake: make(chan struct{}, 1), done: make(chan struct{})} w.tables.Add(1) - go tw.dispatch() + go func() { + if err := safego.Run("dispatch "+name, func() error { tw.dispatch(); return nil }); err != nil { + w.abort(runerr.At(runerr.PhaseWrite, name, err)) + } + }() return tw } @@ -248,9 +259,15 @@ func (tw *tableWriter) write(item queued) { if w.ctx.Err() != nil { return } - if err := insertStrict(w.ctx, w.conn, w.dbType, tw.name, rows, w.batch); err != nil { + err := safego.Run("write "+tw.name, func() error { + if err := faultinject.Hit(w.ctx, "write", tw.name); err != nil { + return err + } + return insertStrict(w.ctx, w.conn, w.dbType, tw.name, rows, w.batch) + }) + if err != nil { if w.ctx.Err() == nil { - w.cancel(err) + w.cancel(runerr.At(runerr.PhaseWrite, tw.name, err)) } return } diff --git a/internal/seeder/writer_test.go b/internal/seeder/writer_test.go index cb29f00..a9e4436 100644 --- a/internal/seeder/writer_test.go +++ b/internal/seeder/writer_test.go @@ -41,6 +41,8 @@ type recording struct { active map[string]int // overlapSelf records a table that ever had two inserts in flight. overlapSelf map[string]bool + // inFlight and peak count inserts running at once across tables. + inFlight, peak int // discard keeps no events (benchmarks). discard bool } @@ -104,6 +106,8 @@ func (c *recordingConn) ExecContext(ctx context.Context, query string, args []dr if r.active[table] > 1 { r.overlapSelf[table] = true } + r.inFlight++ + r.peak = max(r.peak, r.inFlight) ev := insertEvent{table: table, start: time.Now()} if !r.discard { ev.rows = decodeInsert(query, args) @@ -124,6 +128,7 @@ func (c *recordingConn) ExecContext(ctx context.Context, query string, args []dr } r.mu.Lock() r.active[table]-- + r.inFlight-- ev.end = time.Now() if err == nil && !r.discard { r.events = append(r.events, ev) @@ -166,6 +171,12 @@ func (r *recording) rowsOf(table string) []map[string]interface{} { return out } +func (r *recording) peakConcurrency() int { + r.mu.Lock() + defer r.mu.Unlock() + return r.peak +} + func (r *recording) inserts(table string) []insertEvent { r.mu.Lock() defer r.mu.Unlock() @@ -423,3 +434,26 @@ func TestFill_WorkersSplitAChunkButKeepSelfReferencesWhole(t *testing.T) { } }) } + +// A server with few free connections gets fewer writers, and the run says so +// instead of failing with "too many connections" part-way through. +func TestSeed_WritersClampedToTheServersFreeConnections(t *testing.T) { + defer func(old func(context.Context, *sql.DB, string) (int, int, error)) { connectionUsage = old }(connectionUsage) + connectionUsage = func(context.Context, *sql.DB, string) (int, int, error) { return 20, 16, nil } + conn, rec := openRecording(t, map[string]time.Duration{"users": 3 * time.Millisecond, "audit": 3 * time.Millisecond}, nil) + var notices []string + order := []string{"users", "reviewers", "audit", "posts"} + _, err := Seed(withDeadline(t, 20*time.Second), conn, "mysql", usersPostsAudit(), order, order, SeedOptions{ + Rows: 4000, BatchSize: 250, ChunkRows: 1000, Workers: 8, + OnNotice: func(msg string) { notices = append(notices, msg) }, + }) + if err != nil { + t.Fatal(err) + } + if len(notices) != 1 || !strings.Contains(notices[0], "Using 2 writers instead of 8") { + t.Fatalf("notices = %q", notices) + } + if peak := rec.peakConcurrency(); peak > 2 { + t.Fatalf("%d inserts ran at once with 2 writers allowed", peak) + } +} diff --git a/internal/tui/clone.go b/internal/tui/clone.go index 06a79bb..e725c27 100644 --- a/internal/tui/clone.go +++ b/internal/tui/clone.go @@ -4,10 +4,13 @@ import ( "context" "fmt" "strings" + "time" + "github.com/charmbracelet/bubbles/spinner" tea "github.com/charmbracelet/bubbletea" "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/safego" ) type cloneModel struct { @@ -18,6 +21,8 @@ type cloneModel struct { targetDSN string opts db.CloneOptions running bool + spinner spinner.Model + startedAt time.Time done bool confirmed bool result db.CloneResult @@ -31,7 +36,11 @@ type cloneDoneMsg struct { // RunClone presents a small confirmation UI around schema cloning. func RunClone(ctx context.Context, sourceType, sourceDSN, targetType, targetDSN string, opts db.CloneOptions) error { + sp := spinner.New() + sp.Spinner = spinner.Dot + sp.Style = selectedStyle m := cloneModel{ + spinner: sp, ctx: ctx, sourceType: sourceType, sourceDSN: sourceDSN, @@ -50,6 +59,15 @@ func RunClone(ctx context.Context, sourceType, sourceDSN, targetType, targetDSN if !fm.confirmed { return fmt.Errorf("aborted by user") } + // Printed after the screen is released, so it stays in the terminal. + if fm.opts.DryRun { + fmt.Println(strings.Join(fm.result.Statements, ";\n") + ";") + } else { + fmt.Printf("Cloned %d tables (%d statements).\n", fm.result.Tables, len(fm.result.Statements)) + } + for _, skipped := range fm.result.Skipped { + fmt.Printf("Not cloned: %s %s (%s)\n", skipped.Kind, skipped.Name, skipped.Reason) + } return nil } @@ -67,8 +85,16 @@ func (m cloneModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } m.confirmed = true m.running = true - return m, m.run() + m.startedAt = time.Now() + return m, tea.Batch(m.spinner.Tick, m.run()) } + case spinner.TickMsg: + if !m.running { + return m, nil + } + var cmd tea.Cmd + m.spinner, cmd = m.spinner.Update(msg) + return m, cmd case cloneDoneMsg: m.running = false m.done = true @@ -95,7 +121,8 @@ func (m cloneModel) View() string { } sb.WriteString("\n") if m.running { - sb.WriteString(" Cloning schema...\n") + fmt.Fprintf(&sb, " %s Cloning schema: reading the source, then running DDL on the target · %s\n", m.spinner.View(), time.Since(m.startedAt).Round(100*time.Millisecond)) + sb.WriteString(helpStyle.Render(" ctrl+c stops (statements already run stay)")) return sb.String() } if m.done { @@ -112,10 +139,11 @@ func (m cloneModel) View() string { func (m cloneModel) run() tea.Cmd { return func() tea.Msg { - result, err := db.CloneSchema(m.ctx, m.sourceType, m.sourceDSN, m.targetType, m.targetDSN, m.opts) - if m.opts.DryRun && err == nil { - fmt.Println(strings.Join(result.Statements, ";\n") + ";") - } + var result db.CloneResult + err := safego.Run("clone schema", func() (err error) { + result, err = db.CloneSchema(m.ctx, m.sourceType, m.sourceDSN, m.targetType, m.targetDSN, m.opts) + return err + }) return cloneDoneMsg{result: result, err: err} } } diff --git a/internal/tui/execute.go b/internal/tui/execute.go index f97a5db..d033827 100644 --- a/internal/tui/execute.go +++ b/internal/tui/execute.go @@ -13,6 +13,8 @@ import ( "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" "github.com/AxeForging/seedstorm/internal/seeder" ) @@ -22,6 +24,12 @@ type tableSeededMsg struct { rows int } +// seedProgressMsg reports a written piece with run totals (seeder.Progress). +type seedProgressMsg seeder.Progress + +// seedPhaseMsg names a step before rows are written (connecting, truncating). +type seedPhaseMsg string + // seedDoneMsg is sent when the entire seed operation completes. type seedDoneMsg struct { totalRows int @@ -29,6 +37,7 @@ type seedDoneMsg struct { tables []string rowsMap map[string]int err error + syncErr error } // dryRunDoneMsg is sent when dry-run generation completes. @@ -63,6 +72,14 @@ type executeModel struct { seededTables []string seededRows map[string]int height int + // events carries progress from the running seed; nil before it starts. + events chan tea.Msg + startedAt time.Time + phase string + progress seeder.Progress + meter *seeder.Meter + estimate seeder.Estimate + syncErr error } func newExecute(totalTables int, dryRun bool) executeModel { @@ -74,6 +91,7 @@ func newExecute(totalTables int, dryRun bool) executeModel { totalTables: totalTables, dryRun: dryRun, seededRows: make(map[string]int), + startedAt: time.Now(), } } @@ -106,14 +124,32 @@ func (m executeModel) Update(msg tea.Msg) (executeModel, tea.Cmd) { m.completedTables++ m.seededTables = append(m.seededTables, msg.table) m.seededRows[msg.table] = msg.rows + return m, waitSeed(m.events) + case seedPhaseMsg: + m.phase = string(msg) + return m, waitSeed(m.events) + case seedProgressMsg: + p := seeder.Progress(msg) + m.phase = "" + m.progress = p + m.currentTable = p.Table + if m.meter == nil { + m.meter = seeder.NewMeter(m.startedAt) + } + m.estimate = m.meter.Observe(time.Now(), p.RowsDone, p.RowsTotal) + return m, waitSeed(m.events) case seedDoneMsg: m.done = true m.totalRows = msg.totalRows m.elapsed = msg.elapsed m.seededTables = msg.tables m.seededRows = msg.rowsMap + if m.seededRows == nil { + m.seededRows = map[string]int{} + } m.completedTables = len(msg.tables) m.err = msg.err + m.syncErr = msg.syncErr return m, nil case dryRunDoneMsg: m.done = true @@ -235,67 +271,137 @@ func (m executeModel) View() string { if m.done { if m.err != nil { sb.WriteString(errorStyle.Render(fmt.Sprintf(" Error: %v\n", m.err))) + sb.WriteString("\n") + var notWritten []string + for _, t := range m.seededTables { + if n := m.seededRows[t]; n > 0 { + fmt.Fprintf(&sb, " %-30s %d rows\n", t, n) + } else { + notWritten = append(notWritten, t) + } + } + if len(notWritten) > 0 { + sb.WriteString(dimStyle.Render(" not written: " + strings.Join(notWritten, ", "))) + sb.WriteString("\n") + } } else { sb.WriteString(successStyle.Render(fmt.Sprintf(" Seeding complete! %d rows across %d tables in %s\n", m.totalRows, m.completedTables, m.elapsed.Round(time.Millisecond)))) + sb.WriteString("\n") + for _, t := range m.seededTables { + fmt.Fprintf(&sb, " %-30s %d rows\n", t, m.seededRows[t]) + } } - sb.WriteString("\n") - for _, t := range m.seededTables { - fmt.Fprintf(&sb, " %-30s %d rows\n", t, m.seededRows[t]) + if m.syncErr != nil { + sb.WriteString(errorStyle.Render(fmt.Sprintf("\n Sequences not advanced: %v (application inserts may reuse seeded ids)\n", m.syncErr))) } sb.WriteString("\n") sb.WriteString(helpStyle.Render(" q quit")) } else { - pct := 0 - if m.totalTables > 0 { - pct = m.completedTables * 100 / m.totalTables + elapsed := time.Since(m.startedAt).Round(100 * time.Millisecond) + switch { + case m.phase != "": + fmt.Fprintf(&sb, " %s %s · %s\n", m.spinner.View(), m.phase, elapsed) + case m.progress.RowsTotal > 0: + pct := m.progress.RowsDone * 100 / m.progress.RowsTotal + fmt.Fprintf(&sb, " %s Seeding %s %s / %s rows (%d%%) · %d/%d tables · %s\n", + m.spinner.View(), m.currentTable, groupThousands(m.progress.RowsDone), groupThousands(m.progress.RowsTotal), + pct, m.completedTables, m.totalTables, elapsed) + if est := m.estimate.String(); est != "" { + sb.WriteString(dimStyle.Render(" " + est)) + sb.WriteString("\n") + } + case m.currentTable != "": + fmt.Fprintf(&sb, " %s Seeding %s · %d/%d tables · %s\n", m.spinner.View(), m.currentTable, m.completedTables, m.totalTables, elapsed) + default: + fmt.Fprintf(&sb, " %s Starting (%d tables) · %s\n", m.spinner.View(), m.totalTables, elapsed) } - fmt.Fprintf(&sb, " %s Seeding %s (%d/%d tables, %d%%)\n", - m.spinner.View(), m.currentTable, m.completedTables, m.totalTables, pct) + sb.WriteString(helpStyle.Render("\n ctrl+c aborts (rows already inserted stay)")) } return sb.String() } -// startSeed returns a tea.Cmd that runs the seed operation in a goroutine. -func startSeed(ctx context.Context, s *seedParams) tea.Cmd { +// startSeed runs the seed in a goroutine; progress, finished tables and the +// final result arrive on events (read with waitSeed). +func startSeed(ctx context.Context, s *seedParams, events chan tea.Msg) tea.Cmd { + return runSeedInto(ctx, s, s.tables, s.truncate, events) +} + +// runSeedInto seeds tables (preloading keys of preload) and reports on events. +func runSeedInto(ctx context.Context, s *seedParams, preload []string, truncate bool, events chan tea.Msg) tea.Cmd { return func() tea.Msg { start := time.Now() - - batchSize := s.batchSize - if batchSize < 1 { - batchSize = 1 - } - - conn, err := sql.Open(s.dbType, s.dsn) - if err != nil { - return seedDoneMsg{err: fmt.Errorf("failed to open connection: %w", err)} - } - defer conn.Close() - - if err := conn.PingContext(ctx); err != nil { - return seedDoneMsg{err: fmt.Errorf("failed to ping database: %w", err)} - } - - if s.truncate { - if err := db.TruncateConcurrently(ctx, conn, s.dbType, s.tables, seeder.DefaultWorkers, nil); err != nil { - return seedDoneMsg{err: fmt.Errorf("truncate failed: %w", err)} + send := func(msg tea.Msg) { + select { + case events <- msg: + default: // the view keeps up; dropping a tick only skips a redraw } } + done := seedDoneMsg{tables: s.tables} + err := safego.Run("seed", func() error { + batchSize := max(s.batchSize, 1) + send(seedPhaseMsg("Connecting")) + conn, err := sql.Open(s.dbType, s.dsn) + if err != nil { + return runerr.At(runerr.PhaseConnect, "", fmt.Errorf("failed to open connection: %w", err)) + } + defer conn.Close() + pctx, cancel := context.WithTimeout(ctx, 10*time.Second) + err = conn.PingContext(pctx) + cancel() + if err != nil { + return runerr.At(runerr.PhaseConnect, "", fmt.Errorf("database did not answer: %w", err)) + } + if truncate { + send(seedPhaseMsg("Truncating tables")) + if err := db.TruncateConcurrently(ctx, conn, s.dbType, s.tables, seeder.DefaultWorkers, nil); err != nil { + return runerr.At(runerr.PhaseTruncate, "", fmt.Errorf("truncate failed: %w", err)) + } + } + send(seedPhaseMsg("Generating the first rows")) + // Move Postgres sequences past the inserted ids, even if an insert fails. + defer func() { + if _, err := db.SyncSequences(context.WithoutCancel(ctx), conn, s.dbType, s.tables); err != nil { + done.syncErr = err + } + }() + opts := s.seedOptions(batchSize, false, nil) + opts.OnProgress = func(p seeder.Progress) { send(seedProgressMsg(p)) } + opts.OnTable = func(p seeder.Progress) { send(tableSeededMsg{table: p.Table, rows: int(p.Inserted)}) } + res, err := seeder.Seed(ctx, conn, s.dbType, s.schema, preload, s.tables, opts) + done.totalRows, done.rowsMap = res.Total, res.Counts + return err + }) + done.err = err + done.elapsed = time.Since(start) + events <- done + return nil + } +} - // Move Postgres sequences past the inserted ids, even if an insert fails. - defer func() { _, _ = db.SyncSequences(ctx, conn, s.dbType, s.tables) }() - res, err := seeder.Seed(ctx, conn, s.dbType, s.schema, s.tables, s.tables, s.seedOptions(batchSize, false, nil)) - if err != nil { - return seedDoneMsg{err: err} - } - return seedDoneMsg{ - totalRows: res.Total, - elapsed: time.Since(start), - tables: s.tables, - rowsMap: res.Counts, +// waitSeed reads the next message of a running seed. +func waitSeed(events chan tea.Msg) tea.Cmd { + if events == nil { + return nil + } + return func() tea.Msg { return <-events } +} + +// groupThousands formats n as 12,345. +func groupThousands(n int64) string { + s := fmt.Sprint(n) + if n < 0 { + return "-" + groupThousands(-n) + } + var out []byte + for i, c := range []byte(s) { + if i > 0 && (len(s)-i)%3 == 0 { + out = append(out, ',') } + out = append(out, c) } + return string(out) } // startDryRun returns a tea.Cmd that generates data and builds a summary. diff --git a/internal/tui/execute_test.go b/internal/tui/execute_test.go index 3b6fb84..898c433 100644 --- a/internal/tui/execute_test.go +++ b/internal/tui/execute_test.go @@ -2,9 +2,11 @@ package tui import ( "fmt" + "strings" "testing" "time" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/schema" tea "github.com/charmbracelet/bubbletea" ) @@ -486,3 +488,40 @@ func TestStartDryRun_usesPerTableRowOverrides(t *testing.T) { t.Fatalf("rows by table = %+v, want users=2 orders=5", got) } } + +// ── progress while seeding ────────────────────────────────────────────────── + +// The running view used to stay at "0/N tables, 0%" for the whole run: no +// progress ever reached it. Progress messages move it, with rows and rate. +func TestExecute_progressMessagesMoveTheRunningView(t *testing.T) { + m := newExecute(3, false) + m.events = make(chan tea.Msg, 1) + m.startedAt = time.Now().Add(-10 * time.Second) + updated, cmd := m.Update(seedProgressMsg{Table: "orders", TableIndex: 2, Tables: 3, RowsDone: 12000, RowsTotal: 40000}) + view := updated.View() + for _, want := range []string{"orders", "12,000 / 40,000 rows", "30%"} { + if !strings.Contains(view, want) { + t.Errorf("view lacks %q:\n%s", want, view) + } + } + if cmd == nil { + t.Error("after a progress message the model must keep listening for the next one") + } + updated, _ = updated.Update(tableSeededMsg{table: "users", rows: 5000}) + if !strings.Contains(updated.View(), "1/3 tables") { + t.Errorf("finished tables not counted:\n%s", updated.View()) + } +} + +// A failed run keeps what it wrote on screen and says where it stopped. +func TestExecute_failureKeepsPartialResultsAndLocation(t *testing.T) { + m := newExecute(2, false) + cause := runerr.At(runerr.PhaseWrite, "orders", fmt.Errorf("insert into orders failed: duplicate key")) + updated, _ := m.Update(seedDoneMsg{err: cause, tables: []string{"users", "orders"}, rowsMap: map[string]int{"users": 50}, totalRows: 50}) + view := updated.View() + for _, want := range []string{"write · orders", "users", "50 rows", "not written", "orders"} { + if !strings.Contains(view, want) { + t.Errorf("failure view lacks %q:\n%s", want, view) + } + } +} diff --git a/internal/tui/gaps.go b/internal/tui/gaps.go index d10e65f..82874a6 100644 --- a/internal/tui/gaps.go +++ b/internal/tui/gaps.go @@ -2,17 +2,13 @@ package tui import ( "context" - "database/sql" "fmt" "strings" - "time" tea "github.com/charmbracelet/bubbletea" - "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/graph" "github.com/AxeForging/seedstorm/internal/schema" - "github.com/AxeForging/seedstorm/internal/seeder" ) // gapsStep tracks the wizard state for the gaps command. @@ -61,14 +57,14 @@ func RunGaps(ctx context.Context, s *schema.Schema, dbType, dsn string, counts m // Build items: empty tables are selectable, populated are shown but disabled var items []tableItem for _, name := range sortedAll { - count := counts[name] + count, known := counts[name] parents := g.Parents(name) item := tableItem{ name: name, parents: parents, } - if count == 0 { - item.selected = true // empty tables default to selected + if known && count == 0 { + item.selected = true // empty tables default to selected; an unknown count is never assumed empty } // Populated tables are not shown in picker (only empty ones matter for gaps) items = append(items, item) @@ -110,8 +106,12 @@ func RunGaps(ctx context.Context, s *schema.Schema, dbType, dsn string, counts m func newGapsPicker(items []tableItem, counts map[string]int64, height int) tablePickerModel { // Annotate parents with row counts for i := range items { - count := counts[items[i].name] - if count > 0 { + count, known := counts[items[i].name] + if !known { + items[i].name = items[i].name + " (count unknown)" + items[i].selected = false + items[i].autoSelected = false + } else if count > 0 { items[i].name = fmt.Sprintf("%s (%d rows)", items[i].name, count) items[i].selected = false items[i].autoSelected = false @@ -275,7 +275,8 @@ func (m GapsModel) updateReview(msg tea.Msg) (tea.Model, tea.Cmd) { if m.review.dryRun { return m, tea.Batch(m.execute.spinner.Tick, startGapsDryRun(params, m.sortedAll)) } - return m, tea.Batch(m.execute.spinner.Tick, startGapsFill(m.ctx, params, m.sortedAll)) + m.execute.events = make(chan tea.Msg, 64) + return m, tea.Batch(m.execute.spinner.Tick, startGapsFill(m.ctx, params, m.sortedAll, m.execute.events), waitSeed(m.execute.events)) } return m, cmd } @@ -323,39 +324,10 @@ func (m GapsModel) View() string { return sb.String() } -// startGapsFill seeds only the selected gap tables, pre-loading PKs from all tables. -func startGapsFill(ctx context.Context, s *seedParams, allSorted []string) tea.Cmd { - return func() tea.Msg { - start := time.Now() - - batchSize := s.batchSize - if batchSize < 1 { - batchSize = 1 - } - - conn, err := sql.Open(s.dbType, s.dsn) - if err != nil { - return seedDoneMsg{err: fmt.Errorf("failed to open connection: %w", err)} - } - defer conn.Close() - - if err := conn.PingContext(ctx); err != nil { - return seedDoneMsg{err: fmt.Errorf("failed to ping database: %w", err)} - } - - // Move Postgres sequences past the inserted ids, even if an insert fails. - defer func() { _, _ = db.SyncSequences(ctx, conn, s.dbType, s.tables) }() - res, err := seeder.Seed(ctx, conn, s.dbType, s.schema, allSorted, s.tables, s.seedOptions(batchSize, false, nil)) - if err != nil { - return seedDoneMsg{err: err} - } - return seedDoneMsg{ - totalRows: res.Total, - elapsed: time.Since(start), - tables: s.tables, - rowsMap: res.Counts, - } - } +// startGapsFill seeds only the selected gap tables, pre-loading PKs from all +// tables; progress arrives on events. +func startGapsFill(ctx context.Context, s *seedParams, allSorted []string, events chan tea.Msg) tea.Cmd { + return runSeedInto(ctx, s, allSorted, false, events) } // startGapsDryRun generates data for gap tables and returns a preview. diff --git a/internal/tui/mirror.go b/internal/tui/mirror.go index b06c3a1..f375b04 100644 --- a/internal/tui/mirror.go +++ b/internal/tui/mirror.go @@ -10,6 +10,7 @@ import ( "github.com/goccy/go-yaml" "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/safego" "github.com/AxeForging/seedstorm/internal/seeder" ) @@ -32,8 +33,12 @@ type mirrorModel struct { phase mirrorPhase showPreview bool preview string - offset int - height int + // previewLoading: sample rows are being generated (reads the target) in + // the background, so the screen stays responsive. + previewLoading bool + truncating string + offset int + height int progress seeder.Progress events chan tea.Msg @@ -44,6 +49,12 @@ type mirrorModel struct { type mirrorProgressMsg seeder.Progress +// mirrorPreviewMsg carries sample rows generated in the background. +type mirrorPreviewMsg string + +// mirrorTruncateMsg reports the table being truncated in reset mode. +type mirrorTruncateMsg string + type mirrorDoneMsg struct { result seeder.Result err error @@ -61,8 +72,15 @@ func RunMirror(ctx context.Context, job *seeder.MirrorJob, opts seeder.Options, fm := final.(mirrorModel) switch { case fm.err != nil: + // Keep what was written visible after the screen closes. + if len(fm.result.Tables) > 0 { + seeder.RenderResult(os.Stdout, fm.result) + } return fm.err case fm.aborted: + if len(fm.result.Tables) > 0 { + seeder.RenderResult(os.Stdout, fm.result) + } return fmt.Errorf("aborted by user") case fm.phase == mirrorDone: seeder.RenderResult(os.Stdout, fm.result) @@ -84,8 +102,16 @@ func (m mirrorModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case tea.WindowSizeMsg: m.height = msg.Height case mirrorProgressMsg: + m.truncating = "" m.progress = seeder.Progress(msg) return m, waitMirror(m.events) + case mirrorTruncateMsg: + m.truncating = string(msg) + return m, waitMirror(m.events) + case mirrorPreviewMsg: + m.previewLoading = false + m.preview = string(msg) + return m, nil case mirrorDoneMsg: m.phase = mirrorDone m.result = msg.result @@ -124,10 +150,11 @@ func (m mirrorModel) handleKey(msg tea.KeyMsg) (tea.Model, tea.Cmd) { m.offset++ case "p": m.showPreview = !m.showPreview - if m.showPreview && m.preview == "" { - m.preview = m.renderPreview() - } m.offset = 0 + if m.showPreview && m.preview == "" && !m.previewLoading { + m.previewLoading = true + return m, m.loadPreview() + } case "y", "enter": if m.dryRun || m.job.Plan.TotalInsert == 0 { return m, nil @@ -153,12 +180,32 @@ func (m mirrorModel) start() tea.Cmd { default: } } - result, err := job.Run(ctx, opts, nil) + var result seeder.Result + err := safego.Run("mirror", func() (err error) { + result, err = job.Run(ctx, opts, func(_, _ int, table string) { + select { + case events <- mirrorTruncateMsg(table): + default: + } + }) + return err + }) events <- mirrorDoneMsg{result: result, err: err} return nil } } +// loadPreview generates sample rows off the UI goroutine. +func (m mirrorModel) loadPreview() tea.Cmd { + return func() tea.Msg { + var text string + if err := safego.Run("mirror preview", func() error { text = m.renderPreview(); return nil }); err != nil { + text = "Sample rows unavailable: " + err.Error() + } + return mirrorPreviewMsg(text) + } +} + func waitMirror(events chan tea.Msg) tea.Cmd { return func() tea.Msg { return <-events } } @@ -187,7 +234,9 @@ func (m mirrorModel) View() string { switch m.phase { case mirrorRunning: p := m.progress - if p.Tables == 0 { + if m.truncating != "" { + fmt.Fprintf(&sb, "\n Truncating %s…\n", m.truncating) + } else if p.Tables == 0 { fmt.Fprintf(&sb, "\n Preparing %d rows into %d tables…\n", m.job.Plan.TotalInsert, len(m.job.Plan.Entries)) } else { fmt.Fprintf(&sb, "\n Filling %s (%d/%d tables) %d / %d rows\n", p.Table, p.TableIndex, p.Tables, p.Inserted, p.Requested) @@ -199,7 +248,9 @@ func (m mirrorModel) View() string { } var body strings.Builder - if m.showPreview { + if m.showPreview && m.previewLoading { + body.WriteString("Generating sample rows from the target… (the plan stays usable: p to go back)") + } else if m.showPreview { body.WriteString(m.preview) } else { compare.RenderPlan(&body, m.job.Plan) diff --git a/internal/tui/mirror_test.go b/internal/tui/mirror_test.go index ca23993..979ec92 100644 --- a/internal/tui/mirror_test.go +++ b/internal/tui/mirror_test.go @@ -125,3 +125,25 @@ func TestMirrorModel_ProgressAndDoneMessages(t *testing.T) { t.Fatalf("done: phase = %v result = %+v", m.phase, m.result) } } + +// Sample rows read the target database. Generating them inside Update froze the +// whole screen; now the key returns at once with a loading state and the rows +// arrive as a message. +func TestMirrorModel_PreviewLoadsWithoutBlockingTheScreen(t *testing.T) { + m := newMirrorModel(context.Background(), mirrorJob(compare.ModeTopUp, 40), seeder.Options{}, 3, true) + m, cmd := press(m, "p") + if cmd == nil { + t.Fatal("p did not start loading the preview in the background") + } + if !m.previewLoading || !strings.Contains(m.View(), "Generating sample rows") { + t.Fatalf("no loading state:\n%s", m.View()) + } + next, _ := m.Update(mirrorPreviewMsg("Sample rows (up to 3 per table, nothing written)\n\nusers: []")) + m = next.(mirrorModel) + if m.previewLoading || !strings.Contains(m.View(), "nothing written") { + t.Fatalf("preview not shown after it arrived:\n%s", m.View()) + } + if m, cmd = press(m, "p"); cmd != nil || m.showPreview { + t.Fatal("toggling back to the plan should not reload anything") + } +} diff --git a/internal/tui/tui.go b/internal/tui/tui.go index e1b5507..64a6890 100644 --- a/internal/tui/tui.go +++ b/internal/tui/tui.go @@ -272,7 +272,8 @@ func (m Model) updateReview(msg tea.Msg) (tea.Model, tea.Cmd) { if m.review.dryRun { return m, tea.Batch(m.execute.spinner.Tick, startDryRun(params)) } - return m, tea.Batch(m.execute.spinner.Tick, startSeed(m.ctx, params)) + m.execute.events = make(chan tea.Msg, 64) + return m, tea.Batch(m.execute.spinner.Tick, startSeed(m.ctx, params, m.execute.events), waitSeed(m.execute.events)) } return m, cmd diff --git a/internal/tuning/host.go b/internal/tuning/host.go new file mode 100644 index 0000000..d20becb --- /dev/null +++ b/internal/tuning/host.go @@ -0,0 +1,122 @@ +// Package tuning decides how hard a run may push: how many generators the +// machine running seedstorm can feed, and (later) how many writers a database +// can take. +package tuning + +import ( + "math" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" +) + +// cgroupRoot is where the process's own cgroup is mounted in a container. +const cgroupRoot = "/sys/fs/cgroup" + +// MaxGenerators is how many tables may generate at once on this machine: its +// cores, or fewer when a container CPU quota limits the process. +func MaxGenerators() int { + quota, limited := cpuQuota(cgroupRoot) + return maxGenerators(runtime.NumCPU(), quota, limited) +} + +// ClampGenerators bounds a requested generator count to MaxGenerators. +func ClampGenerators(n int) int { + return max(1, min(n, MaxGenerators())) +} + +// maxGenerators rounds a fractional quota down: generation is CPU-bound, and a +// half core cannot run a second generator. +func maxGenerators(hostCPUs int, quota float64, limited bool) int { + n := max(1, hostCPUs) + if limited { + n = min(n, max(1, int(math.Floor(quota)))) + } + return n +} + +// cpuQuota reads the CPU limit of the cgroup mounted at root, in cores: cgroup +// v2 cpu.max ("200000 100000" is 2 cores, "max ..." is none), else cgroup v1 +// cfs quota and period (-1 is none). +func cpuQuota(root string) (float64, bool) { + if raw, err := os.ReadFile(filepath.Join(root, "cpu.max")); err == nil { + fields := strings.Fields(string(raw)) + if len(fields) == 2 && fields[0] != "max" { + return ratio(fields[0], fields[1]) + } + return 0, false + } + for _, dir := range []string{"cpu", "cpu,cpuacct"} { + quota, qerr := os.ReadFile(filepath.Join(root, dir, "cpu.cfs_quota_us")) + period, perr := os.ReadFile(filepath.Join(root, dir, "cpu.cfs_period_us")) + if qerr == nil && perr == nil { + return ratio(strings.TrimSpace(string(quota)), strings.TrimSpace(string(period))) + } + } + return 0, false +} + +func ratio(quota, period string) (float64, bool) { + q, qerr := strconv.ParseFloat(quota, 64) + p, perr := strconv.ParseFloat(period, 64) + if qerr != nil || perr != nil || q <= 0 || p <= 0 { + return 0, false + } + return q / p, true +} + +// DetectHost describes the machine running seedstorm: usable CPUs (a container +// quota when there is one) and memory (a container limit, else the total). +func DetectHost() Host { + cpus := float64(runtime.NumCPU()) + if quota, ok := cpuQuota(cgroupRoot); ok && quota < cpus { + cpus = quota + } + return Host{CPUs: cpus, MemoryMB: memoryLimitMB("/")} +} + +// memoryLimitMB reads the memory limit under root: cgroup v2 memory.max, cgroup +// v1 memory.limit_in_bytes, else /proc/meminfo MemTotal. 0 when unknown. +func memoryLimitMB(root string) int { + total := meminfoTotalMB(root) + for _, rel := range []string{"sys/fs/cgroup/memory.max", "sys/fs/cgroup/memory/memory.limit_in_bytes"} { + raw, err := os.ReadFile(filepath.Join(root, rel)) + if err != nil { + continue + } + v := strings.TrimSpace(string(raw)) + if v == "max" { + break + } + n, err := strconv.ParseInt(v, 10, 64) + if err != nil || n <= 0 { + break + } + mb := int(n >> 20) + // cgroup v1 reports "unlimited" as a huge page-aligned number. + if total > 0 && mb >= total { + break + } + return mb + } + return total +} + +func meminfoTotalMB(root string) int { + raw, err := os.ReadFile(filepath.Join(root, "proc/meminfo")) + if err != nil { + return 0 + } + for _, line := range strings.Split(string(raw), "\n") { + if strings.HasPrefix(line, "MemTotal:") { + fields := strings.Fields(line) + if len(fields) >= 2 { + kb, _ := strconv.Atoi(fields[1]) + return kb / 1024 + } + } + } + return 0 +} diff --git a/internal/tuning/host_test.go b/internal/tuning/host_test.go new file mode 100644 index 0000000..86e2c08 --- /dev/null +++ b/internal/tuning/host_test.go @@ -0,0 +1,103 @@ +package tuning + +import ( + "os" + "path/filepath" + "testing" +) + +func writeFile(t *testing.T, root, rel, body string) { + t.Helper() + p := filepath.Join(root, rel) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte(body), 0o644); err != nil { + t.Fatal(err) + } +} + +// Containers limit CPU with a quota, which runtime.NumCPU ignores (it reported +// 16 inside --cpus=2). Go's GOMAXPROCS follows the quota but never goes below 2, +// so a 1-CPU or half-CPU container still looked like 2 cores. +func TestCPUQuota_ReadsCgroupLimits(t *testing.T) { + cases := []struct { + name string + files map[string]string + want float64 + found bool + }{ + {"v2 two cpus", map[string]string{"cpu.max": "200000 100000\n"}, 2, true}, + {"v2 half a cpu", map[string]string{"cpu.max": "50000 100000\n"}, 0.5, true}, + {"v2 unlimited", map[string]string{"cpu.max": "max 100000\n"}, 0, false}, + {"v1 one and a half", map[string]string{"cpu/cpu.cfs_quota_us": "150000\n", "cpu/cpu.cfs_period_us": "100000\n"}, 1.5, true}, + {"v1 unlimited", map[string]string{"cpu/cpu.cfs_quota_us": "-1\n", "cpu/cpu.cfs_period_us": "100000\n"}, 0, false}, + {"no cgroup files", map[string]string{}, 0, false}, + {"garbage", map[string]string{"cpu.max": "abc def\n"}, 0, false}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + root := t.TempDir() + for rel, body := range c.files { + writeFile(t, root, rel, body) + } + got, found := cpuQuota(root) + if got != c.want || found != c.found { + t.Fatalf("cpuQuota = %v, %v; want %v, %v", got, found, c.want, c.found) + } + }) + } +} + +func TestMaxGenerators_FollowsTheQuotaNotTheHostCores(t *testing.T) { + cases := []struct { + hostCPUs int + quota float64 + limited bool + want int + }{ + {16, 2, true, 2}, + {16, 0.5, true, 1}, + {16, 1.5, true, 1}, + {16, 0, false, 16}, + {4, 8, true, 4}, // a quota above the host's cores changes nothing + {1, 0, false, 1}, + } + for _, c := range cases { + if got := maxGenerators(c.hostCPUs, c.quota, c.limited); got != c.want { + t.Errorf("maxGenerators(host=%d, quota=%v, limited=%v) = %d, want %d", c.hostCPUs, c.quota, c.limited, got, c.want) + } + } +} + +// A container's memory limit (cgroup v2 memory.max, v1 limit_in_bytes) is the +// memory seedstorm may use, not the host's total. +func TestMemoryLimit_ReadsCgroupThenMeminfo(t *testing.T) { + root := t.TempDir() + writeFile(t, root, "sys/fs/cgroup/memory.max", "536870912\n") + writeFile(t, root, "proc/meminfo", "MemTotal: 32000000 kB\nMemAvailable: 9000000 kB\n") + if got := memoryLimitMB(root); got != 512 { + t.Fatalf("cgroup v2 limit: %dMB, want 512", got) + } + + root = t.TempDir() + writeFile(t, root, "sys/fs/cgroup/memory.max", "max\n") + writeFile(t, root, "proc/meminfo", "MemTotal: 32000000 kB\n") + if got := memoryLimitMB(root); got != 31250 { + t.Fatalf("no cgroup limit: %dMB, want MemTotal 31250", got) + } + + root = t.TempDir() + writeFile(t, root, "sys/fs/cgroup/memory/memory.limit_in_bytes", "1073741824\n") + writeFile(t, root, "proc/meminfo", "MemTotal: 32000000 kB\n") + if got := memoryLimitMB(root); got != 1024 { + t.Fatalf("cgroup v1 limit: %dMB, want 1024", got) + } + + root = t.TempDir() + writeFile(t, root, "sys/fs/cgroup/memory/memory.limit_in_bytes", "9223372036854771712\n") + writeFile(t, root, "proc/meminfo", "MemTotal: 2048000 kB\n") + if got := memoryLimitMB(root); got != 2000 { + t.Fatalf("cgroup v1 'unlimited' sentinel: %dMB, want MemTotal 2000", got) + } +} diff --git a/internal/tuning/recommend.go b/internal/tuning/recommend.go new file mode 100644 index 0000000..0161d80 --- /dev/null +++ b/internal/tuning/recommend.go @@ -0,0 +1,248 @@ +package tuning + +import ( + "fmt" + "math" + "strings" +) + +// StorageType is the kind of disk behind a database. +type StorageType string + +const ( + StorageUnknown StorageType = "" + StorageLocalSSD StorageType = "local-ssd" + StorageNetworkSSD StorageType = "network-ssd" + StorageHDD StorageType = "hdd" +) + +// Host is the machine running seedstorm. +type Host struct { + CPUs float64 + MemoryMB int +} + +// Database is what is known about the server a run writes to: detected over +// SQL (connections, buffers, used space) or entered by the user (size, disk). +type Database struct { + Engine string // postgres | mysql + VCPU float64 + MemoryMB int + Storage StorageType + StorageGB int + IOPS int + Shared bool // other workloads use the server + HA bool // synchronous replication (high availability) + Production bool + MaxConnections int + UsedConnections int + // Memory the server reserves (MySQL buffer pool + log buffer, Postgres + // shared buffers); what is left serves connections. + BufferPoolBytes int64 + LogBufferBytes int64 + SharedBuffersBytes int64 + MaxAllowedPacket int64 + UsedBytes int64 +} + +// Run is the shape of the rows to write. +type Run struct { + Rows int64 + AvgRowBytes int + // IndexFactor is index bytes per data byte (0.5: indexes add half). + IndexFactor float64 +} + +// GrowthStatus is the verdict of the disk growth check. +type GrowthStatus string + +const ( + GrowthOK GrowthStatus = "ok" + GrowthWarn GrowthStatus = "warn" + GrowthRefuse GrowthStatus = "refuse" + GrowthUnknown GrowthStatus = "unknown" +) + +// Growth estimates how much the run adds to the database's disk. +type Growth struct { + Status GrowthStatus `json:"status"` + ExpectedBytes int64 `json:"expectedBytes"` + FreeBytes int64 `json:"freeBytes"` + Message string `json:"message"` +} + +// Recommendation is a starting point, with the reason for each value. +type Recommendation struct { + Writers int `json:"writers"` + Generators int `json:"generators"` + ChunkBytes int `json:"chunkBytes"` + BatchBytes int `json:"batchBytes"` + Reasons []string `json:"reasons"` + Growth Growth `json:"growth"` +} + +const ( + maxWriters = 32 + defaultChunkBytes = 32 << 20 + minChunkBytes = 4 << 20 + defaultBatchBytes = 1 << 20 + // iopsPerWriter: a writer's batches need about this many IOPS on a + // network disk before a second writer only adds contention. The probe on a + // 300-IOPS Cloud SQL-sized MySQL still gained from 4 writers. + iopsPerWriter = 75 + // perWriterMB is a connection's working memory on the server. + perWriterMB = 8 + // serverOverheadMB is what the server needs besides buffers and connections. + serverOverheadMB = 150 + // singleCoreRowsPerSec is about what one generator produces + // (docs/benchmarks.md); only a database faster than that benefits from more. + fastWriterThreshold = 8 +) + +// Recommend suggests writers, generators, chunk and batch sizes for a run, and +// checks the disk has room. Constants are starting points calibrated against +// docs/benchmarks.md and the loadsim measurements; the numbers depend on the +// database's load, so the page and CLI always say so. +func Recommend(host Host, db Database, run Run) Recommendation { + r := Recommendation{ChunkBytes: defaultChunkBytes, BatchBytes: defaultBatchBytes} + type cap struct { + n int + why string + } + var caps []cap + add := func(n int, why string) { caps = append(caps, cap{max(1, n), why}) } + + if db.VCPU > 0 { + perCPU := 2.0 + why := fmt.Sprintf("%g vCPU on the database, 2 writers each", db.VCPU) + if db.Shared { + perCPU = 1 + why = fmt.Sprintf("%g vCPU on the database, shared with other workloads", db.VCPU) + } + add(int(math.Ceil(db.VCPU*perCPU)), why) + } + if db.MaxConnections > 0 { + free := db.MaxConnections - db.UsedConnections + headroom := max(5, min(db.MaxConnections/5, free/2)) + add(free-headroom, fmt.Sprintf("%d free connections of %d, leaving %d for applications", free, db.MaxConnections, headroom)) + } + if db.MemoryMB > 0 { + reserved := (db.BufferPoolBytes+db.LogBufferBytes+db.SharedBuffersBytes)/(1<<20) + serverOverheadMB + if left := int64(db.MemoryMB) - reserved; left > 0 { + add(int(left/perWriterMB), fmt.Sprintf("%dMB of database memory left for connections", left)) + } else { + add(1, "database memory is fully reserved by its buffers") + } + } + storageCap := 0 + storageWhy := "" + switch db.Storage { + case StorageHDD: + storageCap, storageWhy = 2, "spinning disk (storage): writes wait on seeks" + case StorageNetworkSSD: + if db.IOPS > 0 { + storageCap, storageWhy = max(2, db.IOPS/iopsPerWriter), fmt.Sprintf("%d IOPS on a network disk", db.IOPS) + } else { + storageCap, storageWhy = 4, "network disk of unknown IOPS" + } + } + if storageCap > 0 && db.HA { + storageCap, storageWhy = max(1, storageCap/2), storageWhy+", halved for high availability (synchronous replication)" + } + if storageCap > 0 { + add(storageCap, storageWhy) + } + add(maxWriters, fmt.Sprintf("at most %d writers per run", maxWriters)) + + writers, why := maxWriters, "" + for _, c := range caps { + if c.n < writers { + writers, why = c.n, c.why + } + } + if db.Production { + limit := 4 + if db.Shared { + limit = 2 + } + if writers > limit { + writers, why = limit, fmt.Sprintf("production database: at most %d writers", limit) + } + } + r.Writers = writers + r.Reasons = append(r.Reasons, fmt.Sprintf("writers %d: %s", writers, why)) + + // Generators help only when the database takes rows faster than one core + // makes them. + r.Generators = 1 + genWhy := "one generator keeps up with this database" + fast := strings.EqualFold(db.Engine, "postgres") && writers >= fastWriterThreshold && !db.Shared && db.Storage != StorageHDD + hostCores := int(math.Floor(host.CPUs)) + if fast && hostCores > 2 { + r.Generators = hostCores - 1 + genWhy = fmt.Sprintf("a fast Postgres (COPY) outpaces one core; %d of %d host cores", r.Generators, hostCores) + } + if host.MemoryMB > 0 { + budget := int64(host.MemoryMB) * (1 << 20) / 4 + if limit := int(budget/defaultChunkBytes) - 2; r.Generators > max(1, limit) { + r.Generators = max(1, limit) + genWhy = fmt.Sprintf("host memory %dMB fits %d generators' chunks", host.MemoryMB, r.Generators) + } + if int64(r.Generators+2)*int64(r.ChunkBytes) > budget { + r.ChunkBytes = max(minChunkBytes, int(budget/int64(r.Generators+2))) + r.Reasons = append(r.Reasons, fmt.Sprintf("chunks %dMB: queued rows stay under a quarter of the host's %dMB", r.ChunkBytes>>20, host.MemoryMB)) + } + } + r.Reasons = append(r.Reasons, fmt.Sprintf("generators %d: %s", r.Generators, genWhy)) + + if db.MaxAllowedPacket > 0 && db.MaxAllowedPacket/4 < int64(r.BatchBytes) { + r.BatchBytes = int(db.MaxAllowedPacket / 4) + r.Reasons = append(r.Reasons, fmt.Sprintf("batches %dKB: a quarter of max_allowed_packet", r.BatchBytes>>10)) + } + + r.Growth = growth(db, run) + return r +} + +// growth compares the rows a run adds with the free space on the database's disk. +func growth(db Database, run Run) Growth { + expected := int64(float64(run.Rows) * float64(max(run.AvgRowBytes, 1)) * (1 + run.IndexFactor) * 1.5) + g := Growth{ExpectedBytes: expected} + if db.StorageGB <= 0 { + g.Status = GrowthUnknown + g.Message = fmt.Sprintf("About %s will be written; enter the database's storage size to check it fits.", gigabytes(expected)) + return g + } + g.FreeBytes = int64(db.StorageGB)*(1<<30) - db.UsedBytes + share := float64(expected) / float64(max(g.FreeBytes, 1)) + switch { + case g.FreeBytes <= 0 || share > 0.9: + g.Status = GrowthRefuse + g.Message = fmt.Sprintf("About %s would be written (rows, indexes and logs) but only %s is free: the disk could fill up.", gigabytes(expected), gigabytes(max(g.FreeBytes, 0))) + case share > 0.5: + g.Status = GrowthWarn + g.Message = fmt.Sprintf("About %s of the %s free will be used (%.0f%%).", gigabytes(expected), gigabytes(g.FreeBytes), share*100) + default: + g.Status = GrowthOK + g.Message = fmt.Sprintf("About %s of the %s free will be used.", gigabytes(expected), gigabytes(g.FreeBytes)) + } + return g +} + +func gigabytes(b int64) string { + return fmt.Sprintf("%.1fGB", float64(b)/(1<<30)) +} + +// ClampWriters lowers the requested writers only when the server lacks free +// connections for the whole run: writers, generators (they read parent keys) +// and one for sequence updates. An unknown limit changes nothing. +func ClampWriters(requested, generators, maxConnections, usedConnections int) int { + if maxConnections <= 0 || requested <= 1 { + return requested + } + free := maxConnections - usedConnections + if free >= requested+generators+1 { + return requested + } + return max(1, free-generators-1) +} diff --git a/internal/tuning/recommend_test.go b/internal/tuning/recommend_test.go new file mode 100644 index 0000000..9e52d0f --- /dev/null +++ b/internal/tuning/recommend_test.go @@ -0,0 +1,114 @@ +package tuning + +import ( + "strings" + "testing" +) + +const mb = int64(1 << 20) + +func micro() Database { + return Database{ + Engine: "mysql", VCPU: 1, MemoryMB: 629, Storage: StorageNetworkSSD, StorageGB: 10, IOPS: 300, + MaxConnections: 280, UsedConnections: 3, BufferPoolBytes: 53477376, LogBufferBytes: 67108864, + MaxAllowedPacket: 33554432, UsedBytes: 200 * mb, + } +} + +func TestRecommend_WritersFollowTheDatabaseNotItsConnectionLimit(t *testing.T) { + host := Host{CPUs: 8, MemoryMB: 16000} + cases := []struct { + name string + db func() Database + want int + why string + }{ + // max_connections 280 on 629MB says nothing about capacity: storage and + // CPU decide. The probe measured 4 writers faster than 2 at 300 IOPS. + {"cloud sql micro", micro, 2, "vCPU"}, + {"few free connections", func() Database { + d := micro() + d.VCPU, d.MaxConnections, d.UsedConnections, d.IOPS = 8, 100, 90, 3000 + return d + }, 5, "connections"}, + {"hdd", func() Database { d := micro(); d.VCPU, d.Storage, d.IOPS = 16, StorageHDD, 0; return d }, 2, "storage"}, + {"network ssd by iops", func() Database { d := micro(); d.VCPU, d.IOPS = 16, 450; return d }, 6, "IOPS"}, + {"production shared", func() Database { + d := micro() + d.VCPU, d.IOPS, d.Production, d.Shared = 16, 30000, true, true + return d + }, 2, "production"}, + {"high availability halves storage", func() Database { d := micro(); d.VCPU, d.IOPS, d.HA = 16, 600, true; return d }, 4, "high availability"}, + {"never above 32", func() Database { + d := micro() + d.VCPU, d.Storage, d.IOPS, d.MaxConnections = 64, StorageLocalSSD, 0, 5000 + return d + }, 32, ""}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + r := Recommend(host, c.db(), Run{Rows: 10000, AvgRowBytes: 200}) + if r.Writers != c.want { + t.Fatalf("writers = %d, want %d (reasons %v)", r.Writers, c.want, r.Reasons) + } + if c.why != "" && !strings.Contains(strings.Join(r.Reasons, " | "), c.why) { + t.Fatalf("reasons %v do not mention %q", r.Reasons, c.why) + } + }) + } +} + +func TestRecommend_GeneratorsChunksAndBatches(t *testing.T) { + r := Recommend(Host{CPUs: 2, MemoryMB: 16000}, micro(), Run{Rows: 1000, AvgRowBytes: 100}) + if r.Generators != 1 { + t.Errorf("generators on a 2-core host = %d, want 1", r.Generators) + } + if r.BatchBytes > int(33554432/4) || r.BatchBytes > 1<<20 { + t.Errorf("batch bytes = %d, want <= 1MB and <= max_allowed_packet/4", r.BatchBytes) + } + + fast := Database{Engine: "postgres", VCPU: 16, MemoryMB: 64000, Storage: StorageLocalSSD, MaxConnections: 400} + r = Recommend(Host{CPUs: 16, MemoryMB: 32000}, fast, Run{Rows: 10_000_000, AvgRowBytes: 200}) + if r.Generators < 2 || r.Generators > 15 { + t.Errorf("generators for a fast Postgres = %d, want several but below the host's cores", r.Generators) + } + + small := Recommend(Host{CPUs: 16, MemoryMB: 512}, fast, Run{Rows: 10_000_000, AvgRowBytes: 200}) + if int64(small.Generators+2)*int64(small.ChunkBytes) > 512*mb/4 { + t.Errorf("chunks do not fit a 512MB host: %d generators × %d bytes", small.Generators, small.ChunkBytes) + } +} + +func TestRecommend_GrowthCheck(t *testing.T) { + d := micro() // 10GB, 200MB used + ok := Recommend(Host{CPUs: 2, MemoryMB: 4000}, d, Run{Rows: 1_000_000, AvgRowBytes: 2000, IndexFactor: 0.5}) + if ok.Growth.Status != GrowthWarn && ok.Growth.Status != GrowthOK { + t.Fatalf("1M×2KB on 10GB: %+v", ok.Growth) + } + d.StorageGB = 2 + refused := Recommend(Host{CPUs: 2, MemoryMB: 4000}, d, Run{Rows: 1_000_000, AvgRowBytes: 2000, IndexFactor: 0.5}) + if refused.Growth.Status != GrowthRefuse || !strings.Contains(refused.Growth.Message, "GB") { + t.Fatalf("1M×2KB on 2GB: %+v", refused.Growth) + } + d.StorageGB = 0 + unknown := Recommend(Host{CPUs: 2, MemoryMB: 4000}, d, Run{Rows: 1_000_000, AvgRowBytes: 2000}) + if unknown.Growth.Status != GrowthUnknown { + t.Fatalf("no storage size: %+v", unknown.Growth) + } +} + +// Run-start clamps only lower writers when the server truly lacks free +// connections for the run (writers + generators + one for sequences). +func TestClampWriters_OnlyWhenConnectionsRunOut(t *testing.T) { + cases := []struct{ requested, generators, max, used, want int }{ + {4, 1, 100, 85, 4}, // 15 free: 4+1+1 fit + {4, 1, 100, 96, 2}, // 4 free: 2 writers + 1 generator + 1 + {8, 1, 25, 24, 1}, // 1 free: never below 1 + {4, 1, 0, 0, 4}, // limit unknown: unchanged + } + for _, c := range cases { + if got := ClampWriters(c.requested, c.generators, c.max, c.used); got != c.want { + t.Errorf("ClampWriters(%d, gen %d, max %d, used %d) = %d, want %d", c.requested, c.generators, c.max, c.used, got, c.want) + } + } +} diff --git a/internal/web/counts_cache_test.go b/internal/web/counts_cache_test.go new file mode 100644 index 0000000..3de5c0e --- /dev/null +++ b/internal/web/counts_cache_test.go @@ -0,0 +1,73 @@ +package web + +import ( + "database/sql" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +// Leaving the workspace and coming back re-ran COUNT(*) on every table before +// the graph appeared. The graph is structure only; counts are counted once per +// session, cached with when they were taken, and recounted on request or after +// a run that writes. +func TestWorkspaceCounts_GraphDoesNotCountAndCountsAreCached(t *testing.T) { + registerServeRunnerTestDriver() + conn, err := sql.Open(serveRunnerTestDriverName, "counted") + if err != nil { + t.Fatal(err) + } + defer conn.Close() + s, err := New(testOptions(t)) + if err != nil { + t.Fatal(err) + } + sess := &Session{ID: "ws", DBType: "pgx", conn: conn, schema: runnerRowCountSchema()} + s.sessions.add(sess) + get := func(path string) (*httptest.ResponseRecorder, map[string]any) { + r := httptest.NewRequest(http.MethodGet, path, nil) + r.AddCookie(&http.Cookie{Name: sessionCookieName, Value: sess.ID}) + w := httptest.NewRecorder() + s.Handler().ServeHTTP(w, r) + body := map[string]any{} + _ = json.Unmarshal(w.Body.Bytes(), &body) + return w, body + } + + countedQueries.Store(0) + _, graph := get("/api/graph") + if n := countedQueries.Load(); n != 0 { + t.Fatalf("the graph ran %d queries; it must serve structure without counting", n) + } + if graph["countsTakenAt"] != nil && graph["countsTakenAt"] != "" { + t.Fatalf("graph claims counts before any were taken: %v", graph["countsTakenAt"]) + } + + w, counts := get("/api/counts") + if counts["users"] != float64(3) || w.Header().Get("X-Counts-Taken-At") == "" { + t.Fatalf("counts = %v, header %q", counts, w.Header().Get("X-Counts-Taken-At")) + } + first := countedQueries.Load() + + get("/api/counts") + if n := countedQueries.Load(); n != first { + t.Fatalf("a second counts request recounted (%d queries)", n-first) + } + _, graph = get("/api/graph") + if graph["countsTakenAt"] == "" || graph["countsTakenAt"] == nil { + t.Fatal("the graph does not carry the cached counts' time") + } + + get("/api/counts?refresh=1") + if n := countedQueries.Load(); n == first { + t.Fatal("refresh=1 did not recount") + } + + sess.InvalidateCounts() + before := countedQueries.Load() + get("/api/counts") + if countedQueries.Load() == before { + t.Fatal("counts were not recounted after being invalidated") + } +} diff --git a/internal/web/failures.go b/internal/web/failures.go new file mode 100644 index 0000000..065e49e --- /dev/null +++ b/internal/web/failures.go @@ -0,0 +1,56 @@ +package web + +import ( + "errors" + + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" + "github.com/AxeForging/seedstorm/internal/seeder" +) + +// failureView describes where a run failed, for the page to show next to the +// error text. +func failureView(err error) map[string]any { + out := map[string]any{"message": err.Error()} + if e, ok := runerr.As(err); ok { + out["side"], out["phase"], out["table"] = e.Side, string(e.Phase), e.Table + } + var p *safego.PanicError + if errors.As(err, &p) { + out["internal"] = true + out["errorId"] = p.ID + } + return out +} + +// partialSeedResult is what a failed seed or fill had written when it stopped: +// rows per table (0 for tables not reached), where it failed, and what to do. +func partialSeedResult(res seeder.SeedResult, order []string, err error) map[string]any { + counts := make(map[string]int, len(order)) + var written, notWritten []string + for _, t := range order { + counts[t] = res.Counts[t] + if res.Counts[t] > 0 { + written = append(written, t) + } else { + notWritten = append(notWritten, t) + } + } + next := "Fix the cause shown above, then run again." + if e, ok := runerr.As(err); ok && e.Phase == runerr.PhaseWrite && len(written) > 0 { + next = "Fill empty tables continues with the tables that were not written; seeding everything again adds rows to the ones that were." + } + var p *safego.PanicError + if errors.As(err, &p) { + next = "This is a bug in seedstorm: the server log has the details under error id " + p.ID + "." + } + return map[string]any{ + "partial": true, + "totalRows": res.Total, + "tableCounts": counts, + "written": written, + "notWritten": notWritten, + "failure": failureView(err), + "nextStep": next, + } +} diff --git a/internal/web/handlers_api.go b/internal/web/handlers_api.go index 8a28086..387834c 100644 --- a/internal/web/handlers_api.go +++ b/internal/web/handlers_api.go @@ -8,6 +8,7 @@ import ( "sort" "strconv" "strings" + "time" "github.com/AxeForging/seedstorm/internal/db" "github.com/AxeForging/seedstorm/internal/graph" @@ -35,6 +36,9 @@ type graphPayload struct { Edges []graphEdge `json:"edges"` Order []string `json:"order,omitempty"` Cycle bool `json:"cycle"` + // CountsTakenAt is when the node counts were taken; empty when the graph + // carries no counts yet (the page asks /api/counts). + CountsTakenAt string `json:"countsTakenAt,omitempty"` } type tablePreviewPayload struct { @@ -57,17 +61,14 @@ func (s *Server) handleGraphJSON(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusInternalServerError, err.Error()) return } - // Row counts are best-effort; if the COUNT(*) sweep fails (e.g., a missing - // table after a DDL edit) we still serve the structure. - counts := map[string]int64{} - tableNames := make([]string, 0, len(sc.Tables)) - for n := range sc.Tables { - tableNames = append(tableNames, n) - } - if c, cerr := db.GetTableRowCounts(r.Context(), sess.Conn(), sess.DBType, tableNames); cerr == nil { - counts = c - } + // The graph is structure only, with counts the session already holds: + // counting every table here made each return to the workspace wait on a + // COUNT(*) per table before anything was drawn. + counts, at := sess.CachedCounts() payload := buildGraphPayload(sc, counts) + if counts != nil { + payload.CountsTakenAt = at.Format(time.RFC3339) + } writeJSON(w, http.StatusOK, payload) } @@ -89,11 +90,14 @@ func (s *Server) handleCountsJSON(w http.ResponseWriter, r *http.Request) { tables = append(tables, n) } sort.Strings(tables) - counts, cerr := db.GetTableRowCounts(r.Context(), sess.Conn(), sess.DBType, tables) - if cerr != nil { - writeError(w, http.StatusInternalServerError, cerr.Error()) + // Counts are cached per session; ?refresh=1 recounts. Tables whose count + // fails are left out: the page keeps them uncounted instead of empty. + counts, at := sess.Counts(r.Context(), tables, r.URL.Query().Get("refresh") == "1") + if len(counts) == 0 && len(tables) > 0 { + writeError(w, http.StatusInternalServerError, "no table could be counted (see the server log)") return } + w.Header().Set("X-Counts-Taken-At", at.Format(time.RFC3339)) writeJSON(w, http.StatusOK, counts) } @@ -218,7 +222,17 @@ func clampQueryInt(r *http.Request, key string, def, min, max int) int { return n } -func loadTablePreview(ctx context.Context, conn *sql.DB, dbType, tableName string, columns []string, limit, offset int) (tablePreviewPayload, error) { +// loadTablePreview reads one page of a table in a read-only transaction that +// never waits behind a lock. +func loadTablePreview(ctx context.Context, conn *sql.DB, dbType, tableName string, columns []string, limit, offset int) (payload tablePreviewPayload, err error) { + err = db.ReadOnce(ctx, conn, dbType, db.DefaultCountLimits, func(ctx context.Context, q db.Querier) error { + payload, err = readTablePreview(ctx, q, dbType, tableName, columns, limit, offset) + return err + }) + return payload, err +} + +func readTablePreview(ctx context.Context, conn db.Querier, dbType, tableName string, columns []string, limit, offset int) (tablePreviewPayload, error) { payload := tablePreviewPayload{ Table: tableName, Limit: limit, diff --git a/internal/web/handlers_compare.go b/internal/web/handlers_compare.go index 489576b..aa10175 100644 --- a/internal/web/handlers_compare.go +++ b/internal/web/handlers_compare.go @@ -8,6 +8,7 @@ import ( "github.com/AxeForging/seedstorm/internal/compare" "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/seeder" ) @@ -38,7 +39,11 @@ func (s *Server) resolveConnection(ref ConnRef, role string) (*Session, error) { return nil, err } info.Label = saved.Label - return s.sessions.OpenDSN(driver, dsn, info) + sess, err := s.sessions.OpenDSN(driver, dsn, info) + if err == nil && sess.SavedID == "" { + sess.SavedID = saved.ID + } + return sess, err } if ref.ID != "" { sess, ok := s.sessions.Get(ref.ID) @@ -101,17 +106,17 @@ func (s *Server) runCompare(ctx context.Context, _ *Session, req CompareRequest, return nil, err } jc.Phase("connect") - srcEP, source, err := s.sourceEndpoint(ctx, req.Source, req.SourceSnapshot) + source, target, err := s.connectBoth(ctx, log, req.Source, req.SourceSnapshot, req.Target) if err != nil { return nil, err } - target, err := s.resolveConnection(req.Target, "target") + srcEP, err := snapshotOrSessionEndpoint(ctx, source, req.SourceSnapshot) if err != nil { - return nil, err + return nil, runerr.OnSide(runerr.SideSource, runerr.At(runerr.PhaseIntrospect, "", err)) } tgtEP, _, err := endpointFor(ctx, target, false) if err != nil { - return nil, err + return nil, runerr.OnSide(runerr.SideTarget, runerr.At(runerr.PhaseIntrospect, "", err)) } jc.Phase("count") log.Info().Str("source", srcEP.Label).Str("target", tgtEP.Label).Str("counts", string(mode)).Msg("Reading table volumes") @@ -127,27 +132,6 @@ func (s *Server) runCompare(ctx context.Context, _ *Session, req CompareRequest, return map[string]any{"report": report, "sameConnection": source != nil && source.ID == target.ID, "sourceIsSnapshot": source == nil}, nil } -// sourceEndpoint is the source side of a compare or mirror: an imported -// snapshot when one is given (the returned session is nil), else a connection. -func (s *Server) sourceEndpoint(ctx context.Context, ref ConnRef, snap *compare.Snapshot) (seeder.Endpoint, *Session, error) { - if snap != nil { - if len(snap.Tables) == 0 { - return seeder.Endpoint{}, nil, fmt.Errorf("the imported counts have no tables") - } - label := snap.Label - if label == "" { - label = "imported counts" - } - return seeder.Endpoint{Snapshot: snap, Label: label, DBType: snap.DBType}, nil, nil - } - sess, err := s.resolveConnection(ref, "source") - if err != nil { - return seeder.Endpoint{}, nil, err - } - ep, _, err := endpointFor(ctx, sess, false) - return ep, sess, err -} - // MirrorRequest seeds the target so its volumes follow the source. type MirrorRequest struct { Source ConnRef `json:"source"` @@ -166,10 +150,18 @@ type MirrorRequest struct { StopOnError bool `json:"stopOnError"` DryRun bool `json:"dryRun"` PreviewRows int `json:"previewRows"` + // ConfirmProduction is the target's label, typed to write to a + // production connection. + ConfirmProduction string `json:"confirmProduction,omitempty"` } func (s *Server) handleMirrorRun(w http.ResponseWriter, r *http.Request) { - startRun(s, w, r, "mirror", s.runMirror) + startGuardedRun(s, w, r, "mirror", s.runMirror, func(req MirrorRequest, _ *Session) *productionRefusal { + if req.DryRun { + return nil + } + return s.guardProduction(s.refTarget(req.Target), req.ConfirmProduction, "mirror into it") + }) } func (s *Server) runMirror(ctx context.Context, _ *Session, req MirrorRequest, jc JobControl) (map[string]any, error) { @@ -187,17 +179,20 @@ func (s *Server) runMirror(ctx context.Context, _ *Session, req MirrorRequest, j return nil, err } jc.Phase("connect") - srcEP, _, err := s.sourceEndpoint(ctx, req.Source, req.SourceSnapshot) + source, target, err := s.connectBoth(ctx, log, req.Source, req.SourceSnapshot, req.Target) if err != nil { return nil, err } - target, err := s.resolveConnection(req.Target, "target") + srcEP, err := snapshotOrSessionEndpoint(ctx, source, req.SourceSnapshot) if err != nil { - return nil, err + return nil, runerr.OnSide(runerr.SideSource, runerr.At(runerr.PhaseIntrospect, "", err)) + } + if !req.DryRun { + defer target.InvalidateCounts() } tgtEP, closeTarget, err := endpointFor(ctx, target, !req.DryRun) if err != nil { - return nil, err + return nil, runerr.OnSide(runerr.SideTarget, runerr.At(runerr.PhaseConnect, "", err)) } defer closeTarget() @@ -282,6 +277,7 @@ func (s *Server) runMirror(ctx context.Context, _ *Session, req MirrorRequest, j log.Warn().Str("table", problem.Table).Str("status", problem.Status).Int64("missing", problem.Missing).Msg(problem.Error) } if runErr != nil { + result["failure"] = failureView(runErr) return result, runErr } jc.Phase("done") diff --git a/internal/web/handlers_connections.go b/internal/web/handlers_connections.go index c510027..8efd6da 100644 --- a/internal/web/handlers_connections.go +++ b/internal/web/handlers_connections.go @@ -13,19 +13,23 @@ import ( // connectForm is everything the connect page collects, kept as strings so a // failed attempt can be re-rendered exactly as the user typed it. type connectForm struct { - ID string - Label string - DBType string - Host string - Port string - DBName string - User string - SSL string - DSN string - Password string - Params []Param - Save bool - SavePassword bool + ID string + Label string + DBType string + Host string + Port string + DBName string + User string + SSL string + DSN string + Password string + Params []Param + Save bool + SavePassword bool + // Production marks the connection as a production database (see + // SavedConnection.Production); ConfirmLabel is typed to clear it. + Production bool + ConfirmLabel string NeedsPassword bool // Origin is "edit" or "duplicate" when the form was opened from a saved // connection, so the page can say which it is. @@ -76,6 +80,8 @@ func parseConnectForm(r *http.Request) connectForm { Password: r.FormValue("password"), Save: isChecked(r.FormValue("save")), SavePassword: isChecked(r.FormValue("savePassword")), + Production: isChecked(r.FormValue("production")), + ConfirmLabel: strings.TrimSpace(r.FormValue("confirmLabel")), } names := r.Form["paramName"] values := r.Form["paramValue"] @@ -118,11 +124,12 @@ func (f connectForm) info() ConnectionInfo { func (f connectForm) saved() SavedConnection { port, _ := strconv.Atoi(f.Port) c := SavedConnection{ - ID: f.ID, - Label: f.Label, - DBType: f.DBType, - DSN: f.DSN, - Params: f.Params, + ID: f.ID, + Label: f.Label, + DBType: f.DBType, + DSN: f.DSN, + Params: f.Params, + Production: f.Production, } if f.DSN == "" { c.Host = f.Host @@ -343,6 +350,7 @@ func (s *Server) handleConnectSaved(w http.ResponseWriter, r *http.Request) { s.renderConnect(w, r, page, err.Error()) return } + sess.SavedID = saved.ID _ = s.store.Touch(saved.ID) setSessionCookie(w, sess.ID) http.Redirect(w, r, "/", http.StatusSeeOther) @@ -356,16 +364,17 @@ func needsPassword(c SavedConnection) bool { func savedToForm(c SavedConnection, password string) connectForm { f := connectForm{ - ID: c.ID, - Label: c.Label, - DBType: c.DBType, - Host: c.Host, - DBName: c.DBName, - User: c.User, - SSL: c.SSL, - DSN: c.DSN, - Params: c.Params, - Password: password, + ID: c.ID, + Label: c.Label, + DBType: c.DBType, + Host: c.Host, + DBName: c.DBName, + User: c.User, + SSL: c.SSL, + DSN: c.DSN, + Params: c.Params, + Password: password, + Production: c.Production, } if c.Port > 0 { f.Port = strconv.Itoa(c.Port) @@ -397,6 +406,7 @@ func (s *Server) handleSavedConnections(w http.ResponseWriter, r *http.Request) SavedConnection Password string `json:"password"` ClearPassword bool `json:"clearPassword"` + ConfirmLabel string `json:"confirmLabel"` } if err := json.NewDecoder(io.LimitReader(r.Body, maxProfileBody)).Decode(&body); err != nil { writeError(w, http.StatusBadRequest, "bad json: "+err.Error()) @@ -415,6 +425,10 @@ func (s *Server) handleSavedConnections(w http.ResponseWriter, r *http.Request) writeError(w, http.StatusBadRequest, err.Error()) return } + if refusal := s.unflagRefusal(conn, body.ConfirmLabel); refusal != nil { + writeProductionRefusal(w, refusal) + return + } out, err := s.store.Save(conn, body.ClearPassword) if err != nil { writeError(w, http.StatusInternalServerError, err.Error()) diff --git a/internal/web/handlers_jobs.go b/internal/web/handlers_jobs.go index 18b27ba..b78bc62 100644 --- a/internal/web/handlers_jobs.go +++ b/internal/web/handlers_jobs.go @@ -3,7 +3,9 @@ package web import ( "fmt" "net/http" + "strconv" "strings" + "time" ) // handleJobsAPI dispatches GET/POST under /api/jobs/. @@ -45,6 +47,37 @@ func (s *Server) handleJobsAPI(w http.ResponseWriter, r *http.Request) { } } +// streamKeepalive is how often a quiet job stream sends a ping event, so the page +// can tell a slow step from a lost server. +var streamKeepalive = 15 * time.Second + +// handleJobList lists the jobs of the requesting session: +// +// GET /api/jobs -> {bootId, jobs: [{id, name, status, phase, progress, ...}]} +func (s *Server) handleJobList(w http.ResponseWriter, r *http.Request) { + sess, err := s.sessions.fromRequest(r) + if err != nil { + writeJSON(w, http.StatusOK, map[string]any{"bootId": s.bootID, "jobs": []any{}}) + return + } + s.writeJobList(w, sess.ID) +} + +func (s *Server) writeJobList(w http.ResponseWriter, owner string) { + jobs := []map[string]any{} + for _, j := range s.jobs.ForOwner(owner) { + view := jobView(j) + delete(view, "result") + phase, progress := j.Position() + view["phase"] = phase + if progress != nil { + view["progress"] = map[string]any{"done": progress.Done, "total": progress.Total, "label": progress.Text} + } + jobs = append(jobs, view) + } + writeJSON(w, http.StatusOK, map[string]any{"bootId": s.bootID, "jobs": jobs}) +} + func (s *Server) streamJob(w http.ResponseWriter, r *http.Request, job *Job) { flusher, ok := w.(http.Flusher) if !ok { @@ -56,34 +89,41 @@ func (s *Server) streamJob(w http.ResponseWriter, r *http.Request, job *Job) { w.Header().Set("Connection", "keep-alive") w.Header().Set("X-Accel-Buffering", "no") + // A reconnecting client passes the last event it saw (?after= or the + // browser's Last-Event-ID header) so the backlog is not replayed twice. + after := lastSeenEvent(r) + ch, backlog := job.Subscribe() defer job.Unsubscribe(ch) - maxSeq := 0 + maxSeq := after for _, ev := range backlog { - writeEvent(w, ev) if ev.Seq > maxSeq { + writeEvent(w, ev) maxSeq = ev.Seq } } flusher.Flush() + keepalive := time.NewTicker(streamKeepalive) + defer keepalive.Stop() for { select { case <-r.Context().Done(): return + case <-keepalive.C: + // A named event, not an SSE comment: the page cannot see comments, + // and it needs the ping to tell a quiet job from a lost server. + writeSSE(w, "ping", "") + flusher.Flush() case ev, alive := <-ch: if !alive { - writeSSE(w, "status", string(job.Status)) - if job.Err != nil { - writeSSE(w, "error", job.Err.Error()) - } - writeSSE(w, "end", "") + writeJobEnd(w, job) flusher.Flush() return } - writeEvent(w, ev) if ev.Seq > maxSeq { + writeEvent(w, ev) maxSeq = ev.Seq } flusher.Flush() @@ -94,20 +134,44 @@ func (s *Server) streamJob(w http.ResponseWriter, r *http.Request, job *Job) { for _, ev := range job.Events() { if ev.Seq > maxSeq { writeEvent(w, ev) + maxSeq = ev.Seq } } - writeSSE(w, "status", string(job.Status)) - if job.Err != nil { - writeSSE(w, "error", job.Err.Error()) - } - writeSSE(w, "end", "") + writeJobEnd(w, job) flusher.Flush() return } } } +// writeJobEnd closes a job stream: status, the failure reason when there is +// one, then end. The reason is sent as "failure", never "error": browsers +// treat a named "error" event as a connection error and fire onerror, which +// closed the stream before "end" arrived. +func writeJobEnd(w http.ResponseWriter, job *Job) { + st := job.State() + writeSSE(w, "status", string(st.Status)) + if st.Err != nil { + writeSSE(w, "failure", st.Err.Error()) + } + writeSSE(w, "end", "") +} + +// lastSeenEvent reads the resume point of a stream request. +func lastSeenEvent(r *http.Request) int { + raw := r.URL.Query().Get("after") + if raw == "" { + raw = r.Header.Get("Last-Event-ID") + } + n, err := strconv.Atoi(strings.TrimSpace(raw)) + if err != nil || n < 0 { + return 0 + } + return n +} + func writeEvent(w http.ResponseWriter, ev Event) { + _, _ = fmt.Fprintf(w, "id: %d\n", ev.Seq) switch ev.Kind { case EventPhase: writeSSE(w, "phase", fmt.Sprintf("[%d] %s", ev.Seq, ev.Text)) diff --git a/internal/web/handlers_jobs_test.go b/internal/web/handlers_jobs_test.go index cb0c417..64f8d58 100644 --- a/internal/web/handlers_jobs_test.go +++ b/internal/web/handlers_jobs_test.go @@ -2,6 +2,7 @@ package web import ( "context" + "errors" "net/http" "net/http/httptest" "strings" @@ -152,3 +153,75 @@ func TestStreamJob_ReplaysBacklog(t *testing.T) { } } } + +// A named SSE event called "error" also fires EventSource.onerror in browsers, +// which closed the stream before "end" arrived: every failed job left the page +// waiting forever. The terminal failure must use another event name. +func TestStreamJob_FailedJobEndsWithFailureNotErrorEvent(t *testing.T) { + m := NewManager() + job := m.Start(context.Background(), "boom", func(ctx context.Context, jc JobControl) (map[string]any, error) { + jc.Phase("connect") + return nil, errors.New("target: ping database: connection refused") + }) + <-job.Done() + + srv := &Server{jobs: m} + w := httptest.NewRecorder() + srv.streamJob(w, httptest.NewRequest(http.MethodGet, "/api/jobs/"+job.ID+"/stream", nil), job) + + body := w.Body.String() + if strings.Contains(body, "event: error\n") { + t.Fatalf("stream uses the reserved 'error' event name:\n%s", body) + } + for _, needle := range []string{ + "event: status\ndata: failed", + "event: failure\ndata: target: ping database: connection refused", + "event: end", + } { + if !strings.Contains(body, needle) { + t.Fatalf("missing %q in:\n%s", needle, body) + } + } + if strings.Index(body, "event: failure") > strings.Index(body, "event: end") { + t.Fatalf("failure must come before end:\n%s", body) + } +} + +// Every job event carries an SSE id so a reconnecting client can resume after +// the last one it saw instead of replaying (and duplicating) the whole log. +func TestStreamJob_ResumesAfterLastSeenEvent(t *testing.T) { + m := NewManager() + job := m.Start(context.Background(), "resume", func(ctx context.Context, jc JobControl) (map[string]any, error) { + jc.Phase("one") + jc.Phase("two") + jc.Phase("three") + return nil, nil + }) + <-job.Done() + srv := &Server{jobs: m} + + w := httptest.NewRecorder() + srv.streamJob(w, httptest.NewRequest(http.MethodGet, "/api/jobs/"+job.ID+"/stream", nil), job) + if body := w.Body.String(); !strings.Contains(body, "id: 1\n") || !strings.Contains(body, "id: 3\n") { + t.Fatalf("events carry no SSE ids:\n%s", body) + } + + for _, req := range []*http.Request{ + httptest.NewRequest(http.MethodGet, "/api/jobs/"+job.ID+"/stream?after=2", nil), + func() *http.Request { + r := httptest.NewRequest(http.MethodGet, "/api/jobs/"+job.ID+"/stream", nil) + r.Header.Set("Last-Event-ID", "2") + return r + }(), + } { + w := httptest.NewRecorder() + srv.streamJob(w, req, job) + body := w.Body.String() + if strings.Contains(body, "] one") || strings.Contains(body, "] two") { + t.Fatalf("resumed stream replayed events already seen:\n%s", body) + } + if !strings.Contains(body, "] three") || !strings.Contains(body, "event: end") { + t.Fatalf("resumed stream lost later events:\n%s", body) + } + } +} diff --git a/internal/web/handlers_pages.go b/internal/web/handlers_pages.go index f012eb8..e97ccf9 100644 --- a/internal/web/handlers_pages.go +++ b/internal/web/handlers_pages.go @@ -79,6 +79,10 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) { fail(err.Error()) return } + if refusal := s.unflagRefusal(form.saved(), form.ConfirmLabel); refusal != nil { + fail(refusal.Error()) + return + } if _, err := s.store.Save(form.saved(), !form.SavePassword); err != nil { fail(err.Error()) return @@ -146,6 +150,8 @@ func (s *Server) handleConnectionsJSON(w http.ResponseWriter, r *http.Request) { ID string `json:"id"` Info ConnectionInfo `json:"info"` Active bool `json:"active"` + // Production: a saved production connection reaches this database. + Production bool `json:"production,omitempty"` } out := []entry{} activeID := "" @@ -153,10 +159,12 @@ func (s *Server) handleConnectionsJSON(w http.ResponseWriter, r *http.Request) { activeID = current.Value } for _, sess := range dedupeConnections(s.sessions.All(), activeID) { + _, production := s.productionConnection(sessionTarget(sess)) out = append(out, entry{ - ID: sess.ID, - Info: sess.Info, - Active: activeID == sess.ID, + ID: sess.ID, + Info: sess.Info, + Active: activeID == sess.ID, + Production: production, }) } writeJSON(w, http.StatusOK, out) diff --git a/internal/web/handlers_runs.go b/internal/web/handlers_runs.go index be3c010..22029a7 100644 --- a/internal/web/handlers_runs.go +++ b/internal/web/handlers_runs.go @@ -8,6 +8,7 @@ import ( "io" "net/http" "os" + "strings" ) // maxRunBody bounds a job request: room for an export document as large as @@ -22,6 +23,19 @@ func startRun[T any]( r *http.Request, jobName string, runner func(ctx context.Context, sess *Session, req T, jc JobControl) (map[string]any, error), +) { + startGuardedRun(s, w, r, jobName, runner, nil) +} + +// startGuardedRun is startRun with a check that runs before the job exists: +// a production refusal answers 409 and nothing starts. +func startGuardedRun[T any]( + s *Server, + w http.ResponseWriter, + r *http.Request, + jobName string, + runner func(ctx context.Context, sess *Session, req T, jc JobControl) (map[string]any, error), + guard func(req T, sess *Session) *productionRefusal, ) { if r.Method != http.MethodPost { writeError(w, http.StatusMethodNotAllowed, "POST required") @@ -46,18 +60,36 @@ func startRun[T any]( writeError(w, http.StatusUnauthorized, err.Error()) return } - job := s.jobs.Start(context.Background(), jobName, func(ctx context.Context, jc JobControl) (map[string]any, error) { + if guard != nil { + if refusal := guard(req, sess); refusal != nil { + writeProductionRefusal(w, refusal) + return + } + } + job := s.jobs.StartFor(context.Background(), sess.ID, jobName, func(ctx context.Context, jc JobControl) (map[string]any, error) { return runner(ctx, sess, req, jc) }) - writeJSON(w, http.StatusAccepted, jobView(job)) + view := jobView(job) + view["bootId"] = s.bootID + writeJSON(w, http.StatusAccepted, view) } func (s *Server) handleSeedRun(w http.ResponseWriter, r *http.Request) { - startRun(s, w, r, "seed", s.runSeed) + startGuardedRun(s, w, r, "seed", s.runSeed, func(req SeedRequest, sess *Session) *productionRefusal { + if req.DryRun { + return nil + } + return s.guardProduction(sessionTarget(sess), req.ConfirmProduction, "seed it") + }) } func (s *Server) handleGapsRun(w http.ResponseWriter, r *http.Request) { - startRun(s, w, r, "gaps", s.runGaps) + startGuardedRun(s, w, r, "gaps", s.runGaps, func(req GapsRequest, sess *Session) *productionRefusal { + if !req.Fill || req.DryRun { + return nil + } + return s.guardProduction(sessionTarget(sess), req.ConfirmProduction, "fill its empty tables") + }) } func (s *Server) handleGenerateRun(w http.ResponseWriter, r *http.Request) { @@ -77,5 +109,19 @@ func (s *Server) handleExportRun(w http.ResponseWriter, r *http.Request) { } func (s *Server) handleCloneSchemaRun(w http.ResponseWriter, r *http.Request) { - startRun(s, w, r, "clone-schema", s.runCloneSchema) + startGuardedRun(s, w, r, "clone-schema", s.runCloneSchema, func(req CloneSchemaRequest, sess *Session) *productionRefusal { + if req.DryRun { + return nil + } + target := s.refTarget(ConnRef{ID: req.TargetID, SavedID: req.TargetSavedID}) + if req.TargetID == "" && req.TargetSavedID == "" { + target = connectionTarget{Info: req.Target} + if strings.TrimSpace(req.TargetDSN) != "" { + if _, _, info, err := buildRawDSN(req.Target.DBType, req.TargetDSN, req.Target.Params); err == nil { + target.Info = info + } + } + } + return s.guardProduction(target, req.ConfirmProduction, "clone a schema into it") + }) } diff --git a/internal/web/jobs.go b/internal/web/jobs.go index a2ec34a..d2fa227 100644 --- a/internal/web/jobs.go +++ b/internal/web/jobs.go @@ -6,8 +6,12 @@ import ( "encoding/hex" "fmt" "io" + "sort" "sync" "time" + + "github.com/AxeForging/seedstorm/internal/faultinject" + "github.com/AxeForging/seedstorm/internal/safego" ) // JobStatus represents the lifecycle state of a job. @@ -74,6 +78,8 @@ type Job struct { cancel context.CancelFunc closed bool closeCh chan struct{} + // Owner is the session that started the job, so its pages can find it again. + Owner string } // JobFunc is the body of a job. @@ -83,15 +89,26 @@ type JobFunc func(ctx context.Context, jc JobControl) (map[string]any, error) type Manager struct { mu sync.RWMutex jobs map[string]*Job + // keepFinished is how many finished jobs stay available (logs, results); + // older ones are evicted when a new job starts. + keepFinished int } +// defaultKeepFinished bounds the memory of finished jobs' logs. +const defaultKeepFinished = 50 + // NewManager constructs an empty job manager. func NewManager() *Manager { - return &Manager{jobs: make(map[string]*Job)} + return &Manager{jobs: make(map[string]*Job), keepFinished: defaultKeepFinished} } // Start registers a new job and runs fn in a goroutine. func (m *Manager) Start(ctx context.Context, name string, fn JobFunc) *Job { + return m.StartFor(ctx, "", name, fn) +} + +// StartFor is Start for a job owned by a session. +func (m *Manager) StartFor(ctx context.Context, owner, name string, fn JobFunc) *Job { jctx, cancel := context.WithCancel(ctx) job := &Job{ ID: newID(), @@ -101,15 +118,26 @@ func (m *Manager) Start(ctx context.Context, name string, fn JobFunc) *Job { subs: make(map[chan Event]struct{}), cancel: cancel, closeCh: make(chan struct{}), + Owner: owner, } m.mu.Lock() + m.evictLocked() m.jobs[job.ID] = job m.mu.Unlock() go func() { job.setStatus(JobRunning) ctrl := &jobWriter{job: job} - result, err := fn(jctx, ctrl) + // A panic ends this job as failed; the server and other jobs go on. + var result map[string]any + err := safego.Run("job "+name, func() error { + if err := faultinject.Hit(jctx, "job", name); err != nil { + return err + } + var ferr error + result, ferr = fn(jctx, ctrl) + return ferr + }) job.mu.Lock() job.EndedAt = time.Now() job.Result = result @@ -306,3 +334,54 @@ func newID() string { } return hex.EncodeToString(b[:]) } + +// evictLocked drops the oldest finished jobs beyond keepFinished. Running jobs +// are never evicted. +func (m *Manager) evictLocked() { + var finished []*Job + for _, j := range m.jobs { + select { + case <-j.closeCh: + finished = append(finished, j) + default: + } + } + if len(finished) <= m.keepFinished { + return + } + sort.Slice(finished, func(a, b int) bool { return finished[a].StartedAt.Before(finished[b].StartedAt) }) + for _, j := range finished[:len(finished)-m.keepFinished] { + delete(m.jobs, j.ID) + } +} + +// ForOwner returns the owner's jobs, newest first. +func (m *Manager) ForOwner(owner string) []*Job { + m.mu.RLock() + defer m.mu.RUnlock() + var out []*Job + for _, j := range m.jobs { + if j.Owner == owner { + out = append(out, j) + } + } + sort.Slice(out, func(a, b int) bool { return out[a].StartedAt.After(out[b].StartedAt) }) + return out +} + +// Position is the latest phase and progress event of a job. +func (j *Job) Position() (phase string, progress *Event) { + j.mu.Lock() + defer j.mu.Unlock() + for i := len(j.events) - 1; i >= 0; i-- { + ev := j.events[i] + if progress == nil && ev.Kind == EventProgress { + e := ev + progress = &e + } + if ev.Kind == EventPhase { + return ev.Text, progress + } + } + return "", progress +} diff --git a/internal/web/jobs_list_test.go b/internal/web/jobs_list_test.go new file mode 100644 index 0000000..575a096 --- /dev/null +++ b/internal/web/jobs_list_test.go @@ -0,0 +1,127 @@ +package web + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +// A page that is left and reopened finds the jobs its session started: running +// ones to reattach to, finished ones with their outcome. The server's boot id +// tells the page when the server restarted and older job ids are gone. +func TestJobsList_ShowsThisSessionsJobsAndTheBootID(t *testing.T) { + s, err := New(testOptions(t)) + if err != nil { + t.Fatal(err) + } + mine := &Session{ID: "sess-mine"} + other := &Session{ID: "sess-other"} + release := make(chan struct{}) + running := s.jobs.StartFor(context.Background(), mine.ID, "seed", func(ctx context.Context, jc JobControl) (map[string]any, error) { + jc.Phase("insert") + jc.Progress(40, 100, "users") + <-release + return nil, nil + }) + defer close(release) + s.jobs.StartFor(context.Background(), other.ID, "seed", func(ctx context.Context, jc JobControl) (map[string]any, error) { return nil, nil }) + + deadline := time.Now().Add(2 * time.Second) + for { + if evs := running.Events(); len(evs) >= 2 { + break + } + if time.Now().After(deadline) { + t.Fatal("job did not report progress") + } + time.Sleep(10 * time.Millisecond) + } + + w := httptest.NewRecorder() + s.writeJobList(w, mine.ID) + var body struct { + BootID string `json:"bootId"` + Jobs []struct { + ID string `json:"id"` + Status string `json:"status"` + Phase string `json:"phase"` + Progress struct { + Done int `json:"done"` + Total int `json:"total"` + Label string `json:"label"` + } `json:"progress"` + } `json:"jobs"` + } + if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil { + t.Fatalf("%v: %s", err, w.Body.String()) + } + if body.BootID == "" || body.BootID != s.bootID { + t.Fatalf("bootId = %q", body.BootID) + } + if len(body.Jobs) != 1 || body.Jobs[0].ID != running.ID { + t.Fatalf("jobs = %+v, want only this session's job", body.Jobs) + } + j := body.Jobs[0] + if j.Status != "running" || j.Phase != "insert" || j.Progress.Done != 40 || j.Progress.Total != 100 || j.Progress.Label != "users" { + t.Fatalf("job = %+v", j) + } +} + +// Finished jobs are evicted: a server running for days must not keep every +// job's log in memory. +func TestManager_EvictsOldFinishedJobs(t *testing.T) { + m := NewManager() + m.keepFinished = 3 + var ids []string + for i := 0; i < 6; i++ { + j := m.Start(context.Background(), "quick", func(ctx context.Context, jc JobControl) (map[string]any, error) { return nil, nil }) + <-j.Done() + ids = append(ids, j.ID) + } + blocking := make(chan struct{}) + live := m.Start(context.Background(), "long", func(ctx context.Context, jc JobControl) (map[string]any, error) { <-blocking; return nil, nil }) + defer close(blocking) + for _, id := range ids[:3] { + if _, ok := m.Get(id); ok { + t.Errorf("old finished job %s was kept", id) + } + } + for _, id := range ids[3:] { + if _, ok := m.Get(id); !ok { + t.Errorf("recent finished job %s was evicted", id) + } + } + if _, ok := m.Get(live.ID); !ok { + t.Fatal("a running job was evicted") + } +} + +// A quiet job still sends a keepalive, so the page can tell a slow step (the +// stream is alive) from a lost server. +func TestStreamJob_SendsKeepalivesWhileQuiet(t *testing.T) { + defer func(old time.Duration) { streamKeepalive = old }(streamKeepalive) + streamKeepalive = 20 * time.Millisecond + m := NewManager() + release := make(chan struct{}) + job := m.Start(context.Background(), "quiet", func(ctx context.Context, jc JobControl) (map[string]any, error) { + <-release + return nil, nil + }) + srv := &Server{jobs: m} + w := httptest.NewRecorder() + done := make(chan struct{}) + go func() { + srv.streamJob(w, httptest.NewRequest(http.MethodGet, "/api/jobs/"+job.ID+"/stream", nil), job) + close(done) + }() + time.Sleep(120 * time.Millisecond) + close(release) + <-done + if n := strings.Count(w.Body.String(), "event: ping\n"); n < 2 { + t.Fatalf("%d keepalives in:\n%s", n, w.Body.String()) + } +} diff --git a/internal/web/preflight.go b/internal/web/preflight.go new file mode 100644 index 0000000..ac83357 --- /dev/null +++ b/internal/web/preflight.go @@ -0,0 +1,81 @@ +package web + +import ( + "context" + "errors" + "sync" + "time" + + "github.com/rs/zerolog" + + "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" + "github.com/AxeForging/seedstorm/internal/seeder" +) + +// preflightTimeout bounds how long a side may take to answer before a two-sided +// run gives up on it. +var preflightTimeout = connectTimeout + +// connectBoth opens and pings the source and the target at the same time, so +// a dead side is reported within preflightTimeout and before anything is read. +// A source given as an imported snapshot is not connected (source is nil). +// Each side logs how it answered. +func (s *Server) connectBoth(ctx context.Context, log zerolog.Logger, srcRef ConnRef, snap *compare.Snapshot, tgtRef ConnRef) (source, target *Session, err error) { + var wg sync.WaitGroup + var srcErr, tgtErr error + side := func(role string, ref ConnRef, out **Session, errp *error) { + defer wg.Done() + *errp = safego.Run("connect "+role, func() error { + start := time.Now() + sess, err := s.resolveConnection(ref, role) + if err == nil { + err = pingSession(ctx, sess) + } + if err != nil { + log.Warn().Str("side", role).Err(err).Msg("Did not answer") + return err + } + log.Info().Str("side", role).Str("database", sessionLabel(sess)).Dur("answered_in", time.Since(start).Round(time.Millisecond)).Msg("Connected") + *out = sess + return nil + }) + *errp = runerr.OnSide(role, runerr.At(runerr.PhaseConnect, "", *errp)) + } + if snap == nil { + wg.Add(1) + go side(runerr.SideSource, srcRef, &source, &srcErr) + } + wg.Add(1) + go side(runerr.SideTarget, tgtRef, &target, &tgtErr) + wg.Wait() + if err := errors.Join(tgtErr, srcErr); err != nil { + return nil, nil, err + } + return source, target, nil +} + +// pingSession checks a live session still answers: a database that went away +// after it was connected must not hang the run on its first query. +func pingSession(ctx context.Context, sess *Session) error { + pctx, cancel := context.WithTimeout(ctx, preflightTimeout) + defer cancel() + return sess.Conn().PingContext(pctx) +} + +// snapshotOrSessionEndpoint builds the source endpoint after connectBoth. +func snapshotOrSessionEndpoint(ctx context.Context, sess *Session, snap *compare.Snapshot) (seeder.Endpoint, error) { + if snap != nil { + if len(snap.Tables) == 0 { + return seeder.Endpoint{}, errors.New("the imported counts have no tables") + } + label := snap.Label + if label == "" { + label = "imported counts" + } + return seeder.Endpoint{Snapshot: snap, Label: label, DBType: snap.DBType}, nil + } + ep, _, err := endpointFor(ctx, sess, false) + return ep, err +} diff --git a/internal/web/production.go b/internal/web/production.go new file mode 100644 index 0000000..06b2476 --- /dev/null +++ b/internal/web/production.go @@ -0,0 +1,119 @@ +package web + +import ( + "fmt" + "net/http" + "strings" +) + +// productionRefusal is a write to a production connection that the user has +// not confirmed by typing its label. +type productionRefusal struct { + Label string + Action string +} + +func (p *productionRefusal) Error() string { + return fmt.Sprintf("%s is marked production: type its label (%s) to %s", p.Label, p.Label, p.Action) +} + +// writeProductionRefusal answers 409 with what the page needs to ask for the +// label and retry. +func writeProductionRefusal(w http.ResponseWriter, p *productionRefusal) { + writeJSON(w, http.StatusConflict, map[string]string{"error": p.Error(), "code": "production_confirm", "label": p.Label}) +} + +// connectionTarget is what a run writes to, as far as it is known before the +// run opens anything. +type connectionTarget struct { + SavedID string + Info ConnectionInfo +} + +// productionConnection returns the saved production connection a target +// points at: the saved entry itself, or any saved production entry reaching +// the same database (a session opened ad hoc, or reused by DSN, still counts). +func (s *Server) productionConnection(t connectionTarget) (SavedConnection, bool) { + if s.store == nil { + return SavedConnection{}, false + } + saved, err := s.store.List() + if err != nil { + return SavedConnection{}, false + } + key := infoKey(t.Info) + for _, c := range saved { + if !c.Production { + continue + } + if t.SavedID != "" && c.ID == t.SavedID { + return c, true + } + if key != "" && savedInfoKey(c) == key { + return c, true + } + } + return SavedConnection{}, false +} + +// guardProduction refuses a write to a production target unless confirm is +// its label. +func (s *Server) guardProduction(t connectionTarget, confirm, action string) *productionRefusal { + c, ok := s.productionConnection(t) + if !ok || strings.TrimSpace(confirm) == c.Label { + return nil + } + return &productionRefusal{Label: c.Label, Action: action} +} + +// unflagRefusal refuses saving a production connection without the flag +// unless its label is typed: an older page or an API call that omits the +// field must not silently remove the protection. +func (s *Server) unflagRefusal(c SavedConnection, confirm string) *productionRefusal { + if s.store == nil || c.ID == "" || c.Production { + return nil + } + existing, ok, err := s.store.Get(c.ID) + if err != nil || !ok || !existing.Production || strings.TrimSpace(confirm) == existing.Label { + return nil + } + return &productionRefusal{Label: existing.Label, Action: "remove the production mark"} +} + +// sessionTarget is the target of a run on a live session. +func sessionTarget(sess *Session) connectionTarget { + if sess == nil { + return connectionTarget{} + } + return connectionTarget{SavedID: sess.SavedID, Info: sess.Info} +} + +// refTarget is the target of a run given by a ConnRef. +func (s *Server) refTarget(ref ConnRef) connectionTarget { + if ref.SavedID != "" { + return connectionTarget{SavedID: ref.SavedID} + } + if sess, ok := s.sessions.Get(ref.ID); ok { + return sessionTarget(sess) + } + return connectionTarget{} +} + +// infoKey identifies the database a connection reaches, like savedConnectionKey. +func infoKey(info ConnectionInfo) string { + if info.DBName == "" && info.Host == "" { + return "" + } + return savedConnectionKey(SavedConnection{DBType: info.DBType, Host: info.Host, Port: info.Port, DBName: info.DBName, User: info.User}) +} + +// savedInfoKey is infoKey for a saved connection, reading a raw DSN's parts. +func savedInfoKey(c SavedConnection) string { + if c.DSN != "" && c.Host == "" && c.DBName == "" { + if _, _, info, err := buildRawDSN(c.DBType, c.DSN, c.Params); err == nil { + return infoKey(info) + } + return "" + } + return savedConnectionKey(c) +} diff --git a/internal/web/production_test.go b/internal/web/production_test.go new file mode 100644 index 0000000..23cdc59 --- /dev/null +++ b/internal/web/production_test.go @@ -0,0 +1,134 @@ +package web + +import ( + "database/sql" + "encoding/json" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" +) + +// productionServer has a saved production connection and a live session +// connected to the same database (opened ad hoc, not from the saved entry). +func productionServer(t *testing.T) (*Server, *httptest.Server, *Session) { + t.Helper() + registerServeRunnerTestDriver() + prev := sqlOpen + sqlOpen = func(_, dsn string) (*sql.DB, error) { return sql.Open(serveRunnerTestDriverName, dsn) } + t.Cleanup(func() { sqlOpen = prev }) + + s, err := New(testOptions(t)) + if err != nil { + t.Fatal(err) + } + if _, err := s.store.Save(SavedConnection{Label: "billing-prod", DBType: "postgres", Host: "db.internal", Port: 5432, DBName: "billing", User: "svc", Production: true}, false); err != nil { + t.Fatal(err) + } + sess, err := s.sessions.OpenDSN("pgx", "live-billing", ConnectionInfo{DBType: "postgres", Host: "DB.internal", Port: 5432, DBName: "billing", User: "svc"}) + if err != nil { + t.Fatal(err) + } + srv := httptest.NewServer(s.Handler()) + t.Cleanup(srv.Close) + return s, srv, sess +} + +func postJSON(t *testing.T, srv *httptest.Server, sess *Session, path, body string) (int, map[string]any) { + t.Helper() + req, _ := http.NewRequest(http.MethodPost, srv.URL+path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Seedstorm-Request", "1") + if sess != nil { + req.AddCookie(&http.Cookie{Name: sessionCookieName, Value: sess.ID}) + } + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer res.Body.Close() + out := map[string]any{} + _ = json.NewDecoder(res.Body).Decode(&out) + return res.StatusCode, out +} + +// Writing to a connection marked production needs its label typed back. The +// refusal happens before a job exists, so nothing starts and nothing is written. +func TestProduction_WritesNeedTheTypedLabel(t *testing.T) { + _, srv, sess := productionServer(t) + s := srv.Config.Handler + + cases := []struct { + name, path, body string + want int + }{ + {"seed refused", "/api/seed", `{"rows":5}`, http.StatusConflict}, + {"seed with a wrong label", "/api/seed", `{"rows":5,"confirmProduction":"billing"}`, http.StatusConflict}, + {"seed dry run is not a write", "/api/seed", `{"rows":5,"dryRun":true}`, http.StatusAccepted}, + {"seed confirmed", "/api/seed", `{"rows":5,"confirmProduction":"billing-prod"}`, http.StatusAccepted}, + {"fill gaps refused", "/api/gaps", `{"rows":5,"fill":true}`, http.StatusConflict}, + {"scanning gaps is not a write", "/api/gaps", `{"rows":5}`, http.StatusAccepted}, + {"mirror into it refused", "/api/mirror", `{"source":{"id":"other"},"target":{"id":"` + sess.ID + `"}}`, http.StatusConflict}, + {"mirror dry run allowed", "/api/mirror", `{"source":{"id":"other"},"target":{"id":"` + sess.ID + `"},"dryRun":true}`, http.StatusAccepted}, + {"clone into it refused", "/api/clone-schema", `{"targetDsn":"postgres://svc@db.internal:5432/billing","target":{"dbType":"postgres","dbName":"billing","user":"svc","host":"db.internal","port":5432}}`, http.StatusConflict}, + } + _ = s + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + status, body := postJSON(t, srv, sess, c.path, c.body) + if status != c.want { + t.Fatalf("status = %d (%v), want %d", status, body, c.want) + } + if c.want == http.StatusConflict { + if body["code"] != "production_confirm" || body["label"] != "billing-prod" || !strings.Contains(body["error"].(string), "production") { + t.Fatalf("refusal body = %v", body) + } + } + }) + } +} + +// Saving a connection without the flag (an older page, a PUT that omits it) +// must not silently clear it; clearing it needs the label too. +func TestProduction_FlagIsNotClearedWithoutTheLabel(t *testing.T) { + s, srv, _ := productionServer(t) + list, _ := s.store.List() + id := list[0].ID + put := func(body string) int { + req, _ := http.NewRequest(http.MethodPut, srv.URL+"/api/saved-connections?id="+id, strings.NewReader(body)) + req.Header.Set("X-Seedstorm-Request", "1") + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + res.Body.Close() + return res.StatusCode + } + if code := put(`{"label":"billing-prod","dbType":"postgres","host":"db.internal","port":5432,"dbName":"billing","user":"svc"}`); code != http.StatusConflict { + t.Fatalf("unflag without label: status %d, want 409", code) + } + if got, _, _ := s.store.Get(id); !got.Production { + t.Fatal("production flag was cleared") + } + if code := put(`{"label":"billing-prod","dbType":"postgres","host":"db.internal","port":5432,"dbName":"billing","user":"svc","confirmLabel":"billing-prod"}`); code != http.StatusOK { + t.Fatalf("unflag with label: status %d", code) + } + if got, _, _ := s.store.Get(id); got.Production { + t.Fatal("production flag not cleared with the label") + } +} + +// The connect form carries the flag, so saving from it keeps it. +func TestProduction_ConnectFormSavesTheFlag(t *testing.T) { + s, srv := newConnectTestServer(t) + form := url.Values{"action": {"save"}, "label": {"orders-prod"}, "dbType": {"postgres"}, "host": {"db"}, "port": {"5432"}, "dbName": {"orders"}, "user": {"svc"}, "production": {"on"}} + res := postForm(t, srv, "/connect", form) + if res.StatusCode != http.StatusSeeOther { + t.Fatalf("status = %d", res.StatusCode) + } + list, _ := s.store.List() + if len(list) != 1 || !list[0].Production { + t.Fatalf("saved = %+v", list) + } +} diff --git a/internal/web/recover.go b/internal/web/recover.go new file mode 100644 index 0000000..3b4efcf --- /dev/null +++ b/internal/web/recover.go @@ -0,0 +1,23 @@ +package web + +import ( + "errors" + "net/http" + + "github.com/AxeForging/seedstorm/internal/safego" +) + +// recoverHandler answers a panicking request with a JSON 500 carrying an error +// id (the stack is logged under it) instead of dropping the connection. +func recoverHandler(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + err := safego.Run(r.Method+" "+r.URL.Path, func() error { + next.ServeHTTP(w, r) + return nil + }) + var p *safego.PanicError + if errors.As(err, &p) { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": p.Error(), "errorId": p.ID}) + } + }) +} diff --git a/internal/web/recover_test.go b/internal/web/recover_test.go new file mode 100644 index 0000000..ebf8f42 --- /dev/null +++ b/internal/web/recover_test.go @@ -0,0 +1,74 @@ +package web + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/safego" +) + +// A panicking job used to crash `serve`: every session and every other job +// went with it. It must end as a failed job while other jobs keep running. +func TestManager_PanickingJobFailsAloneAndOthersFinish(t *testing.T) { + m := NewManager() + release := make(chan struct{}) + other := m.Start(context.Background(), "steady", func(ctx context.Context, jc JobControl) (map[string]any, error) { + <-release + return map[string]any{"ok": true}, nil + }) + boom := m.Start(context.Background(), "boom", func(ctx context.Context, jc JobControl) (map[string]any, error) { + jc.Phase("write") + var list []int + i := 3 + _ = list[i] + return nil, nil + }) + select { + case <-boom.Done(): + case <-time.After(2 * time.Second): + t.Fatal("panicking job never ended") + } + st := boom.State() + var p *safego.PanicError + if st.Status != JobFailed || !errors.As(st.Err, &p) { + t.Fatalf("boom state = %+v, want failed with a recovered panic", st) + } + close(release) + select { + case <-other.Done(): + case <-time.After(2 * time.Second): + t.Fatal("the other job did not finish") + } + if st := other.State(); st.Status != JobDone { + t.Fatalf("other job = %+v", st) + } +} + +// A handler panic returns a JSON error the page can show, with an id to find +// the stack in the server log, instead of a dropped connection. +func TestRecoverHandler_AnswersJSON500(t *testing.T) { + h := recoverHandler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + panic("handler exploded") + })) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/graph", nil)) + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status = %d", rec.Code) + } + var body map[string]string + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not JSON: %q", rec.Body.String()) + } + if !strings.Contains(body["error"], "internal error") || len(body["errorId"]) != 8 { + t.Fatalf("body = %v", body) + } + if strings.Contains(rec.Body.String(), "goroutine") { + t.Fatal("the stack leaked into the response") + } +} diff --git a/internal/web/runners.go b/internal/web/runners.go index ae942e4..e109cec 100644 --- a/internal/web/runners.go +++ b/internal/web/runners.go @@ -5,7 +5,6 @@ import ( "database/sql" "fmt" "io" - "runtime" "strings" "time" @@ -15,8 +14,10 @@ import ( "github.com/AxeForging/seedstorm/internal/faker" "github.com/AxeForging/seedstorm/internal/graph" "github.com/AxeForging/seedstorm/internal/rules" + "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/schema" "github.com/AxeForging/seedstorm/internal/seeder" + "github.com/AxeForging/seedstorm/internal/tuning" "github.com/goccy/go-yaml" "github.com/rs/zerolog" ) @@ -58,6 +59,9 @@ type SeedRequest struct { Workers int `json:"workers,omitempty"` // GenWorkers is how many tables generate at once (0 or 1: one). GenWorkers int `json:"genWorkers,omitempty"` + // ConfirmProduction is the label of a production connection, typed to + // write to it. + ConfirmProduction string `json:"confirmProduction,omitempty"` } type CloneSchemaRequest struct { @@ -68,6 +72,9 @@ type CloneSchemaRequest struct { Password string `json:"password,omitempty"` DropExisting bool `json:"dropExisting"` DryRun bool `json:"dryRun"` + // ConfirmProduction is the target's label, typed to write to a + // production connection. + ConfirmProduction string `json:"confirmProduction,omitempty"` // Objects adds views, routines and triggers to the tables clone. Objects db.CloneObjects `json:"objects"` } @@ -191,6 +198,9 @@ func (s *Server) resolveCloneTarget(req CloneSchemaRequest, source *Session) (*S func (s *Server) runSeed(ctx context.Context, sess *Session, req SeedRequest, jc JobControl) (map[string]any, error) { log := jobLogger(jc) + if !req.DryRun { + defer sess.InvalidateCounts() + } tableRows := cleanTableRows(req.TableRows) truncateOnly := req.Truncate && req.Rows == 0 && req.EnumRows == 0 && len(tableRows) == 0 && req.ProfileID == "" if req.Rows < 0 || (req.Rows == 0 && !truncateOnly) { @@ -260,6 +270,10 @@ func (s *Server) runSeed(ctx context.Context, sess *Session, req SeedRequest, jc return nil, err } log.Info().Str("order", strings.Join(targetTables, " → ")).Msg("Seed order resolved") + // Refuse tables that cannot be generated before anything is truncated. + if err := faker.CheckSeedable(sc, targetTables, profile.overrides); err != nil { + return nil, err + } if req.Truncate && !req.DryRun { jc.Phase("truncate") @@ -270,7 +284,7 @@ func (s *Server) runSeed(ctx context.Context, sess *Session, req SeedRequest, jc } jc.Progress(done, total, table) }); err != nil { - return nil, fmt.Errorf("truncate: %w", err) + return nil, runerr.At(runerr.PhaseTruncate, "", err) } log.Info().Msg("Truncate complete") } @@ -330,6 +344,7 @@ func (s *Server) runSeed(ctx context.Context, sess *Session, req SeedRequest, jc Overrides: overrides, OnWarning: collectWarning(&warnings, log), }, + OnNotice: func(msg string) { log.Warn().Msg(msg) }, OnTableStart: func(table string) error { log.Info().Str("table", table).Msg("Seeding table"); return nil }, OnRows: func(table string, rows []map[string]interface{}) error { if req.DryRun { @@ -339,7 +354,8 @@ func (s *Server) runSeed(ctx context.Context, sess *Session, req SeedRequest, jc }, }) if err != nil { - return nil, err + finishProgress() + return partialSeedResult(res, targetTables, err), err } finishProgress() totalRows := res.Total @@ -388,10 +404,16 @@ type GapsRequest struct { ProfileID string `json:"profileId,omitempty"` Workers int `json:"workers,omitempty"` GenWorkers int `json:"genWorkers,omitempty"` + // ConfirmProduction is the label of a production connection, typed to + // write to it. + ConfirmProduction string `json:"confirmProduction,omitempty"` } func (s *Server) runGaps(ctx context.Context, sess *Session, req GapsRequest, jc JobControl) (map[string]any, error) { log := jobLogger(jc) + if req.Fill && !req.DryRun { + defer sess.InvalidateCounts() + } if req.Rows <= 0 { req.Rows = 100 } @@ -420,24 +442,26 @@ func (s *Server) runGaps(ctx context.Context, sess *Session, req GapsRequest, jc } jc.Phase("scan") log.Info().Int("tables", len(allSorted)).Msg("Scanning row counts") - counts, err := db.GetTableRowCounts(ctx, conn, sess.DBType, allSorted) - if err != nil { + counts, failed := db.CountTables(ctx, conn, sess.DBType, allSorted, func(done, total int, table string) { + jc.Progress(done, total, "count "+table) + }) + if err := ctx.Err(); err != nil { return nil, err } - - // Default gap set: every empty table, in topological order. - var gapTables []string for _, t := range allSorted { - if counts[t] == 0 { - gapTables = append(gapTables, t) + if ferr, ok := failed[t]; ok { + log.Warn().Str("table", t).Err(ferr).Msg("Row count failed: left out of the gaps (unknown is not empty)") } } + + // Default gap set: every table known to be empty, in topological order. + gapTables := seeder.GapTables(allSorted, counts, nil) // If the caller scoped the fill, intersect with empty tables and resolve // non-nullable parents (which may themselves be empty). if len(req.Tables) > 0 { selected := make(map[string]bool, len(req.Tables)) for _, t := range req.Tables { - if counts[t] == 0 { + if seeder.KnownEmpty(counts, t) { selected[t] = true } } @@ -447,7 +471,7 @@ func (s *Server) runGaps(ctx context.Context, sess *Session, req GapsRequest, jc // parents do not need re-seeding. gapTables = gapTables[:0] for _, t := range resolved { - if counts[t] == 0 { + if seeder.KnownEmpty(counts, t) { gapTables = append(gapTables, t) } } @@ -487,10 +511,18 @@ func (s *Server) runGaps(ctx context.Context, sess *Session, req GapsRequest, jc Overrides: profile.overrides, OnWarning: collectWarning(&warnings, log), }, + OnNotice: func(msg string) { log.Warn().Msg(msg) }, OnTableStart: func(table string) error { log.Info().Str("table", table).Msg("Filling table"); return nil }, }) if err != nil { - return nil, err + finishProgress() + partial := partialSeedResult(res, gapTables, err) + for k, v := range result { + if _, set := partial[k]; !set { + partial[k] = v + } + } + return partial, err } finishProgress() totalRows := res.Total @@ -719,9 +751,10 @@ func requestWorkers(n int) int { // maxWorkers caps connections one web run may open against a database. const maxWorkers = 32 -// requestGenWorkers bounds a requested generator count to the machine's cores. +// requestGenWorkers bounds a requested generator count to the cores this +// process may use (a container CPU quota, not the host's cores). func requestGenWorkers(n int) int { - return max(1, min(n, runtime.NumCPU())) + return tuning.ClampGenerators(n) } func generationWarningsView(warnings []faker.GenerationWarning) []map[string]any { diff --git a/internal/web/runners_test.go b/internal/web/runners_test.go index 1e8aefe..2c62f3a 100644 --- a/internal/web/runners_test.go +++ b/internal/web/runners_test.go @@ -12,7 +12,9 @@ import ( "reflect" "strings" "sync" + "sync/atomic" "testing" + "time" "github.com/AxeForging/seedstorm/internal/schema" ) @@ -303,6 +305,9 @@ func containsAll(value string, parts ...string) bool { const serveRunnerTestDriverName = "seedstorm_web_runner_test" +// countedQueries counts queries run on a "counted" connection. +var countedQueries atomic.Int64 + var registerServeRunnerDriverOnce sync.Once func registerServeRunnerTestDriver() { @@ -328,16 +333,42 @@ func (c *serveRunnerTestConn) Prepare(string) (driver.Stmt, error) { func (c *serveRunnerTestConn) Close() error { return nil } func (c *serveRunnerTestConn) Begin() (driver.Tx, error) { + if c.name == "counted" { + return noopTx{}, nil + } return nil, errors.New("transactions not implemented") } -func (c *serveRunnerTestConn) Ping(context.Context) error { return nil } +func (c *serveRunnerTestConn) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { + return c.Begin() +} + +// noopTx lets reads run inside the read-only transactions seedstorm opens. +type noopTx struct{} + +func (noopTx) Commit() error { return nil } +func (noopTx) Rollback() error { return nil } + +func (c *serveRunnerTestConn) Ping(ctx context.Context) error { + if c.name == "down" { + // A database that stopped answering: the ping waits until its deadline. + <-ctx.Done() + return ctx.Err() + } + return nil +} func (c *serveRunnerTestConn) Query(query string, args []driver.Value) (driver.Rows, error) { return &serveRunnerRows{columns: []string{"id"}}, nil } -func (c *serveRunnerTestConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) { +func (c *serveRunnerTestConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) { + if c.name == "counted" { + countedQueries.Add(1) + if strings.HasPrefix(query, "SELECT COUNT(*)") { + return &oneValueRows{value: int64(3)}, nil + } + } return &serveRunnerRows{columns: []string{"id"}}, nil } @@ -345,6 +376,9 @@ func (c *serveRunnerTestConn) ExecContext(_ context.Context, query string, _ []d if c.name == "stale" && strings.Contains(query, "INSERT") { return nil, errors.New("cache lookup failed for type 34868 (SQLSTATE XX000)") } + if c.name == "fail-orders" && strings.Contains(query, "INSERT") && strings.Contains(query, "orders") { + return nil, errors.New(`insert or update on table "orders" violates foreign key constraint`) + } if c.name == "truncate-only" && strings.Contains(query, "INSERT") { return nil, errors.New("truncate-only run should not insert rows") } @@ -521,3 +555,92 @@ func TestStartRunRejectsOversizedRequestBodies(t *testing.T) { t.Fatalf("status %d, want 413", res.StatusCode) } } + +// A run that fails part-way reports what it wrote before the failure, where it +// failed, and what to do next, instead of only an error. +func TestRunSeedFailureReturnsWhatWasWrittenAndWhereItStopped(t *testing.T) { + registerServeRunnerTestDriver() + oldOpen := sqlOpen + sqlOpen = func(_, dsn string) (*sql.DB, error) { return sql.Open(serveRunnerTestDriverName, dsn) } + defer func() { sqlOpen = oldOpen }() + srv, err := New(testOptions(t)) + if err != nil { + t.Fatal(err) + } + conn, _ := sql.Open(serveRunnerTestDriverName, "fail-orders") + defer conn.Close() + sess := &Session{DBType: "mysql", DSN: "fail-orders", conn: conn, schema: runnerRowCountSchema()} + + result, err := srv.runSeed(context.Background(), sess, SeedRequest{Rows: 7, BatchSize: 100, Workers: 1}, testJobControl{}) + if err == nil { + t.Fatal("runSeed succeeded although orders inserts fail") + } + if result == nil || result["partial"] != true { + t.Fatalf("result = %#v, want a partial result", result) + } + failure, _ := result["failure"].(map[string]any) + if failure["phase"] != "write" || failure["table"] != "orders" { + t.Fatalf("failure = %#v, want write · orders", failure) + } + counts, _ := result["tableCounts"].(map[string]int) + if counts["users"] != 7 || counts["orders"] != 0 { + t.Fatalf("tableCounts = %#v, want users 7 and orders 0", counts) + } + if next, _ := result["nextStep"].(string); !strings.Contains(next, "Fill empty") { + t.Fatalf("nextStep = %q", next) + } +} + +// A compare whose target stopped answering fails fast, before reading the +// source, and says it was the target: it used to count the whole source first +// and then wait on the dead target with no feedback. +func TestRunCompare_UnreachableTargetFailsFirstAndNamesTheSide(t *testing.T) { + registerServeRunnerTestDriver() + defer func(old time.Duration) { preflightTimeout = old }(preflightTimeout) + preflightTimeout = 300 * time.Millisecond + srv, err := New(testOptions(t)) + if err != nil { + t.Fatal(err) + } + src, _ := sql.Open(serveRunnerTestDriverName, "counted") + dead, _ := sql.Open(serveRunnerTestDriverName, "down") + defer src.Close() + defer dead.Close() + source := &Session{ID: "src", DBType: "pgx", conn: src, schema: runnerRowCountSchema(), Info: ConnectionInfo{DBName: "app"}} + target := &Session{ID: "tgt", DBType: "pgx", conn: dead, schema: runnerRowCountSchema(), Info: ConnectionInfo{DBName: "app_copy"}} + srv.sessions.add(source) + srv.sessions.add(target) + countedQueries.Store(0) + + start := time.Now() + _, err = srv.runCompare(context.Background(), source, CompareRequest{Source: ConnRef{ID: "src"}, Target: ConnRef{ID: "tgt"}}, testJobControl{}) + if err == nil { + t.Fatal("compare succeeded with a dead target") + } + if !strings.HasPrefix(err.Error(), "target · connect:") { + t.Fatalf("err = %q, want it to start with target · connect:", err) + } + if d := time.Since(start); d > 2*time.Second { + t.Fatalf("took %s to notice the dead target", d) + } + if n := countedQueries.Load(); n != 0 { + t.Fatalf("the source ran %d queries before the dead target was noticed", n) + } +} + +// oneValueRows is a single-row, single-column result. +type oneValueRows struct { + value driver.Value + read bool +} + +func (r *oneValueRows) Columns() []string { return []string{"n"} } +func (r *oneValueRows) Close() error { return nil } +func (r *oneValueRows) Next(dest []driver.Value) error { + if r.read { + return io.EOF + } + r.read = true + dest[0] = r.value + return nil +} diff --git a/internal/web/server.go b/internal/web/server.go index e0b937b..1bc0859 100644 --- a/internal/web/server.go +++ b/internal/web/server.go @@ -29,6 +29,9 @@ type Server struct { jobs *Manager store *ConnectionStore profiles *profiles.Store + // bootID changes every time the server starts: a page holding job ids + // from an earlier run learns they are gone. + bootID string } // Options configures the Server. @@ -69,6 +72,7 @@ func New(opts Options) (*Server, error) { jobs: NewManager(), store: NewConnectionStore(path), profiles: profiles.NewStore(profilesPath), + bootID: newID()[:12], } s.routes() return s, nil @@ -78,11 +82,11 @@ func New(opts Options) (*Server, error) { func (s *Server) Addr() string { return s.addr } // Handler returns the underlying http.Handler. -func (s *Server) Handler() http.Handler { return s.mux } +func (s *Server) Handler() http.Handler { return recoverHandler(s.mux) } // ListenAndServe starts the HTTP server. func (s *Server) ListenAndServe(ctx context.Context) error { - srv := &http.Server{Addr: s.addr, Handler: s.mux} + srv := &http.Server{Addr: s.addr, Handler: s.Handler()} go func() { <-ctx.Done() shutdownCtx, cancel := context.WithCancel(context.Background()) @@ -121,6 +125,7 @@ func (s *Server) routes() { s.mux.HandleFunc("/api/profiles/ignored", s.handleProfileIgnored) s.mux.HandleFunc("/api/schema", s.handleSchemaJSON) s.mux.HandleFunc("/api/table", s.handleTablePreviewJSON) + s.mux.HandleFunc("/api/jobs", s.handleJobList) s.mux.HandleFunc("/api/jobs/", s.handleJobsAPI) s.mux.HandleFunc("/api/seed", s.handleSeedRun) s.mux.HandleFunc("/api/gaps", s.handleGapsRun) @@ -185,6 +190,9 @@ func templateFuncs() template.FuncMap { } return b }, + // connKey identifies what a session points at, so per-connection page + // state (remembered form values) survives reconnecting. + "connKey": sessionConnectionKey, "connName": func(info ConnectionInfo) string { if info.Label != "" { return info.Label diff --git a/internal/web/session.go b/internal/web/session.go index 0c37289..16a68db 100644 --- a/internal/web/session.go +++ b/internal/web/session.go @@ -42,6 +42,8 @@ type ConnectionInfo struct { // Session holds a live database connection plus the cached schema introspected // from it. The DSN (including the password) is intentionally not retained. type Session struct { + // SavedID is the saved connection this session was opened from, if any. + SavedID string ID string Info ConnectionInfo DBType string // driver name: "pgx" or "mysql" @@ -51,6 +53,15 @@ type Session struct { schema *schema.Schema cachedAt time.Time createdAt time.Time + // loading is the introspection in flight: callers wait for it instead of + // starting another, and the session lock is never held across it. + loading *schemaLoad + + // Row counts of the workspace, cached until a run writes or the user + // refreshes; countsLoading is the count in flight. + counts map[string]int64 + countsAt time.Time + countsLoading *countsLoad accessMu sync.Mutex access *accessView @@ -111,10 +122,14 @@ func (r *SessionRegistry) open(driver, dsn string, info ConnectionInfo) (*Sessio conn: conn, createdAt: time.Now(), } + r.add(s) + return s, nil +} + +func (r *SessionRegistry) add(s *Session) { r.mu.Lock() r.sessions[s.ID] = s r.mu.Unlock() - return s, nil } func (r *SessionRegistry) findByDSN(driver, dsn string) *Session { @@ -225,27 +240,59 @@ func (s *Session) OpenRunConn(ctx context.Context) (*sql.DB, error) { return conn, nil } -// Schema returns the cached schema, introspecting if needed. +// schemaLoadTimeout bounds one introspection of a session's database. +var schemaLoadTimeout = 5 * time.Minute + +type schemaLoad struct { + done chan struct{} + sc *schema.Schema + err error +} + +// Schema returns the cached schema, introspecting if needed. Concurrent calls +// share one introspection, run on the session's own connection, without +// holding the session lock while the database answers. func (s *Session) Schema(force bool) (*schema.Schema, error) { s.mu.Lock() - defer s.mu.Unlock() if !force && s.schema != nil { - return s.schema, nil + sc := s.schema + s.mu.Unlock() + return sc, nil } - tables, err := db.Introspect(s.DBType, s.DSN) - if err != nil { - return nil, err + if l := s.loading; l != nil { + s.mu.Unlock() + <-l.done + return l.sc, l.err } - out := faker.BuildSchema(s.DBType, tables) - s.schema = out - s.cachedAt = time.Now() - return out, nil + l := &schemaLoad{done: make(chan struct{})} + s.loading = l + s.mu.Unlock() + + ctx, cancel := context.WithTimeout(context.Background(), schemaLoadTimeout) + tables, err := db.IntrospectConn(ctx, s.conn, s.DBType, nil) + cancel() + if err == nil { + l.sc = faker.BuildSchema(s.DBType, tables) + } + l.err = err + + s.mu.Lock() + s.loading = nil + if err == nil { + s.schema = l.sc + s.cachedAt = time.Now() + } + s.mu.Unlock() + close(l.done) + return l.sc, l.err } -// RawTables returns the raw introspected tables (with constraint metadata -// such as enum values, CHECK ranges, etc.) — re-introspects on each call. +// RawTables introspects the session's database again (constraint metadata such +// as enum values and CHECK ranges), on its own connection. func (s *Session) RawTables() ([]db.Table, error) { - return db.Introspect(s.DBType, s.DSN) + ctx, cancel := context.WithTimeout(context.Background(), schemaLoadTimeout) + defer cancel() + return db.IntrospectConn(ctx, s.conn, s.DBType, nil) } // SetSchema overrides the cached schema (used by upload/paste flows). @@ -297,3 +344,60 @@ func newSessionID() string { } return hex.EncodeToString(b[:]) } + +type countsLoad struct { + done chan struct{} + counts map[string]int64 + at time.Time +} + +// countConcurrency is how many tables a workspace counts at once. +const countConcurrency = 2 + +// Counts returns the tables' row counts: the cached ones unless force is set +// or they were invalidated. Concurrent calls share one count. A table whose +// count failed is missing from the map (unknown, never 0). +func (s *Session) Counts(ctx context.Context, tables []string, force bool) (map[string]int64, time.Time) { + s.mu.Lock() + if !force && s.counts != nil { + counts, at := s.counts, s.countsAt + s.mu.Unlock() + return counts, at + } + if l := s.countsLoading; l != nil { + s.mu.Unlock() + <-l.done + return l.counts, l.at + } + l := &countsLoad{done: make(chan struct{})} + s.countsLoading = l + s.mu.Unlock() + + lim := db.DefaultCountLimits + lim.Concurrency = countConcurrency + counts, _ := db.CountTablesWithin(ctx, s.conn, s.DBType, tables, lim, nil) + l.counts, l.at = counts, time.Now().UTC() + + s.mu.Lock() + s.countsLoading = nil + if ctx.Err() == nil { + s.counts, s.countsAt = l.counts, l.at + } + s.mu.Unlock() + close(l.done) + return l.counts, l.at +} + +// CachedCounts returns the cached counts without counting (nil when none). +func (s *Session) CachedCounts() (map[string]int64, time.Time) { + s.mu.Lock() + defer s.mu.Unlock() + return s.counts, s.countsAt +} + +// InvalidateCounts drops the cached counts: a run wrote to the database. +func (s *Session) InvalidateCounts() { + s.mu.Lock() + defer s.mu.Unlock() + s.counts, s.countsAt = nil, time.Time{} +} diff --git a/internal/web/session_test.go b/internal/web/session_test.go index 09c42ce..9c0d95f 100644 --- a/internal/web/session_test.go +++ b/internal/web/session_test.go @@ -1,6 +1,8 @@ package web import ( + "database/sql" + "sync" "testing" "time" ) @@ -51,3 +53,41 @@ func TestSessionRegistryOpenReleasesIdleDatabaseConnections(t *testing.T) { t.Fatalf("the session must reconnect on use: %v", err) } } + +// Pages load the schema at the same time (graph, counts, access). They share +// one introspection instead of each running one against the database. +func TestSession_ConcurrentSchemaCallsShareOneIntrospection(t *testing.T) { + registerServeRunnerTestDriver() + conn, err := sql.Open(serveRunnerTestDriverName, "counted") + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + countedQueries.Store(0) + one := &Session{DBType: "pgx", conn: conn} + if _, err := one.Schema(false); err != nil { + t.Fatal(err) + } + perIntrospection := countedQueries.Load() + if perIntrospection == 0 { + t.Fatal("introspection ran no queries on the session connection") + } + + countedQueries.Store(0) + shared := &Session{DBType: "pgx", conn: conn} + var wg sync.WaitGroup + for i := 0; i < 12; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := shared.Schema(false); err != nil { + t.Error(err) + } + }() + } + wg.Wait() + if got := countedQueries.Load(); got != perIntrospection { + t.Fatalf("12 concurrent calls ran %d queries, want one introspection's %d", got, perIntrospection) + } +} diff --git a/internal/web/static/app.js b/internal/web/static/app.js index 040ef7d..72be29a 100644 --- a/internal/web/static/app.js +++ b/internal/web/static/app.js @@ -530,6 +530,269 @@ }); } + // ── starting runs (production confirmation) ──────────────────────────── + // confirmProduction asks the user to type a production connection's label. + // It resolves to the typed label, or null when they cancel. + function confirmProduction(label, message) { + return new Promise((resolve) => { + const dlg = document.createElement("dialog"); + dlg.className = "prod-dialog"; + dlg.dataset.testid = "prod-confirm-dialog"; + dlg.setAttribute("data-testid", "prod-confirm-dialog"); + dlg.innerHTML = ` +
+

Write to a production database?

+

${escapeHTML(message || "")}

+ +
+ + +
+
`; + document.body.appendChild(dlg); + const input = dlg.querySelector("input"); + const ok = dlg.querySelector('[value="ok"]'); + input.addEventListener("input", () => { ok.disabled = input.value !== label; }); + dlg.addEventListener("close", () => { + const typed = dlg.returnValue === "ok" && input.value === label ? input.value : null; + dlg.remove(); + resolve(typed); + }); + dlg.showModal(); + input.focus(); + }); + } + + // postRun starts a job. A refusal to write to a production connection asks + // for its label and retries; a request that never reaches the server, or a + // non-JSON answer, rejects with a readable message instead of hanging. + async function postRun(endpoint, body) { + let res; + try { + res = await fetch(endpoint, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify(body) }); + } catch (err) { + throw new Error(`seedstorm did not answer (${err.message || err}). Is the server still running?`); + } + const j = await res.json().catch(() => ({ error: `unexpected answer (${res.status} ${res.statusText})` })); + if (res.status === 409 && j.code === "production_confirm") { + const typed = await confirmProduction(j.label, j.error); + if (typed == null) { + const err = new Error(`Not written: ${j.label} is a production database and the write was not confirmed.`); + err.cancelled = true; + throw err; + } + return postRun(endpoint, { ...body, confirmProduction: typed }); + } + if (!res.ok) throw new Error(j.error || res.statusText); + return j; + } + + // ── browser storage (per-viewer conveniences only) ───────────────────── + // Storage can be unavailable (private windows, blocked site data): every + // read and write falls back to nothing and the page keeps working. + function readStore(key, fallback) { + try { return JSON.parse(localStorage.getItem(key) || "null") ?? fallback; } catch (_) { return fallback; } + } + function writeStore(key, value) { + try { localStorage.setItem(key, JSON.stringify(value)); return true; } catch (_) { return false; } + } + + // ── run strip: what the current job is doing, next to the action ──────── + // Every page with jobs has one or more [data-run-strip] elements. The strip + // shows the job, its phase and progress, elapsed time, and the connection + // state: running (events arriving), quiet (no update for a while but the + // server still pings), reconnecting, or lost. A failure shows where it + // happened and what to do next. + const QUIET_AFTER_MS = 30000; + const LOST_AFTER_MS = 45000; + const run = { + jobId: "", name: "", status: "idle", phase: "", done: 0, total: 0, label: "", + startedAt: 0, endedAt: 0, lastEventAt: 0, lastPingAt: 0, conn: "running", + failure: null, nextStep: "", message: "", timer: null, + }; + function stripElements() { + return document.querySelectorAll("[data-run-strip]"); + } + function formatElapsed(ms) { + const s = Math.max(0, ms) / 1000; + if (s < 60) return s.toFixed(1) + "s"; + const m = Math.floor(s / 60); + return m + "m " + String(Math.floor(s % 60)).padStart(2, "0") + "s"; + } + function failureText(f, fallback) { + if (!f) return fallback || ""; + const where = [f.side, f.phase, f.table].filter(Boolean).join(" · "); + return where ? `${where}: ${fallback || f.message || ""}` : (fallback || f.message || ""); + } + function renderRunStrip() { + const strips = stripElements(); + if (!strips.length || !run.jobId) return; + const now = Date.now(); + const active = run.status === "running" || run.status === "pending"; + const elapsed = formatElapsed((run.endedAt || now) - run.startedAt); + const pct = run.total > 0 ? Math.min(100, (run.done / run.total) * 100) : null; + let state = run.status; + if (active) state = run.conn; + let detail = run.label || ""; + if (active && run.conn === "quiet") { + detail = `No update for ${formatElapsed(now - run.lastEventAt)} — last: ${[run.phase, run.label].filter(Boolean).join(" · ") || "starting"}. Still connected.`; + } else if (active && run.conn === "reconnecting") { + detail = "Connection to seedstorm dropped; reconnecting…"; + } else if (run.conn === "lost") { + detail = run.message || "Lost contact with seedstorm. Is the server still running?"; + } + const failed = run.status === "failed" || run.status === "canceled"; + for (const el of strips) { + el.hidden = false; + el.dataset.state = state; + el.innerHTML = ` +
+ + ${escapeHTML(stateLabel(state))} + ${escapeHTML(run.name)}${run.phase ? " · " + escapeHTML(run.phase) : ""} + ${elapsed} +
+ ${pct != null && !failed ? ` + ${run.done.toLocaleString()} / ${run.total.toLocaleString()} · ${pct.toFixed(pct < 10 ? 1 : 0)}%` : ""} + ${failed ? `` : ""} + ${failed && run.nextStep ? `

${escapeHTML(run.nextStep)}

` : ""} + ${!failed && detail ? `

${escapeHTML(detail)}

` : ""}`; + } + } + function stateLabel(state) { + return { + running: "Running", pending: "Starting", quiet: "Still working", reconnecting: "Reconnecting", + lost: "Connection lost", done: "Done", failed: "Failed", canceled: "Cancelled", + }[state] || state; + } + function tickRunStrip() { + const active = run.status === "running" || run.status === "pending"; + if (active) { + const now = Date.now(); + if (run.conn !== "reconnecting" && run.conn !== "lost") { + run.conn = now - run.lastEventAt > QUIET_AFTER_MS ? "quiet" : "running"; + } + if (now - Math.max(run.lastPingAt, run.lastEventAt) > LOST_AFTER_MS && run.conn !== "lost") { + run.conn = "lost"; + checkServerAfterLoss(); + } + } + renderRunStrip(); + if (!active && run.timer) { clearInterval(run.timer); run.timer = null; } + } + async function checkServerAfterLoss() { + const stored = readStore(LAST_JOB_KEY, null); + try { + const list = await fetchJSON("/api/jobs", { timeoutMs: 5000 }); + if (stored && list.bootId !== stored.bootId) { + run.status = "failed"; + run.message = `seedstorm restarted: the ${run.name} job was lost. Nothing more will be written by it.`; + } else { + run.conn = "reconnecting"; + } + } catch (_) { + run.message = "Lost contact with seedstorm. Is the server still running?"; + } + renderRunStrip(); + } + function beginRun(jobId, name, bootId) { + Object.assign(run, { + jobId, name, status: "running", phase: "", done: 0, total: 0, label: "", failure: null, nextStep: "", message: "", + startedAt: Date.now(), endedAt: 0, lastEventAt: Date.now(), lastPingAt: Date.now(), conn: "running", + }); + if (bootId) writeStore(LAST_JOB_KEY, { bootId, jobId, name, page: document.body.dataset.active || "" }); + if (!run.timer) run.timer = setInterval(tickRunStrip, 1000); + renderRunStrip(); + } + function endRun(job) { + const stored = readStore(LAST_JOB_KEY, null); + if (stored && stored.jobId === job.id) writeStore(LAST_JOB_KEY, { ...stored, seen: true }); + run.status = job.status || "failed"; + run.endedAt = job.end ? Date.parse(job.end) || Date.now() : Date.now(); + const result = job.result || {}; + run.failure = result.failure || null; + run.nextStep = result.nextStep || ""; + run.message = job.error || ""; + if (run.conn !== "lost") run.conn = "running"; + renderRunStrip(); + } + const LAST_JOB_KEY = "seedstorm.lastJob.v1"; + + // fetchJSON is the one way page scripts read the API: it times out, reads a + // non-JSON answer as an error, and rejects with a message a person can act on. + async function fetchJSON(url, opts = {}) { + const ctrl = new AbortController(); + const timer = setTimeout(() => ctrl.abort(), opts.timeoutMs || 30000); + let res; + try { + res = await fetch(url, { cache: "no-store", ...opts, signal: ctrl.signal }); + } catch (err) { + throw new Error(err.name === "AbortError" ? `seedstorm did not answer ${url} in time` : `seedstorm did not answer (${err.message || err})`); + } finally { + clearTimeout(timer); + } + const body = await res.json().catch(() => null); + if (!res.ok) throw new Error((body && body.error) || `${res.status} ${res.statusText}`); + if (body == null) throw new Error(`unexpected answer from ${url}`); + return body; + } + + // resumeRun reattaches to this session's running job after the page was + // left and opened again, or says the server restarted and the job is gone. + async function resumeRun(names, hooks) { + if (!stripElements().length) return; + const stored = readStore(LAST_JOB_KEY, null); + let list; + try { + list = await fetchJSON("/api/jobs", { timeoutMs: 5000 }); + } catch (_) { return; } + const running = (list.jobs || []).find((j) => (j.status === "running" || j.status === "pending") && names.includes(j.name)); + if (running) { + streamJob(running.id, running.name, hooks, list.bootId); + return; + } + // The job this page started ended while the user was away: show how. + const finished = stored && stored.bootId === list.bootId && (list.jobs || []).find((j) => j.id === stored.jobId && names.includes(j.name)); + if (finished && !stored.seen) { + let job = finished; + try { job = await fetchJSON(`/api/jobs/${finished.id}`, { timeoutMs: 5000 }); } catch (_) { /* the summary is enough */ } + Object.assign(run, { jobId: job.id, name: job.name, phase: "", done: 0, total: 0, label: "", + startedAt: Date.parse(job.start) || Date.now(), conn: "running" }); + endRun(job); + writeStore(LAST_JOB_KEY, { ...stored, seen: true }); + hooks?.onEnd?.(job); + return; + } + if (stored && stored.bootId && stored.bootId !== list.bootId && names.includes(stored.name)) { + Object.assign(run, { jobId: stored.jobId, name: stored.name, status: "failed", startedAt: Date.now(), endedAt: Date.now(), conn: "lost", + message: `seedstorm restarted since this page started the ${stored.name} job: it was stopped and will write nothing more.` }); + writeStore(LAST_JOB_KEY, null); + renderRunStrip(); + } + } + + // A page error must never be silent: running jobs are not affected, and the + // user is told a reload is safe. + function setupGlobalErrorNotice() { + let shown = false; + const show = (detail) => { + console.error("seedstorm page error:", detail); + if (shown) return; + shown = true; + const box = document.createElement("div"); + box.className = "page-error-notice"; + box.setAttribute("role", "alert"); + box.setAttribute("data-testid", "page-error-notice"); + box.innerHTML = `A page error happened; running jobs are not affected. Reload is safe. + `; + box.querySelector("button").addEventListener("click", () => { box.remove(); shown = false; }); + document.body.appendChild(box); + }; + window.addEventListener("error", (ev) => show(ev.error || ev.message)); + window.addEventListener("unhandledrejection", (ev) => show(ev.reason)); + } + // ── shared job streaming ────────────────────────────────────────────── let elapsedTimer = null; function startElapsed() { @@ -665,13 +928,15 @@ const m = /^\[(\d+)\]\s?(.*)$/.exec(s); return m ? m[2] : s; } - function streamJob(jobId, jobName, hooks) { + function streamJob(jobId, jobName, hooks, bootId) { const cancel = document.getElementById("job-cancel"); setStatus("running"); resetPhases(); + beginRun(jobId, jobName, bootId); if (cancel) { cancel.disabled = false; - cancel.onclick = () => fetch(`/api/jobs/${jobId}/cancel`, { method: "POST" }); + cancel.onclick = () => fetch(`/api/jobs/${jobId}/cancel`, { method: "POST" }) + .catch((err) => appendLog("ERROR: could not cancel: " + (err.message || err))); } const expandAll = document.getElementById("job-expand-all"); if (expandAll) { @@ -681,36 +946,72 @@ expandAll.dataset.open = (!open).toString(); }; } + // The browser reconnects on its own after a dropped connection and sends + // Last-Event-ID, so the server resumes after the last event we saw. const es = new EventSource(`/api/jobs/${jobId}/stream`); + let settled = false; + const settle = (job) => { + if (settled) return; + settled = true; + es.close(); + if (cancel) cancel.disabled = true; + endRun(job); + setStatus(job.status); + finalizeLastPhase(job.status); + hooks?.onEnd?.(job); + }; + const settleFromServer = () => + fetch(`/api/jobs/${jobId}`, { cache: "no-store" }) + .then((r) => (r.ok ? r.json() : Promise.reject(new Error(`job lookup failed (${r.status})`)))) + .then(settle) + .catch((err) => settle({ id: jobId, name: jobName, status: "failed", error: `Lost track of the job: ${err.message || err}` })); + const seen = () => { run.lastEventAt = Date.now(); run.lastPingAt = run.lastEventAt; if (run.conn !== "lost") run.conn = "running"; }; + es.addEventListener("ping", () => { run.lastPingAt = Date.now(); if (run.conn === "reconnecting") run.conn = "running"; }); es.addEventListener("log", (e) => { + seen(); const text = stripSeq(e.data); appendLog(text); hooks?.onLog?.(text); }); es.addEventListener("phase", (e) => { + seen(); const text = stripSeq(e.data); + run.phase = text; startPhase(text); }); es.addEventListener("progress", (e) => { // payload: `[seq] done/total label` const m = /^\[\d+\]\s?(\d+)\/(\d+)\s?(.*)$/.exec(e.data); if (!m) return; + seen(); + run.done = Number(m[1]); run.total = Number(m[2]); run.label = m[3]; setProgress(Number(m[1]), Number(m[2]), m[3]); }); es.addEventListener("status", (e) => setStatus(e.data)); - es.addEventListener("error", (e) => { + // The server names a job's failure "failure": a named "error" event would + // trigger onerror below and used to close the stream before "end". + es.addEventListener("failure", (e) => { if (e.data) appendLog("ERROR: " + e.data); }); es.addEventListener("end", () => { es.close(); - if (cancel) cancel.disabled = true; - fetch(`/api/jobs/${jobId}`).then(r => r.json()).then((j) => { - setStatus(j.status); - finalizeLastPhase(j.status); - hooks?.onEnd?.(j); - }); - }); - es.onerror = () => { es.close(); }; + settleFromServer(); + }); + es.onerror = () => { + if (settled) return; + if (es.readyState === EventSource.CONNECTING) { run.conn = "reconnecting"; renderRunStrip(); } + // Connection trouble: if the job is over, finish with its real state; + // if it is still running, let the browser reconnect and resume. + fetch(`/api/jobs/${jobId}`, { cache: "no-store" }) + .then((r) => (r.ok ? r.json() : Promise.reject(new Error(`job lookup failed (${r.status})`)))) + .then((job) => { + if (job.status !== "running" && job.status !== "pending") settle(job); + else if (es.readyState === EventSource.CLOSED) streamJob(jobId, jobName, hooks); + }) + .catch(() => { + if (es.readyState === EventSource.CLOSED) settle({ id: jobId, name: jobName, status: "failed", error: "Connection to seedstorm lost" }); + }); + }; } // ── simple run-form (used by /generate, /enrich, /export pages) ─────── @@ -736,14 +1037,15 @@ form.querySelectorAll('input[type="checkbox"]').forEach((el) => { if (!(el.name in payload)) payload[el.name] = false; }); - const res = await fetch(endpoint, { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify(payload), - }); - const j = await res.json(); document.getElementById("job-panel").hidden = false; - if (!res.ok) { resetPhases(); appendLog("ERROR: " + (j.error || res.statusText)); return; } + let j; + try { + j = await postRun(endpoint, payload); + } catch (err) { + resetPhases(); + appendLog("ERROR: " + (err.message || err)); + return; + } streamJob(j.id, j.name, { onEnd: (job) => { const r = job.result || {}; @@ -751,7 +1053,7 @@ if (!out) return; renderJobResult(out, r, job.name || j.name || "run"); }, - }); + }, j.bootId); }); } @@ -1158,8 +1460,10 @@ b.addEventListener("click", () => { document.querySelectorAll(".ws-mode-pill").forEach(x => x.classList.remove("active")); b.classList.add("active"); + const wasClone = ws.mode === "clone"; ws.mode = b.dataset.mode; updateCloneControls(); + if (ws.mode === "clone" && !wasClone) loadCloneTargetAccess(); recomputeAuto(); refreshSelectionUI(); }); @@ -1169,7 +1473,7 @@ document.querySelector('[data-act="none"]').addEventListener("click", () => clearSelection()); document.querySelector('[data-act="empty"]').addEventListener("click", () => selectEmpty()); document.querySelector('[data-act="invert"]').addEventListener("click", () => invertSelection()); - document.querySelector('[data-act="refresh"]').addEventListener("click", () => refreshCounts()); + document.querySelector('[data-act="refresh"]').addEventListener("click", () => refreshCounts(true)); document.getElementById("ws-search")?.addEventListener("input", (ev) => applySearch(ev.target.value)); document.getElementById("ws-search")?.addEventListener("keydown", (ev) => { if (ev.key === "Enter") { @@ -1195,6 +1499,7 @@ if (ws.onlyMatches) { exitOnlyMatches(); fitGraph(); } else enterOnlyMatches(); }); setupMinimap(); + setupWorkspaceFormMemory(); syncTuningSummary(); document.getElementById("ws-fit")?.addEventListener("click", () => fitGraph()); document.getElementById("ws-zoom-in")?.addEventListener("click", () => zoomGraph(1.18)); @@ -1217,11 +1522,68 @@ document.getElementById("ws-run").addEventListener("click", runMode); loadCloneTargets(); - loadProfileOptions(); + loadProfileOptions().then(() => { + restoreWorkspaceForm(true); + const id = document.getElementById("cfg-profile")?.value; + if (id) loadProfileInsights(id); + }); loadGraph(); + resumeRun(["seed", "gaps", "generate", "clone-schema"], { onLog: (line) => onLogPulse(line), onEnd: (job) => onJobEnd(job) }); loadWorkspaceAccess(false); } + // ── remembered workspace form (per connection) ───────────────────────── + // Volume, tuning, profile and dry-run survive leaving the page. Destructive + // toggles (truncate, disable FK checks, drop existing on clone) always start + // off: they are never remembered. + const WS_FORM_KEY = "seedstorm.workspaceForm.v1"; + const WS_FORM_FIELDS = [ + ["cfg-rows", "value"], ["cfg-workers", "value"], ["cfg-gen-workers", "value"], ["cfg-batch", "value"], + ["cfg-enum", "value"], ["cfg-selfref-depth", "value"], ["cfg-dryrun", "checked"], ["cfg-profile", "value"], + ["cfg-clone-dryrun", "checked"], ["ws-clone-views", "checked"], ["ws-clone-routines", "checked"], ["ws-clone-triggers", "checked"], + ]; + function formConnKey() { + return document.body.dataset.connKey || "default"; + } + function saveWorkspaceForm() { + const all = readStore(WS_FORM_KEY, {}); + const values = {}; + for (const [id, prop] of WS_FORM_FIELDS) { + const el = document.getElementById(id); + if (el) values[id] = el[prop]; + } + all[formConnKey()] = { ...values, savedAt: Date.now() }; + // Keep the ten most recent connections. + const keys = Object.keys(all).sort((a, b) => (all[b].savedAt || 0) - (all[a].savedAt || 0)); + for (const k of keys.slice(10)) delete all[k]; + writeStore(WS_FORM_KEY, all); + } + // restoreWorkspaceForm sets remembered values; the profile select only once + // its options exist (they load asynchronously). + function restoreWorkspaceForm(onlyProfile) { + const values = readStore(WS_FORM_KEY, {})[formConnKey()]; + if (!values) return; + for (const [id, prop] of WS_FORM_FIELDS) { + if ((id === "cfg-profile") !== !!onlyProfile) continue; + const el = document.getElementById(id); + if (!el || !(id in values)) continue; + if (id === "cfg-profile") { + const known = [...el.options].some((o) => o.value === values[id]); + if (!known) continue; // the profile was deleted since + } + el[prop] = values[id]; + } + } + function setupWorkspaceFormMemory() { + for (const [id] of WS_FORM_FIELDS) { + const el = document.getElementById(id); + if (!el) continue; + el.addEventListener("input", saveWorkspaceForm); + el.addEventListener("change", saveWorkspaceForm); + } + restoreWorkspaceForm(false); + } + function syncTuningSummary() { const n = Number(document.getElementById("cfg-workers")?.value || 0); const g = Number(document.getElementById("cfg-gen-workers")?.value || 0); @@ -1253,6 +1615,13 @@ const data = await res.json(); if (!res.ok) throw new Error(data.error || res.statusText); initGraph(data); + if (data.countsTakenAt) { + ws.countsTakenAt = Date.parse(data.countsTakenAt); + const status = document.getElementById("ws-counts-status"); + if (status) { status.dataset.state = "ready"; status.textContent = countsAgeText(); } + } else { + refreshCounts(false); + } } catch (err) { setGraphLoading("Graph failed", err.message || String(err), true); } @@ -1934,10 +2303,11 @@ target.hidden = false; target.innerHTML = '

Loading rows...

'; const q = new URLSearchParams({ table: tableName, limit: "5", offset: "0", _: String(Date.now()) }); - const res = await fetch("/api/table?" + q.toString(), { cache: "no-store" }); - const data = await res.json(); - if (!res.ok) { - target.innerHTML = `

Preview failed: ${escapeHTML(data.error || res.statusText)}

`; + let data; + try { + data = await fetchJSON("/api/table?" + q.toString()); + } catch (err) { + target.innerHTML = `

Preview failed: ${escapeHTML(err.message || String(err))}

`; return; } if (!data.rows || data.rows.length === 0) { @@ -2373,10 +2743,16 @@ if (legend) legend.hidden = noInsert.size === 0; } + // Checking a saved target opens (and pings) that connection on the server, + // so it only happens while clone mode is active. async function loadCloneTargetAccess() { const target = document.getElementById("cfg-clone-target"); const selected = target?.selectedOptions?.[0]; ws.cloneAccess = null; + if (ws.mode !== "clone") { + updateAccessWarnings(); + return; + } if (selected && selected.value && selected.dataset.kind !== "empty") { const key = selected.dataset.kind === "saved" ? "savedId" : "id"; try { @@ -2482,7 +2858,7 @@ ws.preview.offset = 0; const target = document.getElementById("ws-detail"); target.innerHTML = "

loading...

"; - fetch("/api/schema").then(r => r.json()).then((sc) => { + fetchJSON("/api/schema").then((sc) => { const t = (sc.tables && sc.tables[tableName]) || (sc.Tables && sc.Tables[tableName]); if (!t) { target.innerHTML = "

not in schema

"; return; } const entries = Object.entries(t.columns || t.Columns); @@ -2560,6 +2936,8 @@ loadPreview(tableName); }); loadPreview(tableName); + }).catch((err) => { + target.innerHTML = `

Could not load ${escapeHTML(tableName)}: ${escapeHTML(err.message || String(err))}

`; }); } @@ -2603,7 +2981,12 @@ async function ensureSchemaColumns(tableName) { if (ws.schemaColumns[tableName]) return; - const sc = await fetch("/api/schema").then(r => r.json()); + let sc; + try { + sc = await fetchJSON("/api/schema"); + } catch (_) { + return; // column hints are optional: the preview still renders without them + } const t = (sc.tables && sc.tables[tableName]) || (sc.Tables && sc.Tables[tableName]); if (!t) return; const entries = Object.entries(t.columns || t.Columns); @@ -2630,10 +3013,11 @@ offset: String(ws.modal.offset), _: String(Date.now()), }); - const res = await fetch("/api/table?" + q.toString(), { cache: "no-store" }); - const data = await res.json(); - if (!res.ok) { - box.innerHTML = `

Preview failed: ${escapeHTML(data.error || res.statusText)}

`; + let data; + try { + data = await fetchJSON("/api/table?" + q.toString()); + } catch (err) { + box.innerHTML = `

Preview failed: ${escapeHTML(err.message || String(err))}

`; return; } const start = data.total === 0 ? 0 : data.offset + 1; @@ -2659,10 +3043,11 @@ offset: String(ws.preview.offset), _: String(Date.now()), }); - const res = await fetch("/api/table?" + q.toString(), { cache: "no-store" }); - const data = await res.json(); - if (!res.ok) { - box.innerHTML = `

Preview failed: ${escapeHTML(data.error || res.statusText)}

`; + let data; + try { + data = await fetchJSON("/api/table?" + q.toString()); + } catch (err) { + box.innerHTML = `

Preview failed: ${escapeHTML(err.message || String(err))}

`; return; } const start = data.total === 0 ? 0 : data.offset + 1; @@ -2754,20 +3139,17 @@ document.getElementById("job-result").innerHTML = ""; if (ws.cy) ws.cy.nodes().removeClass("seeding done failed"); - const res = await fetch(endpoint, { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify(cfg), - }); - const j = await res.json(); - if (!res.ok) { - appendLog("ERROR: " + (j.error || res.statusText)); + let j; + try { + j = await postRun(endpoint, cfg); + } catch (err) { + appendLog("ERROR: " + (err.message || err)); return; } streamJob(j.id, j.name, { onLog: (line) => onLogPulse(line), onEnd: (job) => onJobEnd(job), - }); + }, j.bootId); } async function runCloneSchema() { @@ -2790,14 +3172,11 @@ activateTab("logs"); resetPhases(); document.getElementById("job-result").innerHTML = ""; - const res = await fetch("/api/clone-schema", { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify(cfg), - }); - const j = await res.json(); - if (!res.ok) { - appendLog("ERROR: " + (j.error || res.statusText)); + let j; + try { + j = await postRun("/api/clone-schema", cfg); + } catch (err) { + appendLog("ERROR: " + (err.message || err)); return; } streamJob(j.id, j.name, { @@ -2815,7 +3194,7 @@ out.querySelector(".result-shell")?.appendChild(box); } }, - }); + }, j.bootId); } function tableRowPayload() { @@ -2849,35 +3228,68 @@ refreshCounts(); } - function refreshCounts() { + // refreshCounts fills node counts without blocking the graph: the page is + // usable while tables are counted. force recounts instead of reusing the + // counts the server cached for this connection. + let countsRequest = 0; + function refreshCounts(force) { if (!ws.cy) return; - setGraphLoading("Refreshing counts", "Reading row counts for every table."); - fetch("/api/counts").then(r => r.json()).then((counts) => { - ws.cy.batch(() => { - ws.cy.nodes().forEach((n) => { - const id = n.id(); - if (id in counts) { - n.data("count", counts[id]); - n.data("counted", true); - n.data("countLabel", formatCount(counts[id])); - } - }); + const status = document.getElementById("ws-counts-status"); + const setState = (state, text) => { + if (!status) return; + status.dataset.state = state; + status.textContent = text; + }; + const seq = ++countsRequest; + setState("loading", `Counting rows in ${ws.nodes.length} tables…`); + fetch("/api/counts" + (force ? "?refresh=1" : ""), { cache: "no-store" }) + .then(async (r) => { + const body = await r.json().catch(() => ({})); + if (!r.ok) throw new Error(body.error || r.statusText); + return { counts: body, takenAt: r.headers.get("X-Counts-Taken-At") }; + }) + .then(({ counts, takenAt }) => { + if (seq !== countsRequest) return; + applyCounts(counts); + const missing = ws.nodes.filter((n) => !(n.id in counts)).length; + ws.countsTakenAt = takenAt ? Date.parse(takenAt) : Date.now(); + setState("ready", countsAgeText() + (missing ? ` · ${missing} could not be counted` : "")); + }) + .catch((err) => { + if (seq !== countsRequest) return; + setState("failed", "Row counts unavailable: " + (err.message || err)); }); - // Keep the JS-side mirror in sync so isPopulated() sees fresh counts. - for (const n of ws.nodes) { - if (n.id in counts) { - n.count = counts[n.id]; - n.counted = true; + } + + function applyCounts(counts) { + ws.cy.batch(() => { + ws.cy.nodes().forEach((n) => { + const id = n.id(); + if (id in counts) { + n.data("count", counts[id]); + n.data("counted", true); + n.data("countLabel", formatCount(counts[id])); } - } - updateStats(); - recomputeAuto(); - refreshSelectionUI(); - renderIgnoredTab(document.getElementById("ws-ignored-profile")?.textContent); - clearGraphLoading(); - }).catch((err) => { - setGraphLoading("Count refresh failed", err.message || String(err), true); + }); }); + // Keep the JS-side mirror in sync so isPopulated() sees fresh counts. + for (const n of ws.nodes) { + if (n.id in counts) { + n.count = counts[n.id]; + n.counted = true; + } + } + updateStats(); + recomputeAuto(); + refreshSelectionUI(); + renderIgnoredTab(document.getElementById("ws-ignored-profile")?.textContent); + } + + function countsAgeText() { + if (!ws.countsTakenAt) return ""; + const s = Math.max(0, Math.round((Date.now() - ws.countsTakenAt) / 1000)); + const age = s < 45 ? "just now" : s < 3600 ? `${Math.round(s / 60)} min ago` : `${Math.round(s / 3600)} h ago`; + return `Row counts from ${age} · ↻ to recount`; } // Narrow screens collapse the top navigation into a drawer. @@ -2928,8 +3340,19 @@ } } + function setupProductionToggle() { + const box = document.getElementById("conn-production"); + const row = document.getElementById("conn-confirm-label-row"); + if (!box || !row) return; + const sync = () => { row.hidden = !(box.dataset.wasProduction && !box.checked); }; + box.addEventListener("change", sync); + sync(); + } + document.addEventListener("DOMContentLoaded", () => { + setupGlobalErrorNotice(); setupNavDrawer(); + setupProductionToggle(); setupAccessBadge(); setupConnectForm(); setupConnectionDialog(); @@ -2947,7 +3370,8 @@ window.seedstorm = { // Shared helpers for page scripts (compare.js, profiles.js). ui: { - streamJob, resetPhases, appendLog, escapeHTML, formatCount, copyText, + streamJob, postRun, confirmProduction, resumeRun, fetchJSON, readStore, writeStore, + resetPhases, appendLog, escapeHTML, formatCount, copyText, fetchConnections, fetchSavedConnections, connectionLabel, connectionKey, fetchAccess, }, state: ws, diff --git a/internal/web/static/compare.js b/internal/web/static/compare.js index 448ae50..c0add19 100644 --- a/internal/web/static/compare.js +++ b/internal/web/static/compare.js @@ -9,18 +9,21 @@ const $ = (id) => document.getElementById(id); const IMPORTED_KEY = "seedstorm.importedCounts.v1"; + // Picks and mirror settings survive leaving the page. Reset mode is never + // remembered: it truncates the target. + const COMPARE_FORM_KEY = "seedstorm.compareForm.v1"; + const COMPARE_FORM_FIELDS = [ + ["cmp-scale", "value"], ["cmp-max-rows", "value"], ["cmp-parent-rows", "value"], ["cmp-profile", "value"], + ["cmp-batch", "value"], ["cmp-selfref", "value"], ["cmp-workers", "value"], ["cmp-stop", "checked"], + ]; const REPORT_KEY = "seedstorm.compareReport.v1"; const MAX_IMPORTED = 8; const MAX_SAVED_REPORTS = 6; - // Browser storage can be unavailable (private windows, blocked site data): - // every read and write falls back to nothing. - function readStore(key, fallback) { - try { return JSON.parse(localStorage.getItem(key) || "null") ?? fallback; } catch (_) { return fallback; } - } - function writeStore(key, value) { - try { localStorage.setItem(key, JSON.stringify(value)); return true; } catch (_) { return false; } - } + // Browser storage helpers are shared (app.js): unavailable storage falls + // back to nothing. + const readStore = (key, fallback) => window.seedstorm.ui.readStore(key, fallback); + const writeStore = (key, value) => window.seedstorm.ui.writeStore(key, value); const state = { report: null, @@ -42,14 +45,14 @@ const liveKeys = new Set(); for (const c of live) { liveKeys.add(ui().connectionKey(c.info)); - options.push({ value: "id:" + c.id, label: ui().connectionLabel(c.info), group: "Connected", active: c.active, dbType: c.info.dbType }); + options.push({ value: "id:" + c.id, label: ui().connectionLabel(c.info) + (c.production ? " · production" : ""), group: "Connected", active: c.active, dbType: c.info.dbType }); } for (const c of saved) { if (liveKeys.has(ui().connectionKey(c))) continue; const locked = !c.hasPassword && !c.dsn; options.push({ value: "saved:" + c.id, - label: ui().connectionLabel(c) + (locked ? " — connect once to store its password" : ""), + label: ui().connectionLabel(c) + (c.production ? " · production" : "") + (locked ? " — connect once to store its password" : ""), group: "Saved", disabled: locked, dbType: c.dbType, }); } @@ -60,11 +63,14 @@ })); fillSelect($("cmp-source"), [...options, ...importedOptions]); fillSelect($("cmp-target"), options); + // URL (a link from another page) > remembered picks > active and next connection. const params = new URLSearchParams(location.search); + const remembered = readStore(COMPARE_FORM_KEY, {}); const active = options.find((o) => o.active); const other = options.find((o) => !o.active && !o.disabled); - setSelect($("cmp-source"), params.get("source") || active?.value); - setSelect($("cmp-target"), params.get("target") || other?.value); + const known = (v) => v && [...$("cmp-source").options].some((o) => o.value === v && !o.disabled); + setSelect($("cmp-source"), params.get("source") || (known(remembered.source) ? remembered.source : active?.value)); + setSelect($("cmp-target"), params.get("target") || ([...$("cmp-target").options].some((o) => o.value === remembered.target && !o.disabled) ? remembered.target : other?.value)); const usable = options.filter((o) => !o.disabled).length; const hint = $("cmp-hint"); @@ -158,14 +164,12 @@ } // ── jobs ──────────────────────────────────────────────────────────── - function runJob(endpoint, body) { - return new Promise(async (resolve, reject) => { - const res = await fetch(endpoint, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify(body) }); - const j = await res.json().catch(() => ({})); - if (!res.ok) return reject(new Error(j.error || res.statusText)); - $("job-name").textContent = j.name; - ui().streamJob(j.id, j.name, { onEnd: (job) => resolve(job) }); - }); + // runJob starts a job and settles when it ends; a request that never reaches + // the server rejects instead of leaving the caller waiting. + async function runJob(endpoint, body) { + const j = await ui().postRun(endpoint, body); + $("job-name").textContent = j.name; + return new Promise((resolve) => ui().streamJob(j.id, j.name, { onEnd: resolve }, j.bootId)); } function setBusy(busy, label) { @@ -297,12 +301,12 @@ function renderStats(r) { const t = r.totals; const pct = t.sourceRows > 0 ? Math.round((t.targetRows / t.sourceRows) * 100) : null; - const attention = t.sourceOnly + t.targetOnly + t.columnDrift; + const attention = t.sourceOnly + t.targetOnly + t.columnDrift + (t.unknown || 0); const tiles = [ - { label: "tables matched", value: fmt(t.same + t.differs), note: `${fmt(t.same)} same · ${fmt(t.differs)} differ` }, + { label: "tables matched", value: fmt(t.same + t.differs + (t.unknown || 0)), note: `${fmt(t.same)} same · ${fmt(t.differs)} differ` + (t.unknown ? ` · ${fmt(t.unknown)} count unknown` : "") }, { label: "rows", value: `${fmt(t.sourceRows)} → ${fmt(t.targetRows)}`, note: (pct == null ? "source is empty" : `target holds ${pct}% of source`) + (hasEstimates(r) ? " · ~ = estimated" : ""), meter: pct }, { label: "size", value: `${bytes(t.sourceBytes)} → ${bytes(t.targetBytes)}`, note: r.source.dbType === "mysql" || r.target.dbType === "mysql" ? "MySQL sizes are cached estimates" : "data + indexes" }, - { label: "needs a look", value: fmt(attention), note: `${t.sourceOnly} source-only · ${t.targetOnly} target-only · ${t.columnDrift} column drift`, warn: attention > 0 }, + { label: "needs a look", value: fmt(attention), note: `${t.sourceOnly} source-only · ${t.targetOnly} target-only · ${t.columnDrift} column drift` + (t.unknown ? ` · ${t.unknown} count unknown` : ""), warn: attention > 0 }, ]; $("cmp-stats").innerHTML = tiles.map((tile, i) => `
@@ -329,7 +333,7 @@ const mirrorable = row.status === "same" || row.status === "differs"; const checked = mirrorable && !state.excluded.has(row.table); const delta = row.delta || 0; - const statusText = { same: "same", differs: "differs", source_only: "not on target", target_only: "not on source" }[row.status]; + const statusText = { same: "same", differs: "differs", source_only: "not on target", target_only: "not on source", unknown: "count unknown" }[row.status]; const driftBadge = drift(row) ? `columns ≠` : ""; @@ -677,7 +681,30 @@ } catch (_) { /* optional */ } } + // A compare or mirror started before leaving the page keeps running on the + // server: reattach to it and show its outcome when it ends. + function resumeRunningJob() { + ui().resumeRun(["compare", "mirror"], { + onEnd: (job) => { + if (job.name === "compare" && job.status === "done" && job.result?.report) { + state.report = job.result.report; + state.pairKey = pairKey(); + saveReport(); + render(); + } else if (job.status !== "done") { + showOutcome("error", `${job.name === "mirror" ? "Mirror" : "Compare"} ${job.status}`, job.error || ""); + $("cmp-results").hidden = false; + } else if (job.name === "mirror") { + const run = job.result?.run || {}; + showOutcome("ok", `Inserted ${fmt(run.inserted || 0)} rows`, "The mirror that was running when you left finished."); + $("cmp-results").hidden = false; + } + }, + }); + } + document.addEventListener("DOMContentLoaded", () => { + resumeRunningJob(); $("cmp-form").addEventListener("submit", (ev) => { ev.preventDefault(); compare(); }); $("cmp-source").addEventListener("change", syncPickers); $("cmp-target").addEventListener("change", syncPickers); @@ -749,6 +776,39 @@ app.querySelectorAll(".cmp-tab").forEach((b) => b.addEventListener("click", () => activateTab(b.dataset.tab))); document.addEventListener("keydown", (ev) => { if (ev.key === "Escape" && !$("cmp-modal").hidden) closeModal(); }); loadPickers().then(() => restoreReport()); - loadProfiles(); + loadProfiles().then(() => restoreCompareForm(true)); + restoreCompareForm(false); + for (const [id] of COMPARE_FORM_FIELDS) { + $(id)?.addEventListener("input", saveCompareForm); + $(id)?.addEventListener("change", saveCompareForm); + } + ["cmp-source", "cmp-target"].forEach((id) => $(id).addEventListener("change", saveCompareForm)); + app.querySelectorAll('input[name="counts"]').forEach((r) => r.addEventListener("change", saveCompareForm)); }); + + function saveCompareForm() { + const values = { source: $("cmp-source").value, target: $("cmp-target").value, counts: app.querySelector('input[name="counts"]:checked')?.value }; + for (const [id, prop] of COMPARE_FORM_FIELDS) { + if ($(id)) values[id] = $(id)[prop]; + } + writeStore(COMPARE_FORM_KEY, values); + } + + // restoreCompareForm sets remembered settings; the profile only once its + // options have loaded, and only if it still exists. + function restoreCompareForm(onlyProfile) { + const values = readStore(COMPARE_FORM_KEY, null); + if (!values) return; + if (!onlyProfile && values.counts) { + const radio = app.querySelector(`input[name="counts"][value="${values.counts}"]`); + if (radio) radio.checked = true; + } + for (const [id, prop] of COMPARE_FORM_FIELDS) { + const el = $(id); + if (!el || !(id in values) || (id === "cmp-profile") !== !!onlyProfile) continue; + if (id === "cmp-profile" && ![...el.options].some((o) => o.value === values[id])) continue; + el[prop] = values[id]; + } + if (!onlyProfile && values["cmp-scale"]) setScale(values["cmp-scale"]); + } })(); diff --git a/internal/web/static/profiles.js b/internal/web/static/profiles.js index d92cac5..2f7a86b 100644 --- a/internal/web/static/profiles.js +++ b/internal/web/static/profiles.js @@ -8,6 +8,7 @@ const app = document.getElementById("profiles-app"); if (!app) return; const $ = (id) => document.getElementById(id); + const ui = () => window.seedstorm.ui; const esc = (v) => window.seedstorm.ui.escapeHTML(v == null ? "" : v); const KINDS = [ @@ -726,12 +727,13 @@ setStatusNote("Give the profile a name first; the CLI uses it with --profile."); return; } - const res = await fetch("/api/profiles", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ id: state.id, rules: state.doc }) }); - const data = await res.json().catch(() => ({})); - if (!res.ok) { + let data; + try { + data = await ui().fetchJSON("/api/profiles", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ id: state.id, rules: state.doc }) }); + } catch (err) { const el = $("pf-status"); el.className = "pf-status small error"; - el.textContent = data.error || res.statusText; + el.textContent = "Not saved: " + (err.message || err); return; } await loadProfiles(); @@ -743,7 +745,14 @@ async function remove() { const p = state.profiles.find((x) => x.id === state.id); if (!p || !window.confirm(`Delete profile “${p.rules.name}”? Runs that reference it by name will stop working.`)) return; - await fetch("/api/profiles?id=" + encodeURIComponent(p.id), { method: "DELETE" }); + try { + await ui().fetchJSON("/api/profiles?id=" + encodeURIComponent(p.id), { method: "DELETE" }); + } catch (err) { + const el = $("pf-status"); + el.className = "pf-status small error"; + el.textContent = "Not deleted: " + (err.message || err); + return; + } await loadProfiles(); openProfile(state.profiles[0]?.id || ""); } @@ -754,8 +763,14 @@ $("pf-yaml-error").hidden = true; dialog.dataset.mode = mode; if (mode === "export") { - const res = await fetch("/api/profiles/yaml", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ rules: state.doc }) }); - const data = await res.json(); + let data; + try { + data = await ui().fetchJSON("/api/profiles/yaml", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ rules: state.doc }) }); + } catch (err) { + data = {}; + $("pf-yaml-error").hidden = false; + $("pf-yaml-error").textContent = "Could not render the YAML: " + (err.message || err); + } text.value = data.yaml || ""; text.readOnly = true; $("pf-yaml-eyebrow").textContent = "export"; @@ -807,12 +822,13 @@ } async function loadYAML() { - const res = await fetch("/api/profiles/yaml", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ yaml: $("pf-yaml-text").value }) }); - const data = await res.json().catch(() => ({})); const err = $("pf-yaml-error"); - if (!res.ok) { + let data; + try { + data = await ui().fetchJSON("/api/profiles/yaml", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ yaml: $("pf-yaml-text").value }) }); + } catch (e) { err.hidden = false; - err.textContent = data.error || res.statusText; + err.textContent = e.message || String(e); return; } const existing = state.profiles.find((p) => p.rules.name && p.rules.name.toLowerCase() === String(data.rules.name || "").toLowerCase()); diff --git a/internal/web/static/style.css b/internal/web/static/style.css index 0ecae57..81e0739 100644 --- a/internal/web/static/style.css +++ b/internal/web/static/style.css @@ -1842,6 +1842,13 @@ textarea { font-family: var(--mono); font-size: 12px; } text-transform: lowercase; white-space: nowrap; } .pill.ok { background: rgba(121,216,179,0.14); border-color: rgba(121,216,179,0.4); color: var(--ok); } +.pill.prod { background: rgba(240,120,110,0.14); border-color: rgba(240,120,110,0.45); color: var(--danger, #f0786e); font-weight: 600; } +.prod-dialog { max-width: min(460px, calc(100vw - 32px)); border: 1px solid var(--line, #333); border-radius: 14px; padding: 0; background: var(--panel, #16181c); color: inherit; } +.prod-dialog::backdrop { background: rgba(0,0,0,0.55); } +.prod-dialog form { display: grid; gap: 12px; padding: 18px; } +.prod-dialog h2 { margin: 0; font-size: 1.05rem; } +.prod-dialog input { width: 100%; box-sizing: border-box; } +.prod-dialog footer { display: flex; justify-content: flex-end; gap: 8px; flex-wrap: wrap; } .pill.param { background: rgba(216,181,111,0.12); border-color: rgba(216,181,111,0.32); color: var(--accent-2); } /* ── Connect: driver parameters ────────────────────────────────────────── */ @@ -2382,3 +2389,24 @@ textarea { font-family: var(--mono); font-size: 12px; } .ws-rail, .ws-tab-body, #ws-detail { min-width: 0; } .preview-table-wrap { max-width: 100%; overflow-x: auto; } + +/* ── run strip: the current job next to its action ─────────────────────── */ +.run-strip { display: grid; gap: 4px; min-width: 0; max-width: 100%; padding: 8px 10px; border: 1px solid var(--line, #2a2d33); border-radius: 10px; background: var(--panel, #16181c); font-size: 0.86rem; } +.run-strip-head { display: flex; flex-wrap: wrap; align-items: baseline; gap: 6px 8px; min-width: 0; } +.run-strip-name { min-width: 0; overflow-wrap: anywhere; } +.run-strip-elapsed { margin-left: auto; font-variant-numeric: tabular-nums; } +.run-strip-dot { width: 8px; height: 8px; border-radius: 50%; background: var(--accent, #79d8b3); flex: none; align-self: center; } +.run-strip[data-state="running"] .run-strip-dot, .run-strip[data-state="pending"] .run-strip-dot { animation: run-strip-pulse 1.2s ease-in-out infinite; } +.run-strip[data-state="quiet"] .run-strip-dot { background: #d8b56f; } +.run-strip[data-state="reconnecting"] .run-strip-dot { background: #d8b56f; animation: run-strip-pulse 0.6s ease-in-out infinite; } +.run-strip[data-state="lost"] .run-strip-dot, .run-strip[data-state="failed"] .run-strip-dot { background: var(--danger, #f0786e); } +.run-strip[data-state="canceled"] .run-strip-dot { background: #9aa0a6; } +.run-strip[data-state="failed"], .run-strip[data-state="lost"] { border-color: rgba(240,120,110,0.45); } +.run-strip-bar { width: 100%; height: 6px; } +.run-strip-count { font-variant-numeric: tabular-nums; } +.run-strip p { margin: 0; overflow-wrap: anywhere; } +.run-strip-failure { color: var(--danger, #f0786e); } +.cmp-run-strip { grid-column: 1 / -1; } +@keyframes run-strip-pulse { 50% { opacity: 0.35; } } +@media (prefers-reduced-motion: reduce) { .run-strip-dot { animation: none !important; } } +.page-error-notice { position: fixed; right: 16px; bottom: 16px; left: 16px; max-width: 520px; margin-left: auto; display: flex; gap: 10px; align-items: center; justify-content: space-between; padding: 10px 12px; border-radius: 10px; border: 1px solid rgba(240,120,110,0.45); background: var(--panel, #16181c); z-index: 1000; } diff --git a/internal/web/store.go b/internal/web/store.go index c1bcaa5..7f083c7 100644 --- a/internal/web/store.go +++ b/internal/web/store.go @@ -40,6 +40,10 @@ type SavedConnection struct { Password string `json:"-" yaml:"password,omitempty"` HasPassword bool `json:"hasPassword" yaml:"-"` + // Production marks a database seedstorm must not write to without the + // user typing its label back. + Production bool `json:"production" yaml:"production,omitempty"` + CreatedAt time.Time `json:"createdAt,omitempty" yaml:"createdAt,omitempty"` UsedAt time.Time `json:"usedAt,omitempty" yaml:"usedAt,omitempty"` } diff --git a/internal/web/templates/compare.html.tmpl b/internal/web/templates/compare.html.tmpl index 12fb900..ef0c6e0 100644 --- a/internal/web/templates/compare.html.tmpl +++ b/internal/web/templates/compare.html.tmpl @@ -31,6 +31,7 @@ +
diff --git a/internal/web/templates/connect.html.tmpl b/internal/web/templates/connect.html.tmpl index 78a594b..9506af2 100644 --- a/internal/web/templates/connect.html.tmpl +++ b/internal/web/templates/connect.html.tmpl @@ -99,6 +99,19 @@ — written unencrypted to {{.StorePath}} on this machine. + + {{if .Form.Production}} + + {{end}}
@@ -160,6 +173,7 @@ {{$conn.Target}}
+ {{if $conn.Production}}production{{end}} {{if $live}}connected{{end}} {{if $conn.HasPassword}}saved password{{end}} {{range $conn.Params}}{{.Name}}{{end}} diff --git a/internal/web/templates/layout.html.tmpl b/internal/web/templates/layout.html.tmpl index 856f401..847227f 100644 --- a/internal/web/templates/layout.html.tmpl +++ b/internal/web/templates/layout.html.tmpl @@ -12,7 +12,7 @@ {{block "scripts" .}}{{end}} - +
From 1c5eaa6f44e72cad92b173533f04a7f668b8fadf Mon Sep 17 00:00:00 2001 From: Lucas Machado Date: Thu, 17 Sep 2026 17:13:03 +0200 Subject: [PATCH 03/20] feat: measure and compare relationship shapes - relations.Scan: per-FK degree histogram, percentiles, index gate, estimates, per-edge timeout - snapshot v2 with relationships; snapshot/compare/introspect --relationships - compare shows per-relationship drift --- integration/relationships_test.go | 262 ++++++++++++++++++++++++++++++ internal/cli/compare.go | 26 ++- internal/cli/introspect.go | 16 +- internal/cli/relationships.go | 112 +++++++++++++ internal/cli/snapshot.go | 14 +- internal/compare/compare.go | 5 + internal/compare/render.go | 57 +++++++ internal/compare/shapes.go | 93 +++++++++++ internal/compare/shapes_test.go | 68 ++++++++ internal/compare/snapshot.go | 57 ++++++- internal/compare/snapshot_test.go | 52 +++++- internal/db/relations.go | 140 ++++++++++++++++ internal/relations/scan.go | 215 ++++++++++++++++++++++++ internal/relations/shape.go | 112 +++++++++++++ internal/relations/shape_test.go | 62 +++++++ internal/seeder/relationships.go | 78 +++++++++ 16 files changed, 1355 insertions(+), 14 deletions(-) create mode 100644 integration/relationships_test.go create mode 100644 internal/cli/relationships.go create mode 100644 internal/compare/shapes.go create mode 100644 internal/compare/shapes_test.go create mode 100644 internal/db/relations.go create mode 100644 internal/relations/scan.go create mode 100644 internal/relations/shape.go create mode 100644 internal/relations/shape_test.go create mode 100644 internal/seeder/relationships.go diff --git a/integration/relationships_test.go b/integration/relationships_test.go new file mode 100644 index 0000000..abefd3c --- /dev/null +++ b/integration/relationships_test.go @@ -0,0 +1,262 @@ +//go:build integration + +package integration_test + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "math" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/relations" +) + +// shapeSchema: 5 teams; players per team 0, 1, 1, 3 and 10 (15 players), plus +// 2 free agents with no team (NULL). Every player has exactly 2 badges except +// one with none. badges.player_id has no index on Postgres. +const shapeSchema = ` + CREATE TABLE teams (id INT PRIMARY KEY, name VARCHAR(20)); + CREATE TABLE players (id INT PRIMARY KEY, team_id INT, FOREIGN KEY (team_id) REFERENCES teams (id)); + CREATE INDEX players_team ON players (team_id); + CREATE TABLE badges (id INT PRIMARY KEY, player_id INT NOT NULL, FOREIGN KEY (player_id) REFERENCES players (id))` + +// fillShapeData inserts the rows shapeSchema describes. +func fillShapeData(t *testing.T, conn *sql.DB) { + t.Helper() + execSQL(t, conn, `INSERT INTO teams (id, name) VALUES (1,'a'),(2,'b'),(3,'c'),(4,'d'),(5,'e')`) + teamOf := []int{2, 3, 4, 4, 4} + for i := 0; i < 10; i++ { + teamOf = append(teamOf, 5) + } + id := 1 + for _, team := range teamOf { + execSQL(t, conn, fmt.Sprintf("INSERT INTO players (id, team_id) VALUES (%d, %d)", id, team)) + id++ + } + execSQL(t, conn, fmt.Sprintf("INSERT INTO players (id, team_id) VALUES (%d, NULL); INSERT INTO players (id, team_id) VALUES (%d, NULL)", id, id+1)) + badge := 1 + for p := 1; p <= 16; p++ { // player 17 has no badges + for k := 0; k < 2; k++ { + execSQL(t, conn, fmt.Sprintf("INSERT INTO badges (id, player_id) VALUES (%d, %d)", badge, p)) + badge++ + } + } +} + +func findShape(t *testing.T, shapes []relations.Shape, child, column string) relations.Shape { + t.Helper() + for _, s := range shapes { + if s.Child == child && s.Column == column { + return s + } + } + t.Fatalf("no shape for %s.%s in %+v", child, column, shapes) + return relations.Shape{} +} + +func TestRelationships_ExactShapesMatchTheData(t *testing.T) { + for _, e := range engines() { + t.Run(e.name, func(t *testing.T) { + _, conn := e.scratchDB(t, "ss_relations") + execSQL(t, conn, shapeSchema) + fillShapeData(t, conn) + + tables, err := db.IntrospectConn(context.Background(), conn, e.driver, nil) + if err != nil { + t.Fatal(err) + } + sc := faker.BuildSchema(e.driver, tables) + res, err := relations.Scan(context.Background(), conn, e.driver, sc, relations.Options{Mode: relations.Exact, IncludeUnindexed: true}) + if err != nil { + t.Fatal(err) + } + + teams := findShape(t, res, "players", "team_id") + if teams.Outcome != db.OutcomeOK || teams.Parents != 5 || teams.Children != 15 || teams.NullRows != 2 { + t.Fatalf("players.team_id = %+v", teams) + } + if teams.Min != 1 || teams.Max != 10 || math.Abs(teams.Avg-3.75) > 0.001 || math.Abs(teams.ZeroShare-0.2) > 0.001 { + t.Fatalf("players.team_id degrees = %+v", teams) + } + if math.Abs(teams.NullShare-2.0/17.0) > 0.001 || teams.P95 != 10 || teams.P50 < 1 || teams.P50 > 3 { + t.Fatalf("players.team_id shares/percentiles = %+v", teams) + } + if sum := histogramParents(teams); sum != 4 { + t.Fatalf("histogram counts %d parents with children, want 4: %+v", sum, teams.Histogram) + } + + badges := findShape(t, res, "badges", "player_id") + if badges.Min != 2 || badges.Max != 2 || badges.Avg != 2 || math.Abs(badges.ZeroShare-1.0/17.0) > 0.001 || badges.NullRows != 0 { + t.Fatalf("badges.player_id = %+v", badges) + } + }) + } +} + +func histogramParents(s relations.Shape) int64 { + var n int64 + for _, b := range s.Histogram { + n += b.Parents + } + return n +} + +// Exact scans of an unindexed foreign key are full table scans: skipped by +// default with the reason, and an estimate where the database has one. +func TestRelationships_UnindexedKeyIsSkippedByDefault(t *testing.T) { + e := postgresEngine() + _, conn := e.scratchDB(t, "ss_relations_gate") + execSQL(t, conn, shapeSchema) + tables, err := db.IntrospectConn(context.Background(), conn, e.driver, nil) + if err != nil { + t.Fatal(err) + } + sc := faker.BuildSchema(e.driver, tables) + res, err := relations.Scan(context.Background(), conn, e.driver, sc, relations.Options{Mode: relations.Exact}) + if err != nil { + t.Fatal(err) + } + if s := findShape(t, res, "badges", "player_id"); s.Outcome != relations.OutcomeSkippedUnindexed { + t.Fatalf("unindexed badges.player_id = %+v", s) + } + if s := findShape(t, res, "players", "team_id"); s.Outcome != db.OutcomeOK { + t.Fatalf("indexed players.team_id = %+v", s) + } +} + +// A statement timeout ends one relationship, not the scan; cancelling keeps +// the relationships already measured. +func TestRelationships_TimeoutAndCancelKeepFinishedEdges(t *testing.T) { + e := postgresEngine() + _, conn := e.scratchDB(t, "ss_relations_limits") + execSQL(t, conn, ` + CREATE TABLE parents (id INT PRIMARY KEY); + CREATE TABLE small_kids (id INT PRIMARY KEY, parent_id INT NOT NULL REFERENCES parents (id)); + CREATE INDEX small_kids_parent ON small_kids (parent_id); + CREATE TABLE big_kids (id INT PRIMARY KEY, parent_id INT NOT NULL REFERENCES parents (id)); + CREATE INDEX big_kids_parent ON big_kids (parent_id); + INSERT INTO parents SELECT g FROM generate_series(1, 2000) g; + INSERT INTO small_kids SELECT g, 1 + g % 2000 FROM generate_series(1, 100) g; + INSERT INTO big_kids SELECT g, 1 + g % 2000 FROM generate_series(1, 3000000) g; + ANALYZE`) + tables, err := db.IntrospectConn(context.Background(), conn, e.driver, nil) + if err != nil { + t.Fatal(err) + } + sc := faker.BuildSchema(e.driver, tables) + res, err := relations.Scan(context.Background(), conn, e.driver, sc, relations.Options{ + Mode: relations.Exact, Limits: db.ReadLimits{StatementTimeout: 30 * time.Millisecond}, + }) + if err != nil { + t.Fatal(err) + } + if s := findShape(t, res, "big_kids", "parent_id"); s.Outcome != db.OutcomeTimedOut { + t.Fatalf("3M-row edge with a 30ms limit = %+v", s) + } + if s := findShape(t, res, "small_kids", "parent_id"); s.Outcome != db.OutcomeOK { + t.Fatalf("small edge = %+v", s) + } + + ctx, cancel := context.WithCancel(context.Background()) + var seen []relations.Shape + res, _ = relations.Scan(ctx, conn, e.driver, sc, relations.Options{ + Mode: relations.Exact, + OnEdge: func(done, total int, s relations.Shape) { + seen = append(seen, s) + if s.Child == "small_kids" { + cancel() + } + }, + }) + if s := findShape(t, res, "small_kids", "parent_id"); s.Outcome != db.OutcomeOK { + t.Fatalf("finished edge lost after cancel: %+v", s) + } + if s := findShape(t, res, "big_kids", "parent_id"); s.Outcome != db.OutcomeCancelled { + t.Fatalf("edge after cancel = %+v", s) + } +} + +// TestRelationships_BinarySnapshotAndCompare exports shapes with the binary, +// compares them against a database shaped differently, and refuses a counts +// file without shapes naming the side. +func TestRelationships_BinarySnapshotAndCompare(t *testing.T) { + for _, e := range engines() { + t.Run(e.name, func(t *testing.T) { + srcDSN, src := e.scratchDB(t, "ss_relations_src") + execSQL(t, src, shapeSchema) + fillShapeData(t, src) + tgtDSN, tgt := e.scratchDB(t, "ss_relations_tgt") + execSQL(t, tgt, shapeSchema) + execSQL(t, tgt, `INSERT INTO teams (id, name) VALUES (1,'a'),(2,'b')`) + execSQL(t, tgt, `INSERT INTO players (id, team_id) VALUES (1, 1)`) + execSQL(t, tgt, `INSERT INTO players (id, team_id) VALUES (2, 2)`) + dir := t.TempDir() + + snapPath := filepath.Join(dir, "source.yaml") + _, stderr, err := runBinResult(t, "snapshot", "--db", e.name, "--dsn", srcDSN, "--relationships", "--out", snapPath) + if err != nil { + t.Fatalf("snapshot --relationships: %v\n%s", err, stderr) + } + if !strings.Contains(stderr, "Relationships measured") { + t.Errorf("no summary line:\n%s", stderr) + } + snap := readSnapshotFile(t, snapPath) + teams := findShape(t, snap.Relationships, "players", "team_id") + if teams.Outcome != db.OutcomeOK || teams.Max != 10 || teams.Parents != 5 { + t.Fatalf("exported players.team_id = %+v", teams) + } + badges := findShape(t, snap.Relationships, "badges", "player_id") + if e.name == "postgres" { + if badges.Outcome != relations.OutcomeSkippedUnindexed || !strings.Contains(stderr, "leads no index") { + t.Fatalf("unindexed badges.player_id = %+v\n%s", badges, stderr) + } + } else if badges.Outcome != db.OutcomeOK || badges.Max != 2 { + t.Fatalf("badges.player_id = %+v", badges) + } + + relPath := filepath.Join(dir, "relationships.json") + runBin(t, "introspect", "--db", e.name, "--dsn", srcDSN, "--out", filepath.Join(dir, "schema.yaml"), "--relationships", relPath, "--scan-unindexed") + fromIntrospect := readSnapshotFile(t, relPath) + if b := findShape(t, fromIntrospect.Relationships, "badges", "player_id"); b.Outcome != db.OutcomeOK || b.Max != 2 { + t.Fatalf("introspect --scan-unindexed badges.player_id = %+v", b) + } + + out := runBin(t, "compare", "--source-snapshot", snapPath, "--target-db", e.name, "--target-dsn", tgtDSN, "--relationships", "--format", "json") + var report compare.Report + if err := json.Unmarshal([]byte(out), &report); err != nil { + t.Fatalf("compare json: %v\n%s", err, out) + } + drift := map[string]compare.ShapeDrift{} + for _, d := range report.Relationships { + drift[d.Child+"."+d.Column] = d + } + if d := drift["players.team_id"]; d.Status != compare.ShapeDiffers || d.MaxDelta != -9 { + t.Fatalf("players.team_id drift = %+v", d) + } + + table := runBin(t, "compare", "--source-snapshot", snapPath, "--target-db", e.name, "--target-dsn", tgtDSN, "--relationships", "--only-diff") + if !strings.Contains(table, "Relationships (children per parent") || !strings.Contains(table, "players.team_id → teams") { + t.Fatalf("table output lacks relationships:\n%s", table) + } + + countsOnly := filepath.Join(dir, "counts.yaml") + runBin(t, "snapshot", "--db", e.name, "--dsn", srcDSN, "--out", countsOnly) + _, stderr, err = runBinResult(t, "compare", "--source-snapshot", countsOnly, "--target-db", e.name, "--target-dsn", tgtDSN, "--relationships") + if err == nil || !strings.Contains(stderr, "source") || !strings.Contains(stderr, "no relationships") { + t.Fatalf("counts-only snapshot with --relationships: err=%v\n%s", err, stderr) + } + if data, _ := os.ReadFile(countsOnly); strings.Contains(string(data), "relationships") { + t.Fatalf("counts-only file mentions relationships:\n%s", data) + } + }) + } +} diff --git a/internal/cli/compare.go b/internal/cli/compare.go index 3bdcc35..336a6f6 100644 --- a/internal/cli/compare.go +++ b/internal/cli/compare.go @@ -10,14 +10,17 @@ import ( "github.com/AxeForging/seedstorm/internal/compare" "github.com/AxeForging/seedstorm/internal/logging" + "github.com/AxeForging/seedstorm/internal/relations" "github.com/AxeForging/seedstorm/internal/seeder" ) func compareCmd() *cli.Command { flags := append(endpointFlags(), &cli.StringFlag{Name: "format", Aliases: []string{"f"}, Usage: "Output format: table or json", Value: "table"}, - &cli.BoolFlag{Name: "only-diff", Usage: "Hide tables whose row counts and columns match"}, + &cli.BoolFlag{Name: "only-diff", Usage: "Hide tables (and relationships) that match"}, + &cli.BoolFlag{Name: "relationships", Usage: "Also compare every foreign key's shape (children per parent) on both sides; a snapshot source must include them"}, ) + flags = append(flags, relationshipFlags()...) return &cli.Command{ Name: "compare", Usage: "Compare row counts, sizes and columns between two databases", @@ -25,7 +28,8 @@ func compareCmd() *cli.Command { per table, row counts, on-disk size, the difference, and column-name drift. Works across engines: tables are matched by name, case-insensitively. The source can be a file made by "seedstorm snapshot" (--source-snapshot) -instead of a live database.`, +instead of a live database. --relationships adds a per-foreign-key shape +comparison (average, p95 and max children per parent), measured read-only.`, Flags: flags, Action: func(ctx context.Context, cmd *cli.Command) error { mode, err := countMode(cmd) @@ -44,11 +48,29 @@ instead of a live database.`, if err != nil { return err } + if cmd.Bool("relationships") { + logging.Log.Info().Msg("Comparing relationship shapes") + opts := relationshipOptions(cmd, mode, "") + steps := map[string]func(int, int, relations.Shape){ + "source": relationshipOptions(cmd, mode, "source").OnEdge, + "target": relationshipOptions(cmd, mode, "target").OnEdge, + } + opts.OnEdge = nil + report.Relationships, err = seeder.CompareShapes(ctx, source, target, opts, func(side string, done, total int, s relations.Shape) { + steps[side](done, total, s) + }) + if err != nil { + return err + } + } switch cmd.String("format") { case "json": return writeJSON(report) case "table", "": compare.RenderReport(os.Stdout, report, cmd.Bool("only-diff")) + if cmd.Bool("relationships") { + compare.RenderShapeDrift(os.Stdout, report.Relationships, cmd.Bool("only-diff")) + } return nil default: return fmt.Errorf("unknown format %q (use table or json)", cmd.String("format")) diff --git a/internal/cli/introspect.go b/internal/cli/introspect.go index 2345204..851e81b 100644 --- a/internal/cli/introspect.go +++ b/internal/cli/introspect.go @@ -11,6 +11,7 @@ import ( "github.com/AxeForging/seedstorm/internal/logging" "github.com/AxeForging/seedstorm/internal/runerr" "github.com/AxeForging/seedstorm/internal/schema" + "github.com/AxeForging/seedstorm/internal/seeder" "github.com/urfave/cli/v3" ) @@ -20,8 +21,10 @@ func introspectCmd() *cli.Command { Usage: "Discover database schema and generate a schema YAML file", Description: `Connects to a MySQL or PostgreSQL database and introspects all tables, columns, data types, primary keys, foreign keys, and enum values. -Outputs a schema.yaml that can be used for seeding or AI enrichment.`, - Flags: []cli.Flag{ +Outputs a schema.yaml that can be used for seeding or AI enrichment. +--relationships also measures every foreign key's shape (read-only, +exact) and writes it with estimated table counts to a snapshot file.`, + Flags: append([]cli.Flag{ &cli.StringFlag{ Name: "db", Usage: "Database type: mysql or postgres", @@ -40,7 +43,11 @@ Outputs a schema.yaml that can be used for seeding or AI enrichment.`, Usage: "Output schema YAML file path", Value: "schema.yaml", }, - }, + &cli.StringFlag{ + Name: "relationships", + Usage: "Also measure foreign-key shapes and write them (with estimated counts) to this snapshot file", + }, + }, relationshipFlags()...), Action: func(ctx context.Context, cmd *cli.Command) error { log := logging.Log dbType := normalizeDBType(cmd.String("db")) @@ -79,6 +86,9 @@ Outputs a schema.yaml that can be used for seeding or AI enrichment.`, Int("tables", len(tables)). Msg("Schema saved") + if path := cmd.String("relationships"); path != "" { + return writeRelationshipsSnapshot(ctx, cmd, seeder.Endpoint{Conn: conn, DBType: dbType, Label: dsnLabel(dbType, dsn), Schema: s}, path) + } return nil }, } diff --git a/internal/cli/relationships.go b/internal/cli/relationships.go new file mode 100644 index 0000000..44917dd --- /dev/null +++ b/internal/cli/relationships.go @@ -0,0 +1,112 @@ +package cli + +import ( + "context" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "time" + + "github.com/urfave/cli/v3" + + "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/logging" + "github.com/AxeForging/seedstorm/internal/relations" + "github.com/AxeForging/seedstorm/internal/seeder" +) + +// relationshipFlags tune a relationship scan (--relationships on snapshot, +// compare and introspect). The scan mode follows --counts. +func relationshipFlags() []cli.Flag { + return []cli.Flag{ + &cli.BoolFlag{Name: "scan-unindexed", Usage: "With --relationships: also scan foreign keys that lead no index (full table scans; estimates otherwise)"}, + &cli.DurationFlag{Name: "read-timeout", Usage: "With --relationships: server-side time limit per relationship (a slower one is reported as timed out)", Value: relations.DefaultStatementTimeout}, + } +} + +// relationshipOptions builds scan options; OnEdge logs throttled progress and +// a warning for every relationship that was not measured exactly. +func relationshipOptions(cmd *cli.Command, mode compare.CountMode, side string) relations.Options { + what := "Scanning relationships" + if side != "" { + what += " (" + side + ")" + } + step := stepLogger(what, time.Now) + scanMode := relations.Exact + if mode == compare.CountEstimate { + scanMode = relations.Estimate + } + return relations.Options{ + Mode: scanMode, + Limits: db.ReadLimits{StatementTimeout: cmd.Duration("read-timeout")}, + IncludeUnindexed: cmd.Bool("scan-unindexed"), + OnEdge: func(done, total int, s relations.Shape) { + logShapeWarning(side, s) + step(done, total, s.Child+"."+s.Column) + }, + } +} + +// logShapeWarning explains a relationship that is estimated, unknown or on a +// large table. Safe to call from scan workers (OnEdge runs one at a time). +func logShapeWarning(side string, s relations.Shape) { + if s.Outcome == db.OutcomeOK || s.Outcome == relations.OutcomeEstimated { + if !s.Large || s.Outcome == relations.OutcomeEstimated { + return + } + s.Detail = "large child table: measured exactly, which reads the whole key" + } + ev := logging.Log.Warn() + if side != "" { + ev = ev.Str("side", side) + } + ev.Str("relationship", s.Child+"."+s.Column).Str("outcome", string(s.Outcome)).Msg(s.Detail) +} + +// logShapeSummary counts relationships by outcome. +func logShapeSummary(shapes []relations.Shape) { + byOutcome := map[db.ReadOutcome]int{} + for _, s := range shapes { + byOutcome[s.Outcome]++ + } + ev := logging.Log.Info().Int("relationships", len(shapes)) + outcomes := make([]string, 0, len(byOutcome)) + for outcome := range byOutcome { + outcomes = append(outcomes, string(outcome)) + } + sort.Strings(outcomes) + for _, outcome := range outcomes { + ev = ev.Int(outcome, byOutcome[db.ReadOutcome(outcome)]) + } + ev.Msg("Relationships measured") +} + +// writeRelationshipsSnapshot saves estimated counts plus exact relationship +// shapes to path, as YAML or JSON by its extension. +func writeRelationshipsSnapshot(ctx context.Context, cmd *cli.Command, ep seeder.Endpoint, path string) error { + logging.Log.Info().Msg("Reading estimated table counts") + snap, err := compare.Take(ctx, ep.Conn, ep.DBType, ep.Label, compare.CountEstimate, stepLogger("Counting", time.Now)) + if err != nil { + return err + } + if snap.Relationships, err = ep.Shapes(ctx, relationshipOptions(cmd, compare.CountExact, "")); err != nil { + return err + } + logShapeSummary(snap.Relationships) + format := compare.FormatYAML + if strings.EqualFold(filepath.Ext(path), ".json") { + format = compare.FormatJSON + } + data, err := compare.EncodeSnapshot(snap, format) + if err != nil { + return err + } + if err := os.WriteFile(path, data, 0o644); err != nil { //nolint:gosec // a counts file holds no secrets + return fmt.Errorf("write relationships: %w", err) + } + logging.Log.Info().Str("path", path).Int("relationships", len(snap.Relationships)).Msg("Relationships saved") + return nil +} diff --git a/internal/cli/snapshot.go b/internal/cli/snapshot.go index c7bd512..aee345b 100644 --- a/internal/cli/snapshot.go +++ b/internal/cli/snapshot.go @@ -19,7 +19,8 @@ func snapshotCmd() *cli.Command { Usage: "Save every table's row count and size to a file for later compare or mirror", Description: `Reads one database (read-only) and writes a table-counts snapshot: row count, on-disk size and column names per table, wrapped in a versioned envelope -(kind: seedstorm.table-counts, version: 1). Pass the file to +(kind: seedstorm.table-counts). --relationships adds each foreign key's shape +(children per parent: min, avg, p50, p95, max, histogram) and writes version 2. Pass the file to "compare --source-snapshot" or "mirror --source-snapshot" to follow a database you cannot (or should not) connect to at mirror time. @@ -28,13 +29,14 @@ A hand-written file with only row counts also works: tables: users: 1200 orders: 5000`, - Flags: []cli.Flag{ + Flags: append([]cli.Flag{ &cli.StringFlag{Name: "db", Usage: "Database type: mysql or postgres", Value: "postgres", Sources: cli.EnvVars("SEEDSTORM_DB")}, &cli.StringFlag{Name: "dsn", Usage: "Data source name (connection string)", Required: true, Sources: cli.EnvVars("SEEDSTORM_DSN")}, &cli.StringFlag{Name: "counts", Usage: "Row counts: exact (COUNT(*)) or estimate (planner statistics, fast on large tables)", Value: "exact"}, &cli.StringFlag{Name: "format", Aliases: []string{"f"}, Usage: "Output format: yaml or json", Value: compare.FormatYAML}, &cli.StringFlag{Name: "out", Aliases: []string{"o"}, Usage: "Write the snapshot to this file (default: stdout)"}, - }, + &cli.BoolFlag{Name: "relationships", Usage: "Also measure every foreign key's shape (read-only; exact or estimate per --counts)"}, + }, relationshipFlags()...), Action: func(ctx context.Context, cmd *cli.Command) error { log := logging.Log mode, err := countMode(cmd) @@ -56,6 +58,12 @@ A hand-written file with only row counts also works: if err != nil { return err } + if cmd.Bool("relationships") { + if snap.Relationships, err = ep.Shapes(ctx, relationshipOptions(cmd, mode, "")); err != nil { + return err + } + logShapeSummary(snap.Relationships) + } data, err := compare.EncodeSnapshot(snap, format) if err != nil { return err diff --git a/internal/compare/compare.go b/internal/compare/compare.go index ec008d7..1691520 100644 --- a/internal/compare/compare.go +++ b/internal/compare/compare.go @@ -12,6 +12,7 @@ import ( "time" "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" ) // CountMode selects how row counts are read. @@ -53,6 +54,8 @@ type Snapshot struct { CountMode CountMode `json:"countMode"` TakenAt time.Time `json:"takenAt"` Tables map[string]TableStat `json:"tables"` + // Relationships are foreign-key shapes, when they were scanned. + Relationships []relations.Shape `json:"relationships,omitempty"` } // Take reads table names, columns, sizes and row counts. progress, if set, is @@ -183,6 +186,8 @@ type Report struct { Target SnapshotInfo `json:"target"` Rows []Row `json:"rows"` Totals Totals `json:"totals"` + // Relationships is the per-foreign-key shape drift, when it was compared. + Relationships []ShapeDrift `json:"relationships,omitempty"` } func info(s Snapshot) SnapshotInfo { diff --git a/internal/compare/render.go b/internal/compare/render.go index f72edfe..00dff64 100644 --- a/internal/compare/render.go +++ b/internal/compare/render.go @@ -7,6 +7,7 @@ import ( "text/tabwriter" "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" ) // RenderReport writes the comparison as an aligned text table. onlyDiff hides @@ -131,3 +132,59 @@ func engineName(driver string) string { } return driver } + +// RenderShapeDrift writes relationship shapes side by side: average, p95 and +// maximum children per parent and the share of parents without children. +// onlyDiff hides relationships whose shapes match. +func RenderShapeDrift(w io.Writer, drift []ShapeDrift, onlyDiff bool) { + _, _ = fmt.Fprintln(w, "\nRelationships (children per parent: avg / p95 / max · parents without children)") + tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0) + _, _ = fmt.Fprintln(tw, "RELATIONSHIP\tSOURCE\tTARGET\tSTATUS") + counts := map[ShapeStatus]int{} + shown := 0 + for _, d := range drift { + counts[d.Status]++ + if onlyDiff && d.Status == ShapeSame { + continue + } + shown++ + parent := "" + if s := firstShape(d); s != nil { + parent = " → " + s.Parent + } + _, _ = fmt.Fprintf(tw, "%s.%s%s\t%s\t%s\t%s\n", d.Child, d.Column, parent, shapeCell(d.Source), shapeCell(d.Target), d.Status) + } + _ = tw.Flush() + if onlyDiff && shown == 0 { + _, _ = fmt.Fprintln(w, " (every relationship matches)") + } + _, _ = fmt.Fprintf(w, "%d relationships · same %d · differs %d · source only %d · target only %d · unknown %d\n", + len(drift), counts[ShapeSame], counts[ShapeDiffers], counts[ShapeSourceOnly], counts[ShapeTargetOnly], counts[ShapeUnknown]) +} + +func firstShape(d ShapeDrift) *relations.Shape { + if d.Source != nil { + return d.Source + } + return d.Target +} + +func shapeCell(s *relations.Shape) string { + switch { + case s == nil: + return "—" + case !measured(*s): + return string(s.Outcome) + } + mark := "" + if s.Outcome == relations.OutcomeEstimated || s.Outcome == relations.OutcomeSkippedUnindexed { + mark = "~" + } + num := func(n int64) string { + if n < 0 { + return "?" + } + return fmt.Sprint(n) + } + return fmt.Sprintf("%s%.2f / %s / %s · %.0f%%", mark, s.Avg, num(s.P95), num(s.Max), s.ZeroShare*100) +} diff --git a/internal/compare/shapes.go b/internal/compare/shapes.go new file mode 100644 index 0000000..425be7e --- /dev/null +++ b/internal/compare/shapes.go @@ -0,0 +1,93 @@ +package compare + +import ( + "math" + "sort" + "strings" + + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" +) + +// ShapeStatus classifies one relationship across two databases. +type ShapeStatus string + +const ( + ShapeSame ShapeStatus = "same" + ShapeDiffers ShapeStatus = "differs" + ShapeSourceOnly ShapeStatus = "source_only" + ShapeTargetOnly ShapeStatus = "target_only" + // ShapeUnknown: a side was not measured (timed out, skipped, failed). + ShapeUnknown ShapeStatus = "unknown" +) + +// ShapeDrift is how a relationship's shape differs from source to target. +// Deltas are target minus source. +type ShapeDrift struct { + Child string `json:"child"` + Column string `json:"column"` + Status ShapeStatus `json:"status"` + Source *relations.Shape `json:"source,omitempty"` + Target *relations.Shape `json:"target,omitempty"` + AvgDelta float64 `json:"avgDelta"` + P95Delta int64 `json:"p95Delta"` + MaxDelta int64 `json:"maxDelta"` + ZeroShareDelta float64 `json:"zeroShareDelta"` +} + +// shapeTolerance is how close two averages and shares must be to count as the +// same shape. +const shapeTolerance = 0.05 + +// DiffShapes matches relationships by child table and column, ignoring case. +func DiffShapes(source, target []relations.Shape) []ShapeDrift { + key := func(s relations.Shape) string { return strings.ToLower(s.Child + "." + s.Column) } + tgt := make(map[string]relations.Shape, len(target)) + for _, s := range target { + tgt[key(s)] = s + } + var out []ShapeDrift + matched := map[string]bool{} + for _, s := range source { + src := s + d := ShapeDrift{Child: s.Child, Column: s.Column, Source: &src} + t, ok := tgt[key(s)] + if !ok { + d.Status = ShapeSourceOnly + out = append(out, d) + continue + } + matched[key(s)] = true + t2 := t + d.Target = &t2 + if !measured(s) || !measured(t) { + d.Status = ShapeUnknown + out = append(out, d) + continue + } + d.AvgDelta = t.Avg - s.Avg + d.P95Delta = t.P95 - s.P95 + d.MaxDelta = t.Max - s.Max + d.ZeroShareDelta = t.ZeroShare - s.ZeroShare + d.Status = ShapeSame + if math.Abs(d.AvgDelta) > shapeTolerance*math.Max(1, s.Avg) || d.P95Delta != 0 || d.MaxDelta != 0 || math.Abs(d.ZeroShareDelta) > shapeTolerance { + d.Status = ShapeDiffers + } + out = append(out, d) + } + for _, t := range target { + if matched[key(t)] { + continue + } + t2 := t + out = append(out, ShapeDrift{Child: t.Child, Column: t.Column, Status: ShapeTargetOnly, Target: &t2}) + } + sort.SliceStable(out, func(i, j int) bool { + return strings.ToLower(out[i].Child+"."+out[i].Column) < strings.ToLower(out[j].Child+"."+out[j].Column) + }) + return out +} + +func measured(s relations.Shape) bool { + return s.Outcome == db.OutcomeOK || s.Outcome == relations.OutcomeEstimated || s.Outcome == relations.OutcomeSkippedUnindexed +} diff --git a/internal/compare/shapes_test.go b/internal/compare/shapes_test.go new file mode 100644 index 0000000..3041485 --- /dev/null +++ b/internal/compare/shapes_test.go @@ -0,0 +1,68 @@ +package compare + +import ( + "math" + "strings" + "testing" + + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" +) + +func TestDiffShapes_MatchesRelationshipsAndMeasuresDrift(t *testing.T) { + source := []relations.Shape{ + {Child: "players", Column: "team_id", Parent: "teams", Avg: 3.75, P95: 10, Max: 10, ZeroShare: 0.2, Outcome: db.OutcomeOK}, + {Child: "badges", Column: "player_id", Parent: "players", Avg: 2, P95: 2, Max: 2, Outcome: db.OutcomeOK}, + {Child: "logs", Column: "user_id", Parent: "users", Outcome: db.OutcomeTimedOut}, + } + target := []relations.Shape{ + {Child: "PLAYERS", Column: "TEAM_ID", Parent: "TEAMS", Avg: 1, P95: 1, Max: 1, ZeroShare: 0, Outcome: db.OutcomeOK}, + {Child: "badges", Column: "player_id", Parent: "players", Avg: 2, P95: 2, Max: 2, Outcome: db.OutcomeOK}, + {Child: "extra", Column: "owner_id", Parent: "users", Avg: 4, Outcome: db.OutcomeOK}, + {Child: "logs", Column: "user_id", Parent: "users", Avg: 4, Outcome: db.OutcomeOK}, + } + drift := DiffShapes(source, target) + byKey := map[string]ShapeDrift{} + for _, d := range drift { + byKey[strings.ToLower(d.Child+"."+d.Column)] = d + } + players := byKey["players.team_id"] + if players.Status != ShapeDiffers || math.Abs(players.AvgDelta-(1-3.75)) > 1e-9 || players.MaxDelta != -9 { + t.Fatalf("players.team_id = %+v", players) + } + if byKey["badges.player_id"].Status != ShapeSame { + t.Fatalf("badges = %+v", byKey["badges.player_id"]) + } + if len(drift) != 4 { + t.Fatalf("drift rows = %d, want 4", len(drift)) + } + if byKey["logs.user_id"].Status != ShapeUnknown || byKey["extra.owner_id"].Status != ShapeTargetOnly { + t.Fatalf("unmatched/unknown = %+v / %+v", byKey["logs.user_id"], byKey["extra.owner_id"]) + } +} + +func TestRenderShapeDrift_MarksEstimatesUnknownAndHidesMatches(t *testing.T) { + drift := DiffShapes( + []relations.Shape{ + {Child: "orders", Column: "user_id", Parent: "users", Avg: 2.5, P95: 6, Max: 9, ZeroShare: 0.25, Outcome: db.OutcomeOK}, + {Child: "same", Column: "p_id", Parent: "p", Avg: 1, P95: 1, Max: 1, Outcome: db.OutcomeOK}, + {Child: "events", Column: "user_id", Parent: "users", Avg: 3, P95: -1, Max: -1, Outcome: relations.OutcomeEstimated}, + }, + []relations.Shape{ + {Child: "orders", Column: "user_id", Parent: "users", Avg: 1, P95: 1, Max: 1, Outcome: db.OutcomeOK}, + {Child: "same", Column: "p_id", Parent: "p", Avg: 1, P95: 1, Max: 1, Outcome: db.OutcomeOK}, + {Child: "events", Column: "user_id", Parent: "users", Outcome: db.OutcomeTimedOut}, + }, + ) + var b strings.Builder + RenderShapeDrift(&b, drift, true) + out := b.String() + for _, want := range []string{"orders.user_id → users", "2.50 / 6 / 9 · 25%", "1.00 / 1 / 1 · 0%", "~3.00 / ? / ?", string(db.OutcomeTimedOut), "same 1 · differs 1", "unknown 1"} { + if !strings.Contains(out, want) { + t.Errorf("output lacks %q:\n%s", want, out) + } + } + if strings.Contains(out, "same.p_id") { + t.Errorf("only-diff output shows a matching relationship:\n%s", out) + } +} diff --git a/internal/compare/snapshot.go b/internal/compare/snapshot.go index 12c358f..23b7cb4 100644 --- a/internal/compare/snapshot.go +++ b/internal/compare/snapshot.go @@ -14,16 +14,22 @@ import ( "github.com/goccy/go-yaml" "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" ) // SnapshotKind identifies a table-counts snapshot file. const SnapshotKind = "seedstorm.table-counts" -// SnapshotVersion is the snapshot file format written by EncodeSnapshot. +// SnapshotVersion is the snapshot file format of counts only. Version 2 adds +// relationships and is written only when there are some, so older binaries +// keep reading counts-only files. const SnapshotVersion = 1 +// snapshotVersionRelationships is the first version with relationships. +const snapshotVersionRelationships = 2 + // supportedSnapshotVersions lists every version ParseSnapshot reads. -var supportedSnapshotVersions = []int64{1} +var supportedSnapshotVersions = []int64{1, 2} // Snapshot file formats accepted by EncodeSnapshot. const ( @@ -33,15 +39,19 @@ const ( // snapshotFields is the order fields are written in, and the set ParseSnapshot // accepts at the top level. -var snapshotFields = []string{"kind", "version", "label", "dbType", "countMode", "takenAt", "tables"} +var snapshotFields = []string{"kind", "version", "label", "dbType", "countMode", "takenAt", "tables", "relationships"} // EncodeSnapshot writes s as a versioned table-counts file in format "json" or // "yaml". Tables are sorted by name so two snapshots of one database diff line // by line. Unknown row counts and sizes are written as db.UnknownCount (-1). func EncodeSnapshot(s Snapshot, format string) ([]byte, error) { + version := SnapshotVersion + if len(s.Relationships) > 0 { + version = snapshotVersionRelationships + } doc := yaml.MapSlice{ {Key: "kind", Value: SnapshotKind}, - {Key: "version", Value: SnapshotVersion}, + {Key: "version", Value: version}, {Key: "label", Value: s.Label}, {Key: "dbType", Value: s.DBType}, {Key: "countMode", Value: string(s.CountMode)}, @@ -63,6 +73,9 @@ func EncodeSnapshot(s Snapshot, format string) ([]byte, error) { tables = append(tables, yaml.MapItem{Key: name, Value: entry}) } doc = append(doc, yaml.MapItem{Key: "tables", Value: tables}) + if len(s.Relationships) > 0 { + doc = append(doc, yaml.MapItem{Key: "relationships", Value: s.Relationships}) + } switch strings.ToLower(strings.TrimSpace(format)) { case FormatYAML, "yml": @@ -174,6 +187,9 @@ func ParseSnapshot(data []byte) (Snapshot, error) { if !ok || !slices.Contains(supportedSnapshotVersions, n) { return snap, fmt.Errorf("unsupported snapshot version %v (supported: %s)", version, versionList()) } + if _, has := top["relationships"]; has && n < snapshotVersionRelationships { + return snap, fmt.Errorf("relationships need version %d of the snapshot format (this file says %d)", snapshotVersionRelationships, n) + } } else if _, present := top["version"]; present { return snap, fmt.Errorf("snapshot has a version but no kind: add kind: %s", SnapshotKind) } @@ -213,6 +229,17 @@ func ParseSnapshot(data []byte) (Snapshot, error) { return snap, fmt.Errorf("snapshot takenAt must be a timestamp, got %s", describe(v)) } + if raw, has := top["relationships"]; has { + if minimal { + return snap, fmt.Errorf("relationships need kind: %s and version: %d at the top", SnapshotKind, snapshotVersionRelationships) + } + shapes, err := parseRelationships(raw) + if err != nil { + return snap, err + } + snap.Relationships = shapes + } + tablesValue, present := top["tables"] if !present || tablesValue == nil { return snap, errors.New("snapshot has no tables") @@ -380,3 +407,25 @@ func versionList() string { } return strings.Join(parts, ", ") } + +// parseRelationships reads the relationships section through JSON, which the +// shape type describes field by field. +func parseRelationships(raw any) ([]relations.Shape, error) { + if raw == nil { + return nil, nil + } + data, err := json.Marshal(raw) + if err != nil { + return nil, fmt.Errorf("snapshot relationships: %w", err) + } + var shapes []relations.Shape + if err := json.Unmarshal(data, &shapes); err != nil { + return nil, fmt.Errorf("snapshot relationships must be a list of foreign-key shapes: %w", err) + } + for i, s := range shapes { + if s.Child == "" || s.Column == "" || s.Parent == "" { + return nil, fmt.Errorf("snapshot relationship %d needs child, column and parent", i+1) + } + } + return shapes, nil +} diff --git a/internal/compare/snapshot_test.go b/internal/compare/snapshot_test.go index bc9aa49..e2007ae 100644 --- a/internal/compare/snapshot_test.go +++ b/internal/compare/snapshot_test.go @@ -9,6 +9,7 @@ import ( "time" "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" ) func sampleSnapshot() Snapshot { @@ -226,9 +227,9 @@ func TestParseSnapshot_RejectsWithReadableErrors(t *testing.T) { rejects("bad estimated", "kind: seedstorm.table-counts\nversion: 1\ntables:\n users: {rows: 1, estimated: maybe}\n", "estimated must be true or false") rejects("wrong kind", "kind: seedstorm.profile\nversion: 1\ntables: {a: 1}", `kind is "seedstorm.profile", want "seedstorm.table-counts"`) rejects("missing rows", "kind: seedstorm.table-counts\nversion: 1\ntables:\n users: {bytes: 3}\n", `table "users": missing rows`) - rejects("missing version", "kind: seedstorm.table-counts\ntables: {a: 1}", "snapshot has no version (supported: 1)") + rejects("missing version", "kind: seedstorm.table-counts\ntables: {a: 1}", "snapshot has no version (supported: 1, 2)") rejects("unknown table field", "kind: seedstorm.table-counts\nversion: 1\ntables:\n users: {rows: 1, size: 3}\n", `table "users": unexpected field "size"`) - rejects("unsupported version", "kind: seedstorm.table-counts\nversion: 2\ntables: {a: 1}", "unsupported snapshot version 2 (supported: 1)") + rejects("unsupported version", "kind: seedstorm.table-counts\nversion: 3\ntables: {a: 1}", "unsupported snapshot version 3 (supported: 1, 2)") rejects("random json", `{"hello": "world"}`, `unexpected field(s) "hello"`) rejects("version without kind", "version: 1\ntables: {a: 1}", "add kind: seedstorm.table-counts") rejects("compare report json", `{"source": {}, "target": {}, "rows": [], "totals": {}}`, `unexpected field(s) "rows", "source", "target", "totals"`) @@ -264,3 +265,50 @@ func TestParseSnapshot_DiffsLikeATakenSnapshot(t *testing.T) { t.Errorf("unknown sizes leaked into totals: %+v", r.Totals) } } + +// Relationship shapes travel in snapshot version 2. A counts-only snapshot is +// still written as version 1, so older seedstorm binaries keep reading it. +func TestSnapshot_RelationshipsNeedVersion2AndRoundTrip(t *testing.T) { + base := Snapshot{Label: "app@db", DBType: "pgx", CountMode: CountExact, Tables: map[string]TableStat{"teams": {Rows: 5, Bytes: 100}, "players": {Rows: 17, Bytes: 200}}} + + for _, format := range []string{FormatYAML, FormatJSON} { + data, err := EncodeSnapshot(base, format) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), "version: 1") && !strings.Contains(string(data), `"version": 1`) { + t.Fatalf("counts-only %s snapshot is not version 1:\n%s", format, data) + } + + withShapes := base + withShapes.Relationships = []relations.Shape{{ + Child: "players", Column: "team_id", Parent: "teams", ParentColumn: "id", + Parents: 5, Children: 15, NullRows: 2, ParentsWithChildren: 4, ZeroShare: 0.2, NullShare: 0.117647, + Min: 1, Max: 10, Avg: 3.75, P50: 1, P95: 10, Indexed: true, Outcome: db.OutcomeOK, + Histogram: []db.DegreeBucket{{Min: 1, Max: 1, Parents: 2, Children: 2}, {Min: 3, Max: 3, Parents: 1, Children: 3}, {Min: 10, Max: 10, Parents: 1, Children: 10}}, + }} + data, err = EncodeSnapshot(withShapes, format) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), "version: 2") && !strings.Contains(string(data), `"version": 2`) { + t.Fatalf("%s snapshot with relationships is not version 2:\n%s", format, data) + } + back, err := ParseSnapshot(data) + if err != nil { + t.Fatalf("parse %s: %v\n%s", format, err, data) + } + if len(back.Relationships) != 1 || !reflect.DeepEqual(back.Relationships[0], withShapes.Relationships[0]) { + t.Fatalf("%s round trip:\n got %+v\nwant %+v", format, back.Relationships, withShapes.Relationships) + } + } + + _, err := ParseSnapshot([]byte("kind: seedstorm.table-counts\nversion: 1\ntables: {users: 1}\nrelationships: []\n")) + if err == nil || !strings.Contains(err.Error(), "version 2") { + t.Fatalf("relationships in a version 1 file: %v", err) + } + _, err = ParseSnapshot([]byte("kind: seedstorm.table-counts\nversion: 2\ntables: {users: 1}\nrelationships:\n - {column: team_id}\n")) + if err == nil || !strings.Contains(err.Error(), "child") { + t.Fatalf("a relationship without its child table: %v", err) + } +} diff --git a/internal/db/relations.go b/internal/db/relations.go new file mode 100644 index 0000000..37bc59a --- /dev/null +++ b/internal/db/relations.go @@ -0,0 +1,140 @@ +package db + +import ( + "context" + "fmt" + "strings" +) + +// DegreeBucket counts parents whose number of children falls in [Min, Max]: +// exact for small degrees, powers of two above. +type DegreeBucket struct { + Min int64 `json:"min" yaml:"min"` + Max int64 `json:"max" yaml:"max"` + Parents int64 `json:"parents" yaml:"parents"` + Children int64 `json:"children" yaml:"children"` +} + +// exactDegrees is the largest degree with a bucket of its own. +const exactDegrees = 16 + +// degreeBucketExpr maps a degree c to its bucket's upper bound, the same SQL +// on both engines (no LOG2 in common). +func degreeBucketExpr() string { + var b strings.Builder + b.WriteString("CASE") + for d := 1; d <= exactDegrees; d++ { + fmt.Fprintf(&b, " WHEN c = %d THEN %d", d, d) + } + for upper := int64(exactDegrees * 2); upper <= 1<<40; upper *= 2 { + fmt.Fprintf(&b, " WHEN c <= %d THEN %d", upper, upper) + } + b.WriteString(" ELSE c END") + return b.String() +} + +// DegreeHistogram aggregates children per parent of child.column in the +// database (no rows travel): one row per bucket, plus how many child rows have +// a NULL key. It is a full pass over the key: callers gate unindexed keys. +func DegreeHistogram(ctx context.Context, q Querier, dbType, child, column string) (buckets []DegreeBucket, nulls int64, err error) { + t, c := QuoteIdent(child, dbType), QuoteIdent(column, dbType) + //nolint:gosec // identifiers are quoted + query := fmt.Sprintf(` + SELECT bucket, COUNT(*), MIN(c), MAX(c), SUM(c) + FROM (SELECT %s AS bucket, c FROM (SELECT %s AS k, COUNT(*) AS c FROM %s WHERE %s IS NOT NULL GROUP BY %s) per_parent) b + GROUP BY bucket ORDER BY bucket`, degreeBucketExpr(), c, t, c, c) + rows, err := q.QueryContext(ctx, query) + if err != nil { + return nil, 0, err + } + for rows.Next() { + var upper, parents, lo, hi, sum int64 + if err := rows.Scan(&upper, &parents, &lo, &hi, &sum); err != nil { + rows.Close() + return nil, 0, err + } + buckets = append(buckets, DegreeBucket{Min: lo, Max: hi, Parents: parents, Children: sum}) + } + rows.Close() + if err := rows.Err(); err != nil { + return nil, 0, err + } + //nolint:gosec // identifiers are quoted + err = q.QueryRowContext(ctx, fmt.Sprintf(`SELECT COUNT(*) - COUNT(%s) FROM %s`, c, t)).Scan(&nulls) + return buckets, nulls, err +} + +// LeadingIndexed reports whether column is the first column of any index on +// table (primary key, unique and partial indexes included): only then can the +// database group by it without a full table scan. +func LeadingIndexed(ctx context.Context, q Querier, dbType, table, column string) (bool, error) { + var n int + var err error + if dbType == "mysql" { + err = q.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM information_schema.STATISTICS + WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = ? AND COLUMN_NAME = ? AND SEQ_IN_INDEX = 1`, table, column).Scan(&n) + } else { + err = q.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM pg_index i + JOIN pg_class t ON t.oid = i.indrelid + JOIN pg_namespace ns ON ns.oid = t.relnamespace + JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = i.indkey[0] + WHERE ns.nspname = 'public' AND t.relname = $1 AND a.attname = $2`, table, column).Scan(&n) + } + return n > 0, err +} + +// DegreeEstimate is what planner statistics say about a key's degrees, without +// reading the table. Unknown values are negative. +type DegreeEstimate struct { + ChildRows int64 + ParentRows int64 + Distinct int64 + NullFraction float64 + // MaxFraction is the most common key value's share of child rows. + MaxFraction float64 +} + +// EstimateDegrees reads statistics: Postgres pg_stats (null_frac, n_distinct, +// most common frequencies); MySQL index cardinality (no null share or maximum). +func EstimateDegrees(ctx context.Context, q Querier, dbType, child, column, parent string) (DegreeEstimate, error) { + est := DegreeEstimate{ChildRows: -1, ParentRows: -1, Distinct: -1, NullFraction: -1, MaxFraction: -1} + if dbType == "mysql" { + if err := q.QueryRowContext(ctx, `SELECT COALESCE(TABLE_ROWS, -1) FROM information_schema.TABLES WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = ?`, child).Scan(&est.ChildRows); err != nil { + return est, err + } + if err := q.QueryRowContext(ctx, `SELECT COALESCE(TABLE_ROWS, -1) FROM information_schema.TABLES WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = ?`, parent).Scan(&est.ParentRows); err != nil { + return est, err + } + err := q.QueryRowContext(ctx, ` + SELECT COALESCE(MAX(CARDINALITY), -1) FROM information_schema.STATISTICS + WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = ? AND COLUMN_NAME = ? AND SEQ_IN_INDEX = 1`, child, column).Scan(&est.Distinct) + return est, err + } + rowsOf := `SELECT COALESCE((SELECT GREATEST(c.reltuples, 0)::bigint FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = 'public' AND c.relname = $1), -1)` + if err := q.QueryRowContext(ctx, rowsOf, child).Scan(&est.ChildRows); err != nil { + return est, err + } + if err := q.QueryRowContext(ctx, rowsOf, parent).Scan(&est.ParentRows); err != nil { + return est, err + } + var nDistinct float64 + rows, err := q.QueryContext(ctx, ` + SELECT null_frac, n_distinct, COALESCE((SELECT MAX(f) FROM unnest(most_common_freqs) f), 0) + FROM pg_stats WHERE schemaname = 'public' AND tablename = $1 AND attname = $2`, child, column) + if err != nil { + return est, err + } + defer rows.Close() + if rows.Next() { + if err := rows.Scan(&est.NullFraction, &nDistinct, &est.MaxFraction); err != nil { + return est, err + } + if nDistinct < 0 && est.ChildRows >= 0 { + nDistinct = -nDistinct * float64(est.ChildRows) + } + est.Distinct = int64(nDistinct) + } + return est, rows.Err() +} diff --git a/internal/relations/scan.go b/internal/relations/scan.go new file mode 100644 index 0000000..55ecfee --- /dev/null +++ b/internal/relations/scan.go @@ -0,0 +1,215 @@ +package relations + +import ( + "context" + "database/sql" + "sort" + "strings" + "sync" + "time" + + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/safego" + "github.com/AxeForging/seedstorm/internal/schema" +) + +// Mode selects how shapes are measured. +type Mode string + +const ( + // Exact aggregates the key in the database: one pass over it. + Exact Mode = "exact" + // Estimate reads planner statistics: instant, approximate, partial. + Estimate Mode = "estimate" +) + +// Outcomes beyond db.ReadOutcome. +const ( + OutcomeSkippedUnindexed db.ReadOutcome = "skipped: unindexed" + OutcomeEstimated db.ReadOutcome = "estimated" +) + +// DefaultStatementTimeout bounds one relationship's exact scan. +const DefaultStatementTimeout = 60 * time.Second + +// DefaultLargeRows is the child table size warned about before an exact scan. +const DefaultLargeRows = 5_000_000 + +// Options configure a scan. +type Options struct { + Mode Mode + // Limits apply to each exact scan; a zero StatementTimeout means + // DefaultStatementTimeout. Concurrency defaults to 2. + Limits db.ReadLimits + // IncludeUnindexed scans keys that lead no index (full table scans). + IncludeUnindexed bool + // LargeRows marks child tables above it as Large (0: DefaultLargeRows). + LargeRows int64 + // OnEdge is called after each relationship, one call at a time. + OnEdge func(done, total int, s Shape) +} + +// Edges lists the schema's single-column foreign keys (self-references +// included), sorted by child table and column. +func Edges(sc *schema.Schema) []Shape { + var out []Shape + for childName, t := range sc.Tables { + for colName, col := range t.Columns { + parent, parentCol := splitFK(col.FK) + if parent == "" { + continue + } + out = append(out, Shape{Child: childName, Column: colName, Parent: parent, ParentColumn: parentCol, SelfRef: parent == childName}) + } + } + sort.Slice(out, func(i, j int) bool { + if out[i].Child != out[j].Child { + return out[i].Child < out[j].Child + } + return out[i].Column < out[j].Column + }) + return out +} + +// Scan measures every relationship of sc, cheapest first, read-only. It +// returns every edge with an outcome: finished edges are kept when ctx is +// cancelled, and one edge timing out does not stop the others. +func Scan(ctx context.Context, conn *sql.DB, dbType string, sc *schema.Schema, opts Options) ([]Shape, error) { + if opts.Mode == "" { + opts.Mode = Exact + } + if opts.Limits.StatementTimeout == 0 { + opts.Limits.StatementTimeout = DefaultStatementTimeout + } + if opts.Limits.LockTimeout == 0 { + opts.Limits.LockTimeout = db.DefaultCountLimits.LockTimeout + } + concurrency := max(opts.Limits.Concurrency, 2) + if opts.Limits.Concurrency == 1 { + concurrency = 1 + } + opts.Limits.Concurrency = 1 // each worker runs one statement at a time + if opts.LargeRows <= 0 { + opts.LargeRows = DefaultLargeRows + } + + edges := Edges(sc) + estimates, _ := db.GetEstimatedRowCounts(ctx, conn, dbType) + size := func(table string) int64 { + if n, ok := estimates[table]; ok && n >= 0 { + return n + } + return 1 << 50 + } + sort.SliceStable(edges, func(i, j int) bool { return size(edges[i].Child) < size(edges[j].Child) }) + + out := make([]Shape, len(edges)) + var mu sync.Mutex + done := 0 + work := make(chan int) + var wg sync.WaitGroup + for w := 0; w < concurrency; w++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := range work { + edge := edges[i] + var shape Shape + err := safego.Run("relationship "+edge.Child+"."+edge.Column, func() error { + shape = measure(ctx, conn, dbType, edge, opts) + return nil + }) + if err != nil { + shape = edge + shape.Outcome, shape.Detail = db.OutcomeFailed, err.Error() + } + shape.Large = size(edge.Child) > opts.LargeRows && size(edge.Child) < 1<<50 + mu.Lock() + out[i] = shape + done++ + if opts.OnEdge != nil { + opts.OnEdge(done, len(edges), shape) + } + mu.Unlock() + } + }() + } + for i := range edges { + work <- i + } + close(work) + wg.Wait() + sort.SliceStable(out, func(i, j int) bool { + if out[i].Child != out[j].Child { + return out[i].Child < out[j].Child + } + return out[i].Column < out[j].Column + }) + return out, nil +} + +// measure scans one relationship. +func measure(ctx context.Context, conn *sql.DB, dbType string, edge Shape, opts Options) Shape { + withEdge := func(s Shape) Shape { + s.Child, s.Column, s.Parent, s.ParentColumn, s.SelfRef = edge.Child, edge.Column, edge.Parent, edge.ParentColumn, edge.SelfRef + return s + } + if ctx.Err() != nil { + s := withEdge(Shape{}) + s.Outcome, s.Detail = db.OutcomeCancelled, ctx.Err().Error() + return s + } + var indexed bool + _ = db.ReadOnce(ctx, conn, dbType, db.DefaultCountLimits, func(ctx context.Context, q db.Querier) (err error) { + indexed, err = db.LeadingIndexed(ctx, q, dbType, edge.Child, edge.Column) + return err + }) + + estimate := func(outcome db.ReadOutcome, detail string) Shape { + var est db.DegreeEstimate + err := db.ReadOnce(ctx, conn, dbType, db.DefaultCountLimits, func(ctx context.Context, q db.Querier) (err error) { + est, err = db.EstimateDegrees(ctx, q, dbType, edge.Child, edge.Column, edge.Parent) + return err + }) + s := withEdge(shapeFromEstimate(est)) + s.Indexed = indexed + s.Outcome, s.Detail = outcome, detail + if err != nil { + s.Outcome, s.Detail = db.ReadOutcomeOf(ctx, err), err.Error() + } + return s + } + if opts.Mode == Estimate { + return estimate(OutcomeEstimated, "") + } + if !indexed && !opts.IncludeUnindexed { + return estimate(OutcomeSkippedUnindexed, "the key leads no index: an exact scan reads the whole table (numbers are estimates)") + } + + var buckets []db.DegreeBucket + var nulls, parents int64 + err := db.ReadOnce(ctx, conn, dbType, opts.Limits, func(ctx context.Context, q db.Querier) (err error) { + if buckets, nulls, err = db.DegreeHistogram(ctx, q, dbType, edge.Child, edge.Column); err != nil { + return err + } + //nolint:gosec // identifier is quoted + return q.QueryRowContext(ctx, "SELECT COUNT(*) FROM "+db.QuoteIdent(edge.Parent, dbType)).Scan(&parents) + }) + s := withEdge(shapeFromHistogram(buckets, parents, nulls)) + s.Indexed = indexed + s.Outcome = db.OutcomeOK + if err != nil { + s = withEdge(Shape{Indexed: indexed}) + s.Outcome, s.Detail = db.ReadOutcomeOf(ctx, err), err.Error() + } + return s +} + +// splitFK splits a schema FK "table.column". +func splitFK(fk string) (string, string) { + parts := strings.SplitN(fk, ".", 2) + if len(parts) != 2 { + return "", "" + } + return parts[0], parts[1] +} diff --git a/internal/relations/shape.go b/internal/relations/shape.go new file mode 100644 index 0000000..fb9c939 --- /dev/null +++ b/internal/relations/shape.go @@ -0,0 +1,112 @@ +// Package relations measures relationship shapes: for every foreign key, how +// many children each parent has (minimum, maximum, average, percentiles and a +// histogram), how many parents have none, and how many keys are NULL. +package relations + +import ( + "math" + + "github.com/AxeForging/seedstorm/internal/db" +) + +// Shape is one foreign key's degree distribution. Estimated shapes come from +// planner statistics; unknown numbers are -1. +type Shape struct { + Child string `json:"child" yaml:"child"` + Column string `json:"column" yaml:"column"` + Parent string `json:"parent" yaml:"parent"` + ParentColumn string `json:"parentColumn" yaml:"parentColumn"` + SelfRef bool `json:"selfRef,omitempty" yaml:"selfRef,omitempty"` + + Parents int64 `json:"parents" yaml:"parents"` + Children int64 `json:"children" yaml:"children"` + NullRows int64 `json:"nullRows" yaml:"nullRows"` + ParentsWithChildren int64 `json:"parentsWithChildren" yaml:"parentsWithChildren"` + ZeroShare float64 `json:"zeroShare" yaml:"zeroShare"` + NullShare float64 `json:"nullShare" yaml:"nullShare"` + Min int64 `json:"min" yaml:"min"` + Max int64 `json:"max" yaml:"max"` + Avg float64 `json:"avg" yaml:"avg"` + P50 int64 `json:"p50" yaml:"p50"` + P95 int64 `json:"p95" yaml:"p95"` + + Histogram []db.DegreeBucket `json:"histogram,omitempty" yaml:"histogram,omitempty"` + Estimated bool `json:"estimated,omitempty" yaml:"estimated,omitempty"` + // Indexed reports the key leads an index (exact scans are cheap). + Indexed bool `json:"indexed" yaml:"indexed"` + // Large marks a child table above Options.LargeRows (warned before exact scans). + Large bool `json:"large,omitempty" yaml:"large,omitempty"` + Outcome db.ReadOutcome `json:"outcome" yaml:"outcome"` + Detail string `json:"detail,omitempty" yaml:"detail,omitempty"` +} + +// shapeFromHistogram derives a shape from bucketed degrees, the parent count +// and NULL child rows. +func shapeFromHistogram(buckets []db.DegreeBucket, parents, nulls int64) Shape { + s := Shape{Parents: parents, NullRows: nulls, Histogram: buckets, Min: 0, Max: 0} + for i, b := range buckets { + s.ParentsWithChildren += b.Parents + s.Children += b.Children + if i == 0 || b.Min < s.Min { + s.Min = b.Min + } + s.Max = max(s.Max, b.Max) + } + if s.ParentsWithChildren > 0 { + s.Avg = float64(s.Children) / float64(s.ParentsWithChildren) + s.P50 = percentile(buckets, s.ParentsWithChildren, 0.50) + s.P95 = percentile(buckets, s.ParentsWithChildren, 0.95) + } + if parents > 0 { + s.ZeroShare = float64(max(parents-s.ParentsWithChildren, 0)) / float64(parents) + } + if total := s.Children + nulls; total > 0 { + s.NullShare = float64(nulls) / float64(total) + } + return s +} + +// percentile walks the buckets to the q-th parent; within a bucket spanning +// several degrees it reports the bucket's largest observed degree. +func percentile(buckets []db.DegreeBucket, parents int64, q float64) int64 { + rank := int64(math.Ceil(q * float64(parents))) + var seen int64 + for _, b := range buckets { + seen += b.Parents + if seen >= rank { + return b.Max + } + } + if len(buckets) == 0 { + return 0 + } + return buckets[len(buckets)-1].Max +} + +// shapeFromEstimate derives what statistics can tell: average degree, the +// share of parents without children, and a maximum from the most common key. +func shapeFromEstimate(e db.DegreeEstimate) Shape { + s := Shape{Estimated: true, Parents: e.ParentRows, Min: -1, Max: -1, P50: -1, P95: -1, NullRows: -1, NullShare: -1, ZeroShare: -1} + if e.ChildRows >= 0 && e.NullFraction >= 0 { + s.NullRows = int64(math.Round(float64(e.ChildRows) * e.NullFraction)) + s.NullShare = e.NullFraction + } + nonNull := e.ChildRows + if s.NullRows > 0 { + nonNull -= s.NullRows + } + s.Children = nonNull + if e.Distinct > 0 { + s.ParentsWithChildren = e.Distinct + if nonNull >= 0 { + s.Avg = float64(nonNull) / float64(e.Distinct) + } + if e.ParentRows > 0 { + s.ZeroShare = math.Max(0, 1-float64(e.Distinct)/float64(e.ParentRows)) + } + } + if e.MaxFraction > 0 && e.ChildRows > 0 { + s.Max = int64(math.Round(e.MaxFraction * float64(e.ChildRows))) + } + return s +} diff --git a/internal/relations/shape_test.go b/internal/relations/shape_test.go new file mode 100644 index 0000000..71ab88e --- /dev/null +++ b/internal/relations/shape_test.go @@ -0,0 +1,62 @@ +package relations + +import ( + "math" + "testing" + + "github.com/AxeForging/seedstorm/internal/db" +) + +func TestShapeFromHistogram_ComputesDegreesSharesAndPercentiles(t *testing.T) { + // Degrees 1, 1, 3, 10 over 5 parents, 2 NULL child rows. + buckets := []db.DegreeBucket{ + {Min: 1, Max: 1, Parents: 2, Children: 2}, + {Min: 3, Max: 3, Parents: 1, Children: 3}, + {Min: 10, Max: 10, Parents: 1, Children: 10}, + } + s := shapeFromHistogram(buckets, 5, 2) + if s.Parents != 5 || s.Children != 15 || s.NullRows != 2 || s.ParentsWithChildren != 4 { + t.Fatalf("counts = %+v", s) + } + if s.Min != 1 || s.Max != 10 || s.Avg != 3.75 || s.P50 != 1 || s.P95 != 10 { + t.Fatalf("degrees = %+v", s) + } + if math.Abs(s.ZeroShare-0.2) > 1e-9 || math.Abs(s.NullShare-2.0/17.0) > 1e-9 { + t.Fatalf("shares = %+v", s) + } +} + +// Above the exact buckets a percentile is the bucket's largest observed degree, +// never beyond the maximum. +func TestShapeFromHistogram_PercentileInAPowerOfTwoBucket(t *testing.T) { + buckets := []db.DegreeBucket{ + {Min: 1, Max: 1, Parents: 90, Children: 90}, + {Min: 40, Max: 60, Parents: 10, Children: 500}, + } + s := shapeFromHistogram(buckets, 100, 0) + if s.P50 != 1 || s.P95 != 60 || s.Max != 60 || s.ZeroShare != 0 { + t.Fatalf("shape = %+v", s) + } +} + +func TestShapeFromHistogram_NoChildren(t *testing.T) { + s := shapeFromHistogram(nil, 7, 3) + if s.ParentsWithChildren != 0 || s.ZeroShare != 1 || s.Avg != 0 || s.NullShare != 1 { + t.Fatalf("shape = %+v", s) + } +} + +// Postgres most_common_freqs are shares of all rows (NULLs included). +func TestShapeFromEstimate(t *testing.T) { + s := shapeFromEstimate(db.DegreeEstimate{ChildRows: 1000, ParentRows: 100, Distinct: 80, NullFraction: 0.1, MaxFraction: 0.05}) + if !s.Estimated || s.Parents != 100 || s.ParentsWithChildren != 80 || math.Abs(s.ZeroShare-0.2) > 1e-9 { + t.Fatalf("estimate = %+v", s) + } + if math.Abs(s.Avg-900.0/80.0) > 1e-9 || s.Max != 50 || s.NullRows != 100 { + t.Fatalf("estimate degrees = %+v", s) + } + unknown := shapeFromEstimate(db.DegreeEstimate{ChildRows: -1, ParentRows: -1, Distinct: -1, NullFraction: -1, MaxFraction: -1}) + if unknown.Avg != 0 || unknown.Max != -1 { + t.Fatalf("unknown estimate = %+v", unknown) + } +} diff --git a/internal/seeder/relationships.go b/internal/seeder/relationships.go new file mode 100644 index 0000000..238bf5f --- /dev/null +++ b/internal/seeder/relationships.go @@ -0,0 +1,78 @@ +package seeder + +import ( + "context" + "errors" + "fmt" + "sync" + + "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/faker" + "github.com/AxeForging/seedstorm/internal/relations" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/safego" +) + +// ErrNoRelationships: a snapshot endpoint was asked for shapes it does not hold. +var ErrNoRelationships = errors.New("the snapshot file has no relationships (take it with --relationships)") + +// Shapes returns the endpoint's relationship shapes: the snapshot's when it +// is one, otherwise a read-only scan (introspecting first when the endpoint +// carries no schema). +func (e Endpoint) Shapes(ctx context.Context, opts relations.Options) ([]relations.Shape, error) { + if e.Snapshot != nil { + if len(e.Snapshot.Relationships) == 0 { + return nil, ErrNoRelationships + } + return e.Snapshot.Relationships, nil + } + if e.Conn == nil { + return nil, errors.New("no database connection or snapshot") + } + sc := e.Schema + if sc == nil { + tables, err := db.IntrospectConn(ctx, e.Conn, e.DBType, nil) + if err != nil { + return nil, runerr.At(runerr.PhaseIntrospect, "", err) + } + sc = faker.BuildSchema(e.DBType, tables) + } + shapes, err := relations.Scan(ctx, e.Conn, e.DBType, sc, opts) + if err != nil { + return nil, fmt.Errorf("relationship scan: %w", err) + } + return shapes, nil +} + +// CompareShapes scans (or reads) both sides at once and diffs them. onEdge +// reports progress per side. +func CompareShapes(ctx context.Context, source, target Endpoint, opts relations.Options, onEdge func(side string, done, total int, s relations.Shape)) ([]compare.ShapeDrift, error) { + sideOpts := func(side string) relations.Options { + o := opts + if onEdge != nil { + o.OnEdge = func(done, total int, s relations.Shape) { onEdge(side, done, total, s) } + } + return o + } + var src, tgt []relations.Shape + var srcErr, tgtErr error + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + srcErr = safego.Run("relationships source", func() (err error) { src, err = source.Shapes(ctx, sideOpts("source")); return err }) + }() + go func() { + defer wg.Done() + tgtErr = safego.Run("relationships target", func() (err error) { tgt, err = target.Shapes(ctx, sideOpts("target")); return err }) + }() + wg.Wait() + if srcErr != nil { + return nil, runerr.OnSide(runerr.SideSource, srcErr) + } + if tgtErr != nil { + return nil, runerr.OnSide(runerr.SideTarget, tgtErr) + } + return compare.DiffShapes(src, tgt), nil +} From ebc3b0ab721d9d8b43872d3abb5467985b133c68 Mon Sep 17 00:00:00 2001 From: Lucas Machado Date: Thu, 17 Sep 2026 17:21:54 +0200 Subject: [PATCH 04/20] feat: analyze and compare relationships in the web UI - workspace job measures shapes; graph edges show avg/max as each finishes - compare page relationships step with per-key drift - counts exports include relationships on request --- e2e/support/selectors.ts | 14 ++ e2e/tests/relationships.spec.ts | 68 ++++++++ internal/compare/report_snapshot.go | 10 ++ internal/web/handlers_relationships.go | 190 +++++++++++++++++++++ internal/web/handlers_snapshots.go | 13 +- internal/web/server.go | 2 + internal/web/session.go | 4 + internal/web/shapes.go | 85 +++++++++ internal/web/shapes_test.go | 101 +++++++++++ internal/web/static/app.js | 138 ++++++++++++++- internal/web/static/compare.css | 39 +++++ internal/web/static/compare.js | 86 +++++++++- internal/web/static/style.css | 3 + internal/web/templates/compare.html.tmpl | 16 ++ internal/web/templates/workspace.html.tmpl | 2 + 15 files changed, 760 insertions(+), 11 deletions(-) create mode 100644 e2e/tests/relationships.spec.ts create mode 100644 internal/web/handlers_relationships.go create mode 100644 internal/web/shapes.go create mode 100644 internal/web/shapes_test.go diff --git a/e2e/support/selectors.ts b/e2e/support/selectors.ts index 92cf7a3..d6c40ee 100644 --- a/e2e/support/selectors.ts +++ b/e2e/support/selectors.ts @@ -131,6 +131,20 @@ export const sel = { preview: "snapshot-preview", download: "snapshot-download", compare: "snapshot-compare", + relationships: "snapshot-relationships", + }, + relationships: { + open: "ws-relationships", + dialog: "shapes-dialog", + unindexed: "shapes-unindexed", + start: "shapes-start", + status: "ws-shapes-status", + section: "cmp-shapes", + compareUnindexed: "cmp-shapes-unindexed", + compareRun: "cmp-shapes-run", + summary: "cmp-shapes-summary", + row: "cmp-shape-row", + exportInclude: "cmp-export-relationships", }, tuning: { open: "ws-recommend", diff --git a/e2e/tests/relationships.spec.ts b/e2e/tests/relationships.spec.ts new file mode 100644 index 0000000..942f7ba --- /dev/null +++ b/e2e/tests/relationships.spec.ts @@ -0,0 +1,68 @@ +// Relationship shapes: measure children per parent on the workspace, carry +// them in a counts file, and compare them between two databases. +import { DB } from "../support/db.helpers"; +import { sel } from "../support/selectors"; +import { connectPostgres, expect, openWorkspace, pgConnectionLabel, test } from "../support/test.fixture"; + +const rel = sel.relationships; +const snap = sel.snapshot; +const cmp = sel.compare; + +// orders.customer_id: 3,400 orders over 1,200 customers, dealt round robin → +// 1,000 customers with 3 orders and 200 with 2 (avg 2.83, p95 3, max 3). The +// key has no index on Postgres, so it is only scanned when asked. +test("analyze relationships, export them and compare them with another database", async ({ page }) => { + await connectPostgres(page, DB.tgt); + await connectPostgres(page, DB.src); + await openWorkspace(page, 2); + + await test.step("an unindexed key is estimated unless included", async () => { + await page.getByTestId(rel.open).click(); + await page.getByTestId(rel.dialog).getByTestId(rel.start).click(); + await expect(page.getByTestId(rel.status)).toContainText(/1 relationship measured · 1 estimated/, { timeout: 30_000 }); + }); + + await test.step("an exact scan measures it", async () => { + await page.getByTestId(rel.open).click(); + const dialog = page.getByTestId(rel.dialog); + await dialog.getByTestId(rel.unindexed).check(); + await dialog.getByTestId(rel.start).click(); + await expect(page.getByTestId(rel.status)).toHaveText("1 relationship measured", { timeout: 30_000 }); + }); + + await test.step("the counts file includes them only when asked", async () => { + await page.getByTestId(snap.open).click(); + const dialog = page.getByTestId(snap.dialog); + const preview = dialog.getByTestId(snap.preview); + await expect(preview).toContainText("kind: seedstorm.table-counts", { timeout: 30_000 }); + await expect(preview).not.toContainText("relationships:"); + await dialog.getByTestId(snap.relationships).check(); + await expect(preview).toContainText("relationships:"); + await expect(preview).toContainText(/child: orders[\s\S]*max: 3/); + await dialog.getByRole("button", { name: "Close" }).click(); + }); + + await test.step("compare shows the drift and exports it", async () => { + await page.goto("/compare"); + await page.getByTestId(cmp.source).selectOption({ label: pgConnectionLabel(DB.src) }); + await page.getByTestId(cmp.target).selectOption({ label: pgConnectionLabel(DB.tgt) }); + await page.getByTestId(cmp.run).click(); + await expect(page.getByTestId(cmp.results)).toBeVisible({ timeout: 30_000 }); + await expect(page.getByTestId(cmp.gauge).first()).toBeVisible(); + + const section = page.getByTestId(rel.section); + await section.getByTestId(rel.compareUnindexed).check(); + await section.getByTestId(rel.compareRun).click(); + const row = section.getByTestId(rel.row).filter({ hasText: "orders.customer_id" }); + await expect(row).toContainText("2.83 / 3 / 3 · 0%", { timeout: 30_000 }); + await expect(row).toHaveAttribute("data-status", "differs"); + await expect(section.getByTestId(rel.summary)).toContainText("1 relationship · 0 same · 1 differ"); + + await page.getByTestId(cmp.exportOpen).click(); + const dialog = page.getByTestId(cmp.exportDialog); + await expect(dialog.getByTestId(cmp.exportPreview)).toContainText("kind: seedstorm.table-counts"); + await expect(dialog.getByTestId(cmp.exportPreview)).not.toContainText("relationships:"); + await dialog.getByTestId(rel.exportInclude).check(); + await expect(dialog.getByTestId(cmp.exportPreview)).toContainText(/relationships:[\s\S]*max: 3/); + }); +}); diff --git a/internal/compare/report_snapshot.go b/internal/compare/report_snapshot.go index 4e07fae..1110996 100644 --- a/internal/compare/report_snapshot.go +++ b/internal/compare/report_snapshot.go @@ -37,5 +37,15 @@ func SnapshotFromReport(r Report, side string) (Snapshot, error) { if len(snap.Tables) == 0 { return snap, fmt.Errorf("the %s side of this comparison has no tables", side) } + // Shapes travel only when relationships were compared: exporting never scans. + for _, d := range r.Relationships { + shape := d.Source + if side == SideTarget { + shape = d.Target + } + if shape != nil { + snap.Relationships = append(snap.Relationships, *shape) + } + } return snap, nil } diff --git a/internal/web/handlers_relationships.go b/internal/web/handlers_relationships.go new file mode 100644 index 0000000..a80862c --- /dev/null +++ b/internal/web/handlers_relationships.go @@ -0,0 +1,190 @@ +package web + +import ( + "context" + "fmt" + "net/http" + + "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" + "github.com/AxeForging/seedstorm/internal/runerr" + "github.com/AxeForging/seedstorm/internal/seeder" +) + +// RelationshipsRequest measures the active connection's relationship shapes. +type RelationshipsRequest struct { + // Counts is exact (aggregate each key) or estimate (planner statistics). + Counts string `json:"counts"` + ScanUnindexed bool `json:"scanUnindexed"` + // ConfirmExact allows exact scans on a production connection. + ConfirmExact bool `json:"confirmExact"` +} + +// handleRelationships: GET returns the shapes measured so far on this +// connection; POST starts a scan job. +func (s *Server) handleRelationships(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + sess, err := s.sessions.fromRequest(r) + if err != nil { + writeError(w, http.StatusUnauthorized, err.Error()) + return + } + _, production := s.productionConnection(sessionTarget(sess)) + writeJSON(w, http.StatusOK, struct { + shapeView + Production bool `json:"production"` + }{sess.shapes.view(), production}) + return + } + startRun(s, w, r, "relationships", s.runRelationships) +} + +// relationshipOptions are the scan options of a request. A production +// connection runs one query at a time, and in estimate mode unless exact +// scans were confirmed. +func (s *Server) relationshipOptions(t connectionTarget, counts string, scanUnindexed, confirmExact bool) (relations.Options, []string, error) { + mode, err := compare.ParseCountMode(counts) + if err != nil { + return relations.Options{}, nil, err + } + opts := relations.Options{Mode: relations.Exact, IncludeUnindexed: scanUnindexed} + if mode == compare.CountEstimate { + opts.Mode = relations.Estimate + } + var notes []string + if c, ok := s.productionConnection(t); ok { + opts.Limits.Concurrency = 1 + if opts.Mode == relations.Exact && !confirmExact { + opts.Mode = relations.Estimate + notes = append(notes, fmt.Sprintf("%s is a production connection: reading estimates only (confirm an exact scan to aggregate every key)", c.Label)) + } + } + return opts, notes, nil +} + +func (s *Server) runRelationships(ctx context.Context, sess *Session, req RelationshipsRequest, jc JobControl) (map[string]any, error) { + log := jobLogger(jc) + opts, notes, err := s.relationshipOptions(sessionTarget(sess), req.Counts, req.ScanUnindexed, req.ConfirmExact) + if err != nil { + return nil, err + } + for _, n := range notes { + log.Warn().Msg(n) + } + jc.Phase("introspect") + ep, _, err := endpointFor(ctx, sess, false) + if err != nil { + return nil, runerr.At(runerr.PhaseIntrospect, "", err) + } + jc.Phase("scan") + gen := sess.shapes.begin(opts.Mode == relations.Estimate) + defer sess.shapes.finish(gen) + log.Info().Str("database", ep.Label).Str("mode", string(opts.Mode)).Bool("scan_unindexed", opts.IncludeUnindexed).Msg("Measuring relationships (read-only)") + opts.OnEdge = func(done, total int, sh relations.Shape) { + sess.shapes.put(gen, total, sh) + logShape(jc, sh) + jc.Progress(done, total, sh.Child+"."+sh.Column) + } + shapes, err := ep.Shapes(ctx, opts) + if err != nil { + return nil, err + } + outcomes := map[db.ReadOutcome]int{} + for _, sh := range shapes { + outcomes[sh.Outcome]++ + } + if ctx.Err() != nil { + return map[string]any{"shapes": shapes, "outcomes": outcomes}, ctx.Err() + } + jc.Phase("done") + log.Info().Int("relationships", len(shapes)).Msg("Relationships measured") + return map[string]any{"shapes": shapes, "outcomes": outcomes}, nil +} + +// logShape writes one relationship's result, warning when it was not measured. +func logShape(jc JobControl, sh relations.Shape) { + log := jobLogger(jc) + name := sh.Child + "." + sh.Column + switch sh.Outcome { + case db.OutcomeOK, relations.OutcomeEstimated: + ev := log.Info() + if sh.Large && sh.Outcome == db.OutcomeOK { + ev = log.Warn().Bool("large", true) + } + ev.Str("relationship", name).Float64("avg", round2(sh.Avg)).Int64("max", sh.Max).Str("outcome", string(sh.Outcome)).Msg("Relationship") + default: + log.Warn().Str("relationship", name).Str("outcome", string(sh.Outcome)).Msg(sh.Detail) + } +} + +func round2(f float64) float64 { return float64(int64(f*100+0.5)) / 100 } + +// CompareRelationshipsRequest compares relationship shapes of two connections +// (or a snapshot with relationships and a connection). +type CompareRelationshipsRequest struct { + CompareRequest + ScanUnindexed bool `json:"scanUnindexed"` + ConfirmExact bool `json:"confirmExact"` +} + +func (s *Server) handleCompareRelationships(w http.ResponseWriter, r *http.Request) { + startRun(s, w, r, "compare relationships", s.runCompareRelationships) +} + +func (s *Server) runCompareRelationships(ctx context.Context, _ *Session, req CompareRelationshipsRequest, jc JobControl) (map[string]any, error) { + log := jobLogger(jc) + if req.SourceSnapshot != nil && len(req.SourceSnapshot.Relationships) == 0 { + return nil, runerr.OnSide(runerr.SideSource, seeder.ErrNoRelationships) + } + jc.Phase("connect") + source, target, err := s.connectBoth(ctx, log, req.Source, req.SourceSnapshot, req.Target) + if err != nil { + return nil, err + } + srcEP, err := snapshotOrSessionEndpoint(ctx, source, req.SourceSnapshot) + if err != nil { + return nil, runerr.OnSide(runerr.SideSource, runerr.At(runerr.PhaseIntrospect, "", err)) + } + tgtEP, _, err := endpointFor(ctx, target, false) + if err != nil { + return nil, runerr.OnSide(runerr.SideTarget, runerr.At(runerr.PhaseIntrospect, "", err)) + } + // The stricter side decides: production anywhere means one query at a + // time, and estimates unless confirmed. + opts, notes, err := s.relationshipOptions(s.refTarget(req.Target), req.Counts, req.ScanUnindexed, req.ConfirmExact) + if err != nil { + return nil, err + } + if source != nil { + srcOpts, srcNotes, _ := s.relationshipOptions(sessionTarget(source), req.Counts, req.ScanUnindexed, req.ConfirmExact) + if srcOpts.Mode == relations.Estimate { + opts.Mode = relations.Estimate + } + if srcOpts.Limits.Concurrency == 1 { + opts.Limits.Concurrency = 1 + } + notes = append(srcNotes, notes...) + } + for _, n := range notes { + log.Warn().Msg(n) + } + jc.Phase("scan") + log.Info().Str("source", srcEP.Label).Str("target", tgtEP.Label).Str("mode", string(opts.Mode)).Msg("Comparing relationships (read-only)") + drift, err := seeder.CompareShapes(ctx, srcEP, tgtEP, opts, func(side string, done, total int, sh relations.Shape) { + if sh.Outcome != db.OutcomeOK && sh.Outcome != relations.OutcomeEstimated { + log.Warn().Str("side", side).Str("relationship", sh.Child+"."+sh.Column).Str("outcome", string(sh.Outcome)).Msg(sh.Detail) + } + jc.Progress(done, total, side+": "+sh.Child+"."+sh.Column) + }) + if err != nil { + return nil, err + } + counts := map[compare.ShapeStatus]int{} + for _, d := range drift { + counts[d.Status]++ + } + jc.Phase("done") + log.Info().Int("same", counts[compare.ShapeSame]).Int("differs", counts[compare.ShapeDiffers]).Int("unknown", counts[compare.ShapeUnknown]).Msg("Relationships compared") + return map[string]any{"relationships": drift}, nil +} diff --git a/internal/web/handlers_snapshots.go b/internal/web/handlers_snapshots.go index 330fd2c..3a4d64e 100644 --- a/internal/web/handlers_snapshots.go +++ b/internal/web/handlers_snapshots.go @@ -34,6 +34,9 @@ func (s *Server) handleSnapshotEncode(w http.ResponseWriter, r *http.Request) { Side string `json:"side"` Snapshot *compare.Snapshot `json:"snapshot"` Format string `json:"format"` + // Relationships keeps the shapes the report or snapshot holds; + // without it the file has counts only. + Relationships bool `json:"relationships"` } if err := decodeLimited(w, r, &req, maxSnapshotBody); err != nil { writeError(w, http.StatusBadRequest, err.Error()) @@ -50,6 +53,9 @@ func (s *Server) handleSnapshotEncode(w http.ResponseWriter, r *http.Request) { return } } + if !req.Relationships { + snap.Relationships = nil + } format := strings.ToLower(strings.TrimSpace(req.Format)) data, err := compare.EncodeSnapshot(snap, format) if err != nil { @@ -57,9 +63,10 @@ func (s *Server) handleSnapshotEncode(w http.ResponseWriter, r *http.Request) { return } writeJSON(w, http.StatusOK, map[string]any{ - "content": string(data), - "filename": snapshotFilename(snap.Label, req.Side, format), - "tables": len(snap.Tables), + "content": string(data), + "filename": snapshotFilename(snap.Label, req.Side, format), + "tables": len(snap.Tables), + "relationships": len(snap.Relationships), }) } diff --git a/internal/web/server.go b/internal/web/server.go index b34ca81..393dbdb 100644 --- a/internal/web/server.go +++ b/internal/web/server.go @@ -137,6 +137,8 @@ func (s *Server) routes() { s.mux.HandleFunc("/api/compare", s.handleCompareRun) s.mux.HandleFunc("/api/mirror", s.handleMirrorRun) s.mux.HandleFunc("/api/snapshot", s.handleSnapshotRun) + s.mux.HandleFunc("/api/relationships", s.handleRelationships) + s.mux.HandleFunc("/api/compare/relationships", s.handleCompareRelationships) s.mux.HandleFunc("/api/snapshots/encode", s.handleSnapshotEncode) s.mux.HandleFunc("/api/snapshots/parse", s.handleSnapshotParse) s.mux.HandleFunc("/api/generators", s.handleGeneratorsJSON) diff --git a/internal/web/session.go b/internal/web/session.go index 16a68db..6409957 100644 --- a/internal/web/session.go +++ b/internal/web/session.go @@ -65,6 +65,9 @@ type Session struct { accessMu sync.Mutex access *accessView + + // Relationship shapes measured on this connection (see shapes.go). + shapes shapeCache } // SessionRegistry holds active sessions keyed by their server-issued ID. @@ -400,4 +403,5 @@ func (s *Session) InvalidateCounts() { s.mu.Lock() defer s.mu.Unlock() s.counts, s.countsAt = nil, time.Time{} + s.shapes.reset() } diff --git a/internal/web/shapes.go b/internal/web/shapes.go new file mode 100644 index 0000000..f5a2176 --- /dev/null +++ b/internal/web/shapes.go @@ -0,0 +1,85 @@ +package web + +import ( + "sort" + "sync" + "time" + + "github.com/AxeForging/seedstorm/internal/relations" +) + +// shapeCache holds a session's relationship shapes as a scan measures them, so +// the workspace can show each one as it finishes and export them later. +type shapeCache struct { + mu sync.Mutex + gen int + byKey map[string]relations.Shape + total int + running bool + takenAt time.Time + estimate bool +} + +// shapeView is what /api/relationships returns. +type shapeView struct { + Shapes []relations.Shape `json:"shapes"` + Total int `json:"total"` + Running bool `json:"running"` + TakenAt string `json:"takenAt,omitempty"` + Estimate bool `json:"estimate,omitempty"` +} + +// begin starts a new scan, dropping earlier shapes; the returned generation +// keeps a superseded scan from writing into the new one. +func (c *shapeCache) begin(estimate bool) int { + c.mu.Lock() + defer c.mu.Unlock() + c.gen++ + c.byKey, c.total, c.running, c.takenAt, c.estimate = map[string]relations.Shape{}, 0, true, time.Now().UTC(), estimate + return c.gen +} + +func (c *shapeCache) put(gen, total int, s relations.Shape) { + c.mu.Lock() + defer c.mu.Unlock() + if gen != c.gen { + return + } + c.total = total + c.byKey[s.Child+"."+s.Column] = s +} + +func (c *shapeCache) finish(gen int) { + c.mu.Lock() + defer c.mu.Unlock() + if gen == c.gen { + c.running = false + } +} + +// reset forgets shapes: a run wrote to the database. +func (c *shapeCache) reset() { + c.mu.Lock() + defer c.mu.Unlock() + c.gen++ + c.byKey, c.total, c.running, c.takenAt = nil, 0, false, time.Time{} +} + +func (c *shapeCache) view() shapeView { + c.mu.Lock() + defer c.mu.Unlock() + v := shapeView{Shapes: make([]relations.Shape, 0, len(c.byKey)), Total: c.total, Running: c.running, Estimate: c.estimate} + for _, s := range c.byKey { + v.Shapes = append(v.Shapes, s) + } + sort.Slice(v.Shapes, func(i, j int) bool { + if v.Shapes[i].Child != v.Shapes[j].Child { + return v.Shapes[i].Child < v.Shapes[j].Child + } + return v.Shapes[i].Column < v.Shapes[j].Column + }) + if !c.takenAt.IsZero() { + v.TakenAt = c.takenAt.Format(time.RFC3339) + } + return v +} diff --git a/internal/web/shapes_test.go b/internal/web/shapes_test.go new file mode 100644 index 0000000..a36e2c8 --- /dev/null +++ b/internal/web/shapes_test.go @@ -0,0 +1,101 @@ +package web + +import ( + "net/http" + "strings" + "testing" + "time" + + "github.com/AxeForging/seedstorm/internal/compare" + "github.com/AxeForging/seedstorm/internal/db" + "github.com/AxeForging/seedstorm/internal/relations" +) + +// A scan fills the cache edge by edge; a newer scan or a write makes an older +// scan's late results disappear instead of mixing into the new ones. +func TestShapeCache_StaleScansAndWritesDoNotLeak(t *testing.T) { + sess := &Session{ID: "s"} + old := sess.shapes.begin(false) + sess.shapes.put(old, 2, relations.Shape{Child: "orders", Column: "user_id", Outcome: db.OutcomeOK}) + if v := sess.shapes.view(); len(v.Shapes) != 1 || !v.Running || v.Total != 2 || v.TakenAt == "" { + t.Fatalf("partial view = %+v", v) + } + + fresh := sess.shapes.begin(true) + sess.shapes.put(old, 2, relations.Shape{Child: "stale", Column: "x_id"}) + sess.shapes.finish(old) + if v := sess.shapes.view(); len(v.Shapes) != 0 || !v.Running || !v.Estimate { + t.Fatalf("a superseded scan leaked into the new one: %+v", v) + } + sess.shapes.put(fresh, 1, relations.Shape{Child: "orders", Column: "user_id", Outcome: relations.OutcomeEstimated}) + sess.shapes.finish(fresh) + if v := sess.shapes.view(); len(v.Shapes) != 1 || v.Running { + t.Fatalf("finished view = %+v", v) + } + + sess.InvalidateCounts() + sess.shapes.put(fresh, 1, relations.Shape{Child: "late", Column: "x_id"}) + if v := sess.shapes.view(); len(v.Shapes) != 0 || v.TakenAt != "" { + t.Fatalf("shapes survived a write: %+v", v) + } +} + +// Production connections read estimates one query at a time unless an exact +// scan is confirmed; other connections get what was asked. +func TestRelationshipOptions_ProductionReadsEstimatesUnlessConfirmed(t *testing.T) { + s, _, prod := productionServer(t) + opts, notes, err := s.relationshipOptions(sessionTarget(prod), "exact", false, false) + if err != nil || opts.Mode != relations.Estimate || opts.Limits.Concurrency != 1 || len(notes) != 1 || !strings.Contains(notes[0], "billing-prod") { + t.Fatalf("production unconfirmed = %+v %v %v", opts, notes, err) + } + opts, notes, _ = s.relationshipOptions(sessionTarget(prod), "exact", true, true) + if opts.Mode != relations.Exact || opts.Limits.Concurrency != 1 || !opts.IncludeUnindexed || len(notes) != 0 { + t.Fatalf("production confirmed = %+v %v", opts, notes) + } + other := sessionTarget(&Session{Info: ConnectionInfo{DBType: "postgres", Host: "localhost", DBName: "scratch"}}) + opts, _, _ = s.relationshipOptions(other, "exact", false, false) + if opts.Mode != relations.Exact || opts.Limits.Concurrency != 0 { + t.Fatalf("ordinary connection = %+v", opts) + } + if _, _, err := s.relationshipOptions(other, "sideways", false, false); err == nil { + t.Fatal("an unknown mode was accepted") + } +} + +// Exported files carry relationships only when asked, from a comparison or a +// snapshot; exporting never scans. +func TestSnapshotsAPI_RelationshipsOnlyWhenAsked(t *testing.T) { + s := profileServer(t) + s.sessions.sessions["sess-snap"] = &Session{ID: "sess-snap", DBType: "pgx"} + shape := func(max int64) *relations.Shape { + return &relations.Shape{Child: "orders", Column: "user_id", Parent: "users", ParentColumn: "id", Avg: 2, Max: max, Outcome: db.OutcomeOK} + } + report := compare.Diff( + compare.Snapshot{Label: "prod", DBType: "pgx", Tables: map[string]compare.TableStat{"orders": {Rows: 10}}}, + compare.Snapshot{Label: "stage", DBType: "pgx", Tables: map[string]compare.TableStat{"orders": {Rows: 1}}}, + ) + report.Relationships = []compare.ShapeDrift{{Child: "orders", Column: "user_id", Status: compare.ShapeDiffers, Source: shape(9), Target: shape(1)}} + + rec, out := snapshotCall(t, s, "/api/snapshots/encode", map[string]any{"report": report, "side": "target", "format": "yaml", "relationships": true}) + if rec.Code != http.StatusOK || out["relationships"].(float64) != 1 || !strings.Contains(out["content"].(string), "max: 1") { + t.Fatalf("target export with relationships = %d %v", rec.Code, out) + } + rec, parsed := snapshotCall(t, s, "/api/snapshots/parse", map[string]any{"data": out["content"]}) + if rec.Code != http.StatusOK { + t.Fatalf("exported shapes do not parse back: %s", rec.Body) + } + if snap := parsed["snapshot"].(map[string]any); len(snap["relationships"].([]any)) != 1 { + t.Fatalf("parsed = %v", parsed) + } + + rec, out = snapshotCall(t, s, "/api/snapshots/encode", map[string]any{"report": report, "side": "source", "format": "yaml"}) + if rec.Code != http.StatusOK || out["relationships"].(float64) != 0 || strings.Contains(out["content"].(string), "relationships") { + t.Fatalf("counts-only export = %d %v", rec.Code, out) + } + + snap := compare.Snapshot{Label: "prod", DBType: "pgx", TakenAt: time.Now().UTC(), Tables: map[string]compare.TableStat{"orders": {Rows: 10}}, Relationships: []relations.Shape{*shape(9)}} + _, out = snapshotCall(t, s, "/api/snapshots/encode", map[string]any{"snapshot": snap, "format": "json"}) + if strings.Contains(out["content"].(string), "relationships") { + t.Fatalf("snapshot export without the flag kept shapes: %v", out["content"]) + } +} diff --git a/internal/web/static/app.js b/internal/web/static/app.js index 6004ce6..ccccf83 100644 --- a/internal/web/static/app.js +++ b/internal/web/static/app.js @@ -1503,6 +1503,7 @@ syncTuningSummary(); document.getElementById("ws-recommend")?.addEventListener("click", openRecommendDialog); document.getElementById("ws-snapshot")?.addEventListener("click", takeSnapshot); + document.getElementById("ws-relationships")?.addEventListener("click", openShapesDialog); fetchConnections().then((conns) => { const active = (conns || []).find((c) => c.active); const link = document.getElementById("ws-calibrate"); @@ -1681,6 +1682,112 @@ refresh(); } + // ── relationship shapes of the active connection ────────────────────── + // Shapes are measured by a job; the server keeps each one as it finishes, so + // the graph polls while a scan runs and draws a badge per foreign key. + ws.shapes = {}; + ws.shapesProduction = false; + let shapesPoll = null; + const shapeMeasured = (sh) => sh.outcome === "ok" || sh.outcome === "estimated" || sh.outcome === "skipped: unindexed"; + function shapeBadge(sh) { + if (!shapeMeasured(sh)) return sh.outcome; + const n = (v) => (v == null || v < 0 ? "?" : String(v)); + return `${sh.outcome === "ok" ? "" : "~"}avg ${Number(sh.avg || 0).toFixed(1)} · max ${n(sh.max)}`; + } + function shapeTitle(sh) { + if (!shapeMeasured(sh)) return `${sh.child}.${sh.column}: ${sh.outcome}${sh.detail ? " — " + sh.detail : ""}`; + const n = (v) => (v == null || v < 0 ? "?" : Number(v).toLocaleString()); + return `${sh.child}.${sh.column} → ${sh.parent}: children per parent min ${n(sh.min)} · avg ${Number(sh.avg || 0).toFixed(2)} · p50 ${n(sh.p50)} · p95 ${n(sh.p95)} · max ${n(sh.max)} · ${Math.round((sh.zeroShare || 0) * 100)}% of parents have none` + + (sh.nullShare ? ` · ${Math.round(sh.nullShare * 100)}% NULL keys` : "") + (sh.outcome === "ok" ? "" : ` · ${sh.detail || "estimated"}`); + } + + async function loadShapes() { + let view; + try { + view = await fetchJSON("/api/relationships", { cache: "no-store" }); + } catch (_) { return; } + ws.shapesProduction = !!view.production; + ws.shapes = Object.fromEntries((view.shapes || []).map((sh) => [sh.child + "." + sh.column, sh])); + applyShapes(); + const status = document.getElementById("ws-shapes-status"); + const shapes = view.shapes || []; + if (status) { + status.hidden = !shapes.length && !view.running; + const unknown = shapes.filter((sh) => !shapeMeasured(sh)).length; + const est = shapes.filter((sh) => sh.outcome !== "ok" && shapeMeasured(sh)).length; + status.dataset.state = view.running ? "loading" : "ready"; + status.textContent = view.running + ? `Measuring relationships… ${shapes.length}${view.total ? "/" + view.total : ""}` + : `${shapes.length} ${shapes.length === 1 ? "relationship" : "relationships"} measured${est ? ` · ${est} estimated (~)` : ""}${unknown ? ` · ${unknown} not measured` : ""}`; + } + if (view.running && !shapesPoll) shapesPoll = setInterval(loadShapes, 1500); + if (!view.running && shapesPoll) { clearInterval(shapesPoll); shapesPoll = null; } + } + + function applyShapes() { + if (!ws.cy) return; + ws.cy.batch(() => { + for (const e of ws.edges) { + const edge = ws.cy.getElementById(e.id); + const sh = ws.shapes[e.target + "." + e.column]; + if (!edge || edge.empty()) continue; + if (sh) { + edge.data("shapeLabel", shapeBadge(sh)); + edge.data("shapeTitle", shapeTitle(sh)); + edge.toggleClass("shape-unknown", !shapeMeasured(sh)); + } else { + edge.removeData("shapeLabel shapeTitle"); + edge.removeClass("shape-unknown"); + } + } + }); + } + + function openShapesDialog() { + const dlg = document.createElement("dialog"); + dlg.className = "prod-dialog shapes-dialog"; + dlg.setAttribute("data-testid", "shapes-dialog"); + const prod = ws.shapesProduction; + dlg.innerHTML = ` +
+

Analyze relationships

+

For every foreign key: how many children each parent has (min, avg, p95, max) and how many parents have none. Read-only, two queries at a time, 60s limit per key; cancel keeps what finished.

+
+ + +
+

Exact aggregates each key in the database. Estimate reads planner statistics: instant, averages only where the database keeps them.

+ + ${prod ? `` : ""} +
+ + +
+
`; + document.body.appendChild(dlg); + dlg.querySelector("[data-start]").addEventListener("click", async () => { + const form = dlg.querySelector("form"); + const body = { + counts: form.querySelector('input[name="counts"]:checked').value, + scanUnindexed: form.querySelector('input[name="unindexed"]').checked, + confirmExact: !!form.querySelector('input[name="confirmExact"]')?.checked, + }; + dlg.close(); + let j; + try { + j = await postRun("/api/relationships", body); + } catch (err) { + appendLog("ERROR: " + (err.message || err)); + return; + } + activateTab("logs"); + streamJob(j.id, j.name, { onEnd: () => loadShapes() }, j.bootId); + setTimeout(loadShapes, 300); + }); + dlg.addEventListener("close", () => dlg.remove()); + dlg.showModal(); + } + // ── counts snapshot of the active connection ────────────────────────── const IMPORTED_COUNTS_KEY = "seedstorm.importedCounts.v1"; async function takeSnapshot() { @@ -1712,6 +1819,7 @@
+
Rendering…
@@ -1726,8 +1834,10 @@ let content = ""; const render = async () => { const format = dlg.querySelector('input[name="format"]:checked').value; + const withShapes = dlg.querySelector('input[name="relationships"]').checked; + const payload = withShapes ? { ...snapshot, relationships: Object.values(ws.shapes) } : snapshot; try { - const out = await fetchJSON("/api/snapshots/encode", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ snapshot, format }) }); + const out = await fetchJSON("/api/snapshots/encode", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ snapshot: payload, format, relationships: withShapes }) }); content = out.content; preview.textContent = content; if (link.dataset.url) URL.revokeObjectURL(link.dataset.url); @@ -1738,7 +1848,7 @@ preview.textContent = "Could not render: " + (err.message || err); } }; - dlg.querySelectorAll('input[name="format"]').forEach((r) => r.addEventListener("change", render)); + dlg.querySelectorAll('input[name="format"], input[name="relationships"]').forEach((r) => r.addEventListener("change", render)); dlg.querySelector("[data-copy]").addEventListener("click", () => copyText(content)); dlg.querySelector("[data-compare]").addEventListener("click", () => { // Hand the snapshot to Compare as an imported source. @@ -1792,6 +1902,7 @@ } else { refreshCounts(false); } + loadShapes(); } catch (err) { setGraphLoading("Graph failed", err.message || String(err), true); } @@ -1921,6 +2032,8 @@ }); ws.cy.on("tap", "node", (ev) => toggleSelect(ev.target.id())); + // A measured relationship's numbers live on its child table's columns. + ws.cy.on("tap", "edge[shapeLabel]", (ev) => showDetail(ev.target.data("target"))); ws.cy.on("cxttap", "node", (ev) => { ev.preventDefault?.(); showDetail(ev.target.id()); @@ -2182,6 +2295,24 @@ "control-point-step-size": 42, }, }, + { + selector: "edge[shapeLabel]", + style: { + "label": "data(shapeLabel)", + "font-size": 9, + "color": "#c9d4cc", + "text-background-color": "#11161a", + "text-background-opacity": 0.85, + "text-background-padding": "2px", + "text-background-shape": "roundrectangle", + "text-rotation": "autorotate", + "min-zoomed-font-size": 7, + }, + }, + { + selector: "edge.shape-unknown", + style: { "color": "#d8b56f" }, + }, { selector: "edge[?nullable]", style: { "line-style": "dashed", "line-color": "#4a5169", "target-arrow-color": "#4a5169" }, @@ -3043,6 +3174,8 @@ const flags = []; if (c.pk || c.PK) flags.push('PK'); if (c.fk || c.FK) flags.push(`FK -> ${escapeHTML(c.fk || c.FK)}`); + const sh = ws.shapes[tableName + "." + col]; + if (sh) flags.push(`${escapeHTML(shapeBadge(sh))}`); if (c.nullable || c.Nullable) flags.push('nullable'); return `${escapeHTML(col)} ${flags.join(" ")}${escapeHTML(c.type || c.Type || "")}`; }).join(""); @@ -3396,6 +3529,7 @@ const out = document.getElementById("job-result"); if (out) renderJobResult(out, job.result || {}, job.name || ws.mode); refreshCounts(); + loadShapes(); } // refreshCounts fills node counts without blocking the graph: the page is diff --git a/internal/web/static/compare.css b/internal/web/static/compare.css index 9387fc4..00c5321 100644 --- a/internal/web/static/compare.css +++ b/internal/web/static/compare.css @@ -86,6 +86,39 @@ .cmp-meter { display: block; height: 4px; border-radius: 2px; background: rgba(216,181,111,0.18); overflow: hidden; } .cmp-meter i { display: block; height: 100%; background: var(--tgt); transition: width 500ms var(--ease-out); } +/* ── relationship shapes ── */ +.cmp-shapes { + border: 1px solid var(--line); + border-radius: var(--radius-lg); + background: var(--panel); + overflow: hidden; +} +.cmp-shapes-head { + display: flex; flex-wrap: wrap; align-items: center; justify-content: space-between; gap: 10px; + padding: 10px 12px; +} +.cmp-shapes-head h2 { margin: 0; font-size: 15px; } +.cmp-shapes-head p { margin: 2px 0 0; max-width: 72ch; } +.cmp-shapes-actions { display: flex; flex-wrap: wrap; align-items: center; gap: 10px 14px; } +.cmp-shapes-body { border-top: 1px solid var(--line); max-height: 52vh; overflow: auto; } +.cmp-shapes-summary { padding: 8px 12px; color: var(--muted); font-size: 12px; border-bottom: 1px solid var(--line); } +.cmp-shape-row { + display: grid; + grid-template-columns: minmax(160px, 1.2fr) minmax(0, 1fr) minmax(0, 1fr) 96px; + align-items: center; gap: 10px; + padding: 7px 12px; + border-bottom: 1px solid rgba(231,236,232,0.05); + font-size: 12px; +} +.cmp-shape-row.head { color: var(--muted); font-family: var(--mono); font-size: 11px; text-transform: uppercase; letter-spacing: 0.06em; background: rgba(11,15,18,0.35); } +.cmp-shape-row code { overflow: hidden; text-overflow: ellipsis; white-space: nowrap; } +.cmp-shape-src { color: var(--src); text-align: right; font-variant-numeric: tabular-nums; } +.cmp-shape-tgt { color: var(--tgt); font-variant-numeric: tabular-nums; } +.cmp-shape-status { font-size: 11px; color: var(--muted); text-align: right; } +.cmp-shape-status.differs { color: var(--warn); } +.cmp-shape-row .cell-unknown { color: var(--muted); } +.cmp-export-shapes { align-self: end; } + /* ── table gauges ── */ .cmp-grid { display: grid; grid-template-columns: minmax(0, 1fr) 340px; gap: 14px; align-items: start; } .cmp-table-panel { @@ -285,6 +318,12 @@ .cmp-empty { flex-direction: column; align-items: flex-start; } .cmp-search, .cmp-search input { width: 100%; } .cmp-gauge-head { display: none; } + .cmp-shape-row.head { display: none; } + .cmp-shape-row { grid-template-columns: minmax(0, 1fr) auto; grid-template-areas: "name status" "src src" "tgt tgt"; row-gap: 3px; } + .cmp-shape-row > code { grid-area: name; } + .cmp-shape-src { grid-area: src; text-align: left; } + .cmp-shape-tgt { grid-area: tgt; } + .cmp-shape-status { grid-area: status; } .cmp-gauge { grid-template-columns: 24px minmax(0, 1fr) auto; grid-template-areas: "check name delta" "check src src" "check tgt tgt"; diff --git a/internal/web/static/compare.js b/internal/web/static/compare.js index a397387..ad5c708 100644 --- a/internal/web/static/compare.js +++ b/internal/web/static/compare.js @@ -290,6 +290,7 @@ } app.querySelectorAll("[data-count]").forEach((b) => { b.textContent = counts[b.dataset.count]; }); renderGauges(); + renderShapes(); $("cmp-writes").innerHTML = `Writes only to ${ui().escapeHTML(targetLabel())}. The source is read.`; $("cmp-plan").disabled = false; } @@ -410,7 +411,7 @@ const job = await runJob("/api/mirror", mirrorRequest(true)); if (job.status !== "done") throw new Error(job.error || "plan " + job.status); state.plan = job.result; - state.report = job.result.report; // counts are fresh from the plan run + state.report = { ...job.result.report, relationships: state.report?.relationships }; // counts are fresh from the plan run state.pairKey = pairKey(); saveReport(); render(); @@ -619,6 +620,74 @@ compare(); } + // ── relationship shapes ───────────────────────────────────────────── + // A side's cell: avg / p95 / max children per parent and parents without any. + function shapeCell(sh) { + if (!sh) return `—`; + const measured = sh.outcome === "ok" || sh.outcome === "estimated" || sh.outcome === "skipped: unindexed"; + if (!measured) return `${ui().escapeHTML(sh.outcome)}`; + const est = sh.outcome !== "ok" ? "~" : ""; + const n = (v) => (v == null || v < 0 ? "?" : Number(v).toLocaleString()); + const title = `${sh.parents >= 0 ? n(sh.parents) + " parents · " : ""}${n(sh.children)} children · min ${n(sh.min)} · p50 ${n(sh.p50)}${sh.nullShare ? ` · ${Math.round(sh.nullShare * 100)}% NULL keys` : ""}${est ? " · estimated" : ""}`; + return `${est}${Number(sh.avg || 0).toFixed(2)} / ${n(sh.p95)} / ${n(sh.max)} · ${Math.round((sh.zeroShare || 0) * 100)}%`; + } + + const shapeStatusText = { same: "same", differs: "differs", source_only: "source only", target_only: "target only", unknown: "unknown" }; + + function renderShapes() { + const body = $("cmp-shapes-body"); + const drift = state.report?.relationships; + const exportBox = $("cmp-export-relationships"); + exportBox.disabled = !drift?.length; + if (exportBox.disabled) exportBox.checked = false; + $("cmp-export-relationships-note").textContent = drift?.length ? `(${drift.length})` : "(compare them first)"; + if (!drift) { body.hidden = true; body.innerHTML = ""; return; } + body.hidden = false; + if (!drift.length) { + body.innerHTML = `

Neither side has foreign keys to compare.

`; + return; + } + const counts = {}; + for (const d of drift) counts[d.status] = (counts[d.status] || 0) + 1; + const rows = drift.map((d) => { + const parent = (d.source || d.target)?.parent; + return `
+ ${ui().escapeHTML(d.child + "." + d.column)}${parent ? ` → ${ui().escapeHTML(parent)}` : ""} + ${shapeCell(d.source)} + ${shapeCell(d.target)} + ${shapeStatusText[d.status] || d.status} +
`; + }).join(""); + body.innerHTML = `

${drift.length} ${drift.length === 1 ? "relationship" : "relationships"} · ${counts.same || 0} same · ${counts.differs || 0} differ · ${(counts.source_only || 0) + (counts.target_only || 0)} on one side · ${counts.unknown || 0} unknown · children per parent: avg / p95 / max · parents without children · ~ = estimated

+ ${rows}`; + } + + async function compareShapes() { + if (state.busy || !state.report) return; + setBusy(true); + const button = $("cmp-shapes-run"); + button.disabled = true; + button.textContent = "Measuring…"; + const counts = app.querySelector('input[name="counts"]:checked').value; + try { + const job = await runJob("/api/compare/relationships", { + source: refOf($("cmp-source")), sourceSnapshot: sourceSnapshot(), target: refOf($("cmp-target")), counts, + scanUnindexed: $("cmp-shapes-unindexed").checked, confirmExact: $("cmp-shapes-exact-prod").checked, + }); + if (job.status !== "done") throw new Error(job.error || "relationships " + job.status); + state.report.relationships = job.result.relationships || []; + saveReport(); + renderShapes(); + } catch (err) { + showOutcome("error", "Relationships could not be compared", err.message); + $("cmp-logs").open = true; + } finally { + button.disabled = false; + button.textContent = "Compare relationships"; + setBusy(false); + } + } + // ── counts export ─────────────────────────────────────────────────── let exportSeq = 0; async function renderExport() { @@ -629,7 +698,7 @@ $("cmp-export-status").textContent = "Rendering…"; let out; try { - const res = await fetch("/api/snapshots/encode", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ report: state.report, side, format }) }); + const res = await fetch("/api/snapshots/encode", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ report: state.report, side, format, relationships: $("cmp-export-relationships").checked }) }); out = await res.json(); if (!res.ok) throw new Error(out.error || res.statusText); } catch (err) { @@ -641,7 +710,7 @@ } if (seq !== exportSeq) return; $("cmp-export-preview").textContent = out.content; - $("cmp-export-status").textContent = `${out.tables} tables · ${out.filename}`; + $("cmp-export-status").textContent = `${out.tables} tables${out.relationships ? ` · ${out.relationships} relationships` : ""} · ${out.filename}`; const link = $("cmp-export-download"); const type = format === "json" ? "application/json" : "text/yaml"; if (link.dataset.url) URL.revokeObjectURL(link.dataset.url); @@ -684,9 +753,13 @@ // A compare or mirror started before leaving the page keeps running on the // server: reattach to it and show its outcome when it ends. function resumeRunningJob() { - ui().resumeRun(["compare", "mirror"], { + ui().resumeRun(["compare", "mirror", "compare relationships"], { onEnd: (job) => { - if (job.name === "compare" && job.status === "done" && job.result?.report) { + if (job.name === "compare relationships" && job.status === "done" && state.report) { + state.report.relationships = job.result?.relationships || []; + saveReport(); + renderShapes(); + } else if (job.name === "compare" && job.status === "done" && job.result?.report) { state.report = job.result.report; state.pairKey = pairKey(); saveReport(); @@ -763,7 +836,8 @@ readImportFile(ev.dataTransfer?.files?.[0]); }); $("cmp-export").addEventListener("click", openExport); - app.querySelectorAll('input[name="export-side"], input[name="export-format"]').forEach((r) => r.addEventListener("change", renderExport)); + app.querySelectorAll('input[name="export-side"], input[name="export-format"], #cmp-export-relationships').forEach((r) => r.addEventListener("change", renderExport)); + $("cmp-shapes-run").addEventListener("click", compareShapes); $("cmp-export-copy").addEventListener("click", async () => { await ui().copyText($("cmp-export-preview").textContent); $("cmp-export-copy").textContent = "Copied"; diff --git a/internal/web/static/style.css b/internal/web/static/style.css index b90dc62..cbee123 100644 --- a/internal/web/static/style.css +++ b/internal/web/static/style.css @@ -1848,6 +1848,8 @@ textarea { font-family: var(--mono); font-size: 12px; } .prod-dialog form { display: grid; gap: 12px; padding: 18px; } .prod-dialog h2 { margin: 0; font-size: 1.05rem; } .prod-dialog input { width: 100%; box-sizing: border-box; } +.prod-dialog input[type="checkbox"], .prod-dialog input[type="radio"] { width: auto; } +.badge.shape { background: rgba(125, 180, 150, 0.14); color: #b9d8c4; } .prod-dialog footer { display: flex; justify-content: flex-end; gap: 8px; flex-wrap: wrap; } .pill.param { background: rgba(216,181,111,0.12); border-color: rgba(216,181,111,0.32); color: var(--accent-2); } @@ -2424,5 +2426,6 @@ textarea { font-family: var(--mono); font-size: 12px; } .ws-recommend-btn { justify-self: start; } .ws-counts-actions { display: flex; flex-wrap: wrap; gap: 6px; min-width: 0; } +.ws-counts-actions .btn-ghost { min-height: 30px; padding: 4px 9px; font-size: 12px; } .snapshot-dialog { max-width: min(640px, calc(100vw - 32px)); } .snapshot-preview { max-height: 40vh; overflow: auto; } diff --git a/internal/web/templates/compare.html.tmpl b/internal/web/templates/compare.html.tmpl index ef0c6e0..a7671a6 100644 --- a/internal/web/templates/compare.html.tmpl +++ b/internal/web/templates/compare.html.tmpl @@ -146,6 +146,21 @@
+
+
+
+

Relationships

+

Children per parent for every foreign key, on both sides. Read-only; keys without an index are estimated unless you include them. Counts mode applies: estimate reads statistics only.

+
+
+ + + +
+
+ +
+
Job log idle
@@ -218,6 +233,7 @@
+