diff --git a/db/db.go b/db/db.go index c53aa364a..c46510496 100644 --- a/db/db.go +++ b/db/db.go @@ -1,17 +1,24 @@ package db import ( + "cmp" "context" "database/sql" "embed" "errors" "fmt" + "os" + "path/filepath" + "strconv" + "strings" "sync" "time" + _ "github.com/jackc/pgx/v5/stdlib" "github.com/mattn/go-sqlite3" "github.com/navidrome/navidrome/conf" _ "github.com/navidrome/navidrome/db/migrations" + "github.com/navidrome/navidrome/db/pglite" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/utils/hasher" "github.com/navidrome/navidrome/utils/natural" @@ -38,6 +45,9 @@ const migrationsFolder = "migrations" // (tests/benchmarks) and rebuilt, but the driver is process-global and registers only once. var registerDriverOnce sync.Once +// embedded is the in-process PGlite instance when DbPath uses the pglite:// scheme (spike). +var embedded *pglite.PGlite + func Db() *sql.DB { return singleton.GetInstance(func() *sql.DB { registerDriverOnce.Do(func() { @@ -51,6 +61,41 @@ func Db() *sql.DB { }) }) Path = conf.Server.DbPath + if dataDir, ok := strings.CutPrefix(Path, "pglite://"); ok { + Dialect = "postgres" + Driver = "pgx" + tarball := cmp.Or(os.Getenv("ND_PGLITE_TARBALL"), "tmp/pglite-wasi-O2-fix.tar.gz") + logFile, _ := os.Create(filepath.Join(dataDir, "..", "pglite.log")) + socketDir := filepath.Dir(dataDir) + pg, err := pglite.New(context.Background(), pglite.Config{ + DataDir: dataDir, Tarball: tarball, Stderr: logFile, SocketDir: socketDir, + }) + if err != nil { + log.Fatal("Error starting embedded PGlite", err) + } + embedded = pg + log.Debug("Opening DataBase", "dbPath", Path, "driver", Driver, "dsn", pg.DSN()) + abs, _ := filepath.Abs(socketDir) + log.Info("PGlite spike: connect with psql", "cmd", "PGSSLMODE=disable psql -h "+abs+" -U postgres postgres") + db, err := sql.Open(Driver, pg.DSN()) + if err != nil { + log.Fatal("Error opening database", err) + } + db.SetMaxOpenConns(1) + return db + } + if isPostgres(Path) { + Dialect = "postgres" + Driver = "pgx" + log.Debug("Opening DataBase", "dbPath", Path, "driver", Driver) + db, err := sql.Open(Driver, Path) + if err != nil { + log.Fatal("Error opening database", err) + } + // One PGlite backend means one PG session; the bridge serializes clients onto it. + db.SetMaxOpenConns(4) + return db + } if Path == ":memory:" { Path = "file::memory:?cache=shared&_foreign_keys=on" conf.Server.DbPath = Path @@ -59,7 +104,11 @@ func Db() *sql.DB { } log.Debug("Opening DataBase", "dbPath", Path, "driver", Driver) db, err := sql.Open(Driver, Path) - db.SetMaxOpenConns(conf.MaxOpenConns()) + maxConns := conf.MaxOpenConns() + if v, err := strconv.Atoi(os.Getenv("ND_TEST_MAXCONNS")); err == nil && v > 0 { + maxConns = v // spike only: simulate the PGlite bridge's serialized session + } + db.SetMaxOpenConns(maxConns) if err != nil { log.Fatal("Error opening database", err) } @@ -76,11 +125,55 @@ func Close(ctx context.Context) { if err != nil { log.Error(ctx, "Error closing Database", err) } + if embedded != nil { + _ = embedded.Close() + } +} + +// applySpikeSchema loads a crudely translated SQLite schema, one statement at a time, tolerating failures. +func applySpikeSchema(ctx context.Context, db *sql.DB, path string) { + data, err := os.ReadFile(path) + if err != nil { + log.Fatal(ctx, "PGlite spike: cannot read schema", err) + } + var ok, failed int + for _, stmt := range strings.Split(string(data), ";\n") { + if strings.TrimSpace(stmt) == "" { + continue + } + if _, err := db.ExecContext(ctx, stmt); err != nil { + failed++ + log.Warn(ctx, "PGlite spike: schema statement failed", "stmt", strings.SplitN(strings.TrimSpace(stmt), "\n", 2)[0], err) + continue + } + ok++ + } + log.Warn(ctx, "PGlite spike: schema applied", "ok", ok, "failed", failed) +} + +// isPostgres reports whether the DbPath is a Postgres connection URL (PGlite spike). +func isPostgres(path string) bool { + return strings.HasPrefix(path, "postgres://") || strings.HasPrefix(path, "postgresql://") } func Init(ctx context.Context) func() { db := Db() + if Dialect == "postgres" { + var version string + if err := db.QueryRowContext(ctx, "SELECT version()").Scan(&version); err != nil { + log.Fatal(ctx, "PGlite spike: cannot reach database", err) + } + log.Warn(ctx, "PGlite spike: connected; migrations are SQLite-only and were SKIPPED", "version", version) + if _, err := db.ExecContext(ctx, "CREATE TABLE IF NOT EXISTS property (id text PRIMARY KEY, value text)"); err != nil { + log.Fatal(ctx, "PGlite spike: cannot create property table", err) + } + if schema := os.Getenv("ND_PGLITE_SCHEMA"); schema != "" { + applySpikeSchema(ctx, db, schema) + } + return func() { Close(ctx) } + } + // Disable foreign_keys to allow re-creating tables in migrations _, err := db.ExecContext(ctx, "PRAGMA foreign_keys=off") defer func() { diff --git a/db/pglite/README.md b/db/pglite/README.md new file mode 100644 index 000000000..de69715fe --- /dev/null +++ b/db/pglite/README.md @@ -0,0 +1,52 @@ +# db/pglite — spike: PostgreSQL embedded in the Navidrome binary + +Runs the [PGlite](https://pglite.dev) WASI build of PostgreSQL 17.5 inside the process with +[wazero](https://wazero.io) (pure Go, no CGO), and exposes it to `database/sql` through the real PostgreSQL wire +protocol. Enabled with `DbPath = "pglite://"`. + +**This is exploratory. It is not a supported way to run Navidrome.** The 131 SQLite migrations are not ported, so the +schema has to be supplied by hand (`ND_PGLITE_SCHEMA`), and several queries Navidrome generates are SQLite-only. + +## How it fits together + + Navidrome → database/sql → pgx → unix socket → bridge → wasm memory → PostgreSQL (wasm) + +The bridge speaks no SQL. It copies wire-protocol bytes, and only inspects message framing, the handshake, and the +ReadyForQuery status byte. `PGlite.OpenDB()` swaps the unix socket for an in-process pipe; both perform the same, so +the socket is the default because external tools can use it. + +## Getting the wasm binary + +Not in the repository: it is a 6.7 MB archive containing a 17 MB wasm module. Build it, then point the tests at it +with `ND_PGLITE_TARBALL`, or drop it at `tmp/pglite-wasi-O2-fix.tar.gz` (the default the tests look for). + +The published WASI builds are all compiled `-O0`; rebuilding with `-O2` is worth 2 to 3x on every query and cuts +startup from 16 s to 1.8 s. The build recipe and the patches, including one C fix of our own, are in the spike notes +under `tmp/pglite-build/` on the `pglite-spike` branch. + +## Connecting with psql + + PGPASSWORD=x PGSSLMODE=disable psql -h /absolute/path/to/DataFolder -U postgres postgres + +The exact command is logged at startup. A password is required even though PGlite skips authentication. Plain SQL +works; `\dt` and `\d` do not, because the build's session schema is `pg_catalog` and its system views are missing. + +## Known limits + +- **One session.** PGlite is a single PostgreSQL backend. The bridge accepts several connections and serializes them, + holding the session across a whole transaction and a whole handshake. Four connections perform exactly like one. +- **One process.** Set `DevExternalScanner = false`. Navidrome otherwise scans in a child process, which would open a + second PGlite on the same data directory and corrupt it. +- **Session state is shared.** `SET`, temp tables and the like leak between clients. +- **Simple protocol only.** The WASI build mishandles extended-protocol portals, so the DSN pins + `default_query_exec_mode=simple_protocol`. +- **No `COPY FROM STDIN`.** The tick loop cannot suspend to wait for the data rows. + +## Environment switches + +| variable | effect | +|---|---| +| `ND_PGLITE_TARBALL` | path to the wasm archive (tests) | +| `ND_PGLITE_SCHEMA` | apply a schema file after connecting, tolerating failures | +| `ND_PGLITE_TRACE` | log wire traffic and per-tick timings | +| `ND_PGLITE_NOCMA` | force the file transport instead of shared memory | diff --git a/db/pglite/concurrency_test.go b/db/pglite/concurrency_test.go new file mode 100644 index 000000000..3a1eae237 --- /dev/null +++ b/db/pglite/concurrency_test.go @@ -0,0 +1,194 @@ +package pglite_test + +import ( + "context" + "database/sql" + "os" + "sync" + "time" + + _ "github.com/jackc/pgx/v5/stdlib" + "github.com/navidrome/navidrome/db/pglite" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("PGlite concurrent connections", func() { + var pg *pglite.PGlite + var db *sql.DB + + BeforeEach(func() { + tarball := os.Getenv("ND_PGLITE_TARBALL") + if tarball == "" { + tarball = "tmp/pglite-wasi-O2-fix.tar.gz" + } + if _, err := os.Stat(tarball); err != nil { + Skip("pglite tarball not found: " + tarball) + } + var err error + pg, err = pglite.New(context.Background(), pglite.Config{ + DataDir: GinkgoT().TempDir(), Tarball: tarball, Stderr: GinkgoWriter, + }) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(pg.Close) + + db, err = sql.Open("pgx", pg.DSN()) + Expect(err).ToNot(HaveOccurred()) + db.SetMaxOpenConns(4) + DeferCleanup(db.Close) + + _, err = db.Exec("CREATE TABLE t (id int PRIMARY KEY, name text)") + Expect(err).ToNot(HaveOccurred()) + for i := range 200 { + _, err = db.Exec("INSERT INTO t VALUES ($1, $2)", i, "row") + Expect(err).ToNot(HaveOccurred()) + } + }) + + It("serves 4 concurrent connections with correct results", func() { + var wg sync.WaitGroup + errs := make(chan error, 4) + for w := range 4 { + wg.Add(1) + go func() { + defer wg.Done() + for i := range 25 { + id := (w*25 + i) % 200 + var name string + if err := db.QueryRow("SELECT name FROM t WHERE id = $1", id).Scan(&name); err != nil { + errs <- err + return + } + if name != "row" { + errs <- sql.ErrNoRows + return + } + } + }() + } + wg.Wait() + close(errs) + for err := range errs { + Expect(err).ToNot(HaveOccurred()) + } + GinkgoWriter.Println("bridge connections opened:", pg.Connections()) + Expect(pg.Connections()).To(BeNumerically(">", 1)) + }) + + It("does not let a second connection enter an open transaction", func() { + tx, err := db.Begin() + Expect(err).ToNot(HaveOccurred()) + _, err = tx.Exec("INSERT INTO t VALUES (1000, 'in-tx')") + Expect(err).ToNot(HaveOccurred()) + + // A second connection must neither see the uncommitted row nor join the transaction. + done := make(chan int, 1) + go func() { + var n int + if err := db.QueryRow("SELECT count(*) FROM t WHERE id = 1000").Scan(&n); err == nil { + done <- n + } + }() + select { + case n := <-done: + Fail("second connection ran inside the open transaction, count=" + string(rune('0'+n))) + case <-time.After(500 * time.Millisecond): + // blocked, as required + } + + Expect(tx.Commit()).To(Succeed()) + Eventually(done, 5*time.Second).Should(Receive(Equal(1))) + }) + + It("shares session state between connections (single backend)", func() { + _, err := db.Exec("SET application_name = 'alpha'") + Expect(err).ToNot(HaveOccurred()) + var seen string + Expect(db.QueryRow("SELECT current_setting('application_name')").Scan(&seen)).To(Succeed()) + GinkgoWriter.Println("application_name seen by another connection:", seen) + Expect(seen).To(Equal("alpha"), "documents the limitation: all clients share one PG session") + }) + + It("makes progress with parallel transactional workers (the scanner's pattern)", func() { + done := make(chan error, 3) + for w := range 3 { + go func() { + for i := range 20 { + tx, err := db.Begin() + if err != nil { + done <- err + return + } + if _, err := tx.Exec("INSERT INTO t VALUES ($1, $2)", 10000+w*100+i, "worker"); err != nil { + _ = tx.Rollback() + done <- err + return + } + if err := tx.Commit(); err != nil { + done <- err + return + } + } + done <- nil + }() + } + for range 3 { + Eventually(done, 60*time.Second).Should(Receive(BeNil())) + } + var n int + Expect(db.QueryRow("SELECT count(*) FROM t WHERE name = 'worker'").Scan(&n)).To(Succeed()) + Expect(n).To(Equal(60)) + }) + + It("serves queries over an in-process pipe, with no unix socket", func() { + piped, err := pg.OpenDB() + Expect(err).ToNot(HaveOccurred()) + defer piped.Close() + piped.SetMaxOpenConns(2) + + var n int + Expect(piped.QueryRow("SELECT count(*) FROM t").Scan(&n)).To(Succeed()) + Expect(n).To(Equal(200)) + + tx, err := piped.Begin() + Expect(err).ToNot(HaveOccurred()) + _, err = tx.Exec("INSERT INTO t VALUES ($1, $2)", 5000, "piped") + Expect(err).ToNot(HaveOccurred()) + Expect(tx.Commit()).To(Succeed()) + + var name string + Expect(piped.QueryRow("SELECT name FROM t WHERE id = $1", 5000).Scan(&name)).To(Succeed()) + Expect(name).To(Equal("piped")) + + // An error must not kill the piped connection either. + _, err = piped.Exec("SELECT * FROM nope") + Expect(err).To(MatchError(ContainSubstring("does not exist"))) + Expect(piped.QueryRow("SELECT count(*) FROM t").Scan(&n)).To(Succeed()) + Expect(n).To(Equal(201)) + }) + + It("measures throughput at 1 vs 4 connections", func() { + bench := func(conns int) time.Duration { + d, err := sql.Open("pgx", pg.DSN()) + Expect(err).ToNot(HaveOccurred()) + defer d.Close() + d.SetMaxOpenConns(conns) + start := time.Now() + var wg sync.WaitGroup + for range conns { + wg.Add(1) + go func() { + defer wg.Done() + for range 100 / conns { + var n int + _ = d.QueryRow("SELECT count(*) FROM t WHERE id < 150").Scan(&n) + } + }() + } + wg.Wait() + return time.Since(start) + } + one, four := bench(1), bench(4) + GinkgoWriter.Printf("100 queries: 1 conn %s, 4 conns %s\n", one.Round(time.Millisecond), four.Round(time.Millisecond)) + }) +}) diff --git a/db/pglite/pglite.go b/db/pglite/pglite.go new file mode 100644 index 000000000..e77f97c34 --- /dev/null +++ b/db/pglite/pglite.go @@ -0,0 +1,801 @@ +// Package pglite runs the PGlite WASI build of PostgreSQL inside the process using wazero. +// Spike: ported from github.com/elliots/go-pglite (wasmtime) to wazero. +package pglite + +import ( + "bytes" + "context" + "crypto/rand" + "database/sql" + "errors" + "fmt" + "io" + "net" + "os" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/stdlib" + "github.com/tetratelabs/wazero" + "github.com/tetratelabs/wazero/api" + "github.com/tetratelabs/wazero/experimental" + "github.com/tetratelabs/wazero/sys" +) + +type Config struct { + DataDir string + Tarball string // path to pglite-wasi-17.tar.gz + Database string + User string + Stderr io.Writer + // ServerSettings are appended to postgresql.conf after initdb (e.g. fsync=off). + ServerSettings map[string]string + // ExtraArgs are appended to the postgres argv (e.g. "-c", "fsync=off"). + ExtraArgs []string + // WasmOverride, if set, is used instead of the pglite.wasi extracted from the tarball. + WasmOverride string + // SocketDir is where the unix socket is created. Defaults to a temp dir removed on Close. + SocketDir string +} + +type PGlite struct { + wasmMu sync.Mutex + cfg Config + ctx context.Context + cancel context.CancelFunc + + runtime wazero.Runtime + mod api.Module + stdin *switchableStdin + trace bool + v2 bool + + connections atomic.Int64 + sessionMu sync.Mutex // one backend means one session: serialize exchanges, and whole transactions + tFiles, tWasm time.Duration // trace-mode timing accumulators + ticks int + paramStatus []byte // ParameterStatus messages from the first handshake, replayed on later ones + + dataDir string + ioBase string + tempSocketDir bool + socketDir string + socketPath string + listener net.Listener + wg sync.WaitGroup + + fnInteractiveOne api.Function + fnInteractiveWrite api.Function + fnInteractiveRead api.Function + fnGetChannel api.Function + + // Shared-memory ("CMA") transport: wire bytes go straight into wasm memory instead of + // through the .in/.out files. Falls back to files when the module does not offer it. + cmaOK bool + cmaAddr uint32 + pendingWireLen uint32 + fnUseWire api.Function + fnClearError api.Function +} + +func New(ctx context.Context, cfg Config) (*PGlite, error) { + if cfg.Database == "" { + cfg.Database = "postgres" + } + if cfg.User == "" { + cfg.User = "postgres" + } + if cfg.Stderr == nil { + cfg.Stderr = io.Discard + } + ctx, cancel := context.WithCancel(ctx) + pg := &PGlite{cfg: cfg, ctx: ctx, cancel: cancel, dataDir: cfg.DataDir, trace: os.Getenv("ND_PGLITE_TRACE") != ""} + + if err := os.MkdirAll(cfg.DataDir, 0o755); err != nil { + cancel() + return nil, err + } + if WASIBinary == nil { + b, err := os.ReadFile(cfg.Tarball) + if err != nil { + cancel() + return nil, fmt.Errorf("reading tarball: %w", err) + } + WASIBinary = b + } + wasmBinary, err := setupEnvironment(pg.dataDir) + if err != nil { + cancel() + return nil, fmt.Errorf("setup environment: %w", err) + } + if cfg.WasmOverride != "" { + if wasmBinary, err = os.ReadFile(cfg.WasmOverride); err != nil { + cancel() + return nil, fmt.Errorf("reading wasm override: %w", err) + } + } + + // v2 layout (pglite4j build): share files embedded via wasi-vfs, pre-initialized + // cluster shipped as pgdata/, initdb+backend already run by wizer. + pg.v2 = fileExists(filepath.Join(pg.dataDir, "pgdata", "PG_VERSION")) + pgdataDir := filepath.Join(pg.dataDir, "pglite", "base") + if pg.v2 { + pgdataDir = filepath.Join(pg.dataDir, "pgdata") + } + devDir := filepath.Join(pg.dataDir, "dev") + if err := os.MkdirAll(pgdataDir, 0o755); err != nil { + cancel() + return nil, err + } + + start := time.Now() + pg.stdin = &switchableStdin{} + tracker := newFDTracker(pg.dataDir, pgdataDir, devDir) + _ = tracker + ctx = experimental.WithFunctionListenerFactory(ctx, tracker) + pg.runtime = wazero.NewRuntime(ctx) + if err := instantiateWASI(ctx, pg.runtime, tracker, pg.stdin); err != nil { + pg.Close() + return nil, fmt.Errorf("instantiating WASI: %w", err) + } + compiled, err := pg.runtime.CompileModule(ctx, wasmBinary) + if err != nil { + pg.Close() + return nil, fmt.Errorf("compiling pglite.wasi: %w", err) + } + compileTime := time.Since(start) + + modCfg := wazero.NewModuleConfig(). + WithArgs(append([]string{"/tmp/pglite/bin/postgres", "--single"}, append(cfg.ExtraArgs, cfg.Database)...)...). + WithEnv("ENVIRONMENT", "wasm32_wasi_preview1"). + WithEnv("PREFIX", "/tmp/pglite"). + WithEnv("PGDATA", pg.guestPGData()). + WithEnv("PGSYSCONFDIR", "/tmp/pglite"). + WithEnv("PGUSER", cfg.User). + WithEnv("PGDATABASE", cfg.Database). + WithEnv("MODE", "REACT"). + WithEnv("REPL", "N"). + WithEnv("TZ", "UTC"). + WithEnv("PGTZ", "UTC"). + WithEnv("PATH", "/tmp/pglite/bin"). + WithFSConfig(wazero.NewFSConfig(). + WithDirMount(pg.dataDir, "/tmp"). + WithDirMount(pgdataDir, pg.guestPGData()). // preopen order must match the wizer snapshot + WithDirMount(devDir, "/dev")). + WithStdin(pg.stdin). + WithStdout(io.Discard). + WithStderr(cfg.Stderr). + WithSysWalltime(). + WithSysNanotime(). + WithSysNanosleep(). + WithRandSource(rand.Reader). + WithStartFunctions() // call _start ourselves so an exit does not abort instantiation + + pg.mod, err = pg.runtime.InstantiateModule(ctx, compiled, modCfg) + if err != nil { + pg.Close() + return nil, fmt.Errorf("instantiating pglite.wasi: %w", err) + } + + startFns := []string{"_start", "pgl_initdb", "pgl_backend"} + if pg.v2 { + startFns = nil // wizer already ran them at build time + } + for _, name := range startFns { + fn := pg.mod.ExportedFunction(name) + if fn == nil { + continue + } + if _, err := fn.Call(ctx); err != nil { + var exitErr *sys.ExitError + if errors.As(err, &exitErr) && exitErr.ExitCode() == 0 { + continue + } + pg.Close() + return nil, fmt.Errorf("%s: %w", name, err) + } + if name == "pgl_initdb" && len(cfg.ServerSettings) > 0 { + if err := appendSettings(filepath.Join(pgdataDir, "postgresql.conf"), cfg.ServerSettings); err != nil { + pg.Close() + return nil, err + } + } + } + if pg.v2 { + if fn := pg.mod.ExportedFunction("interactive_write"); fn != nil { + if _, err := fn.Call(ctx, api.EncodeI32(0)); err != nil { + pg.Close() + return nil, fmt.Errorf("interactive_write(0): %w", err) + } + } + } + fmt.Fprintf(cfg.Stderr, "# pglite: wazero compile=%s init=%s v2=%v\n", compileTime, time.Since(start), pg.v2) + + pg.fnInteractiveOne = pg.mod.ExportedFunction("interactive_one") + pg.fnInteractiveWrite = pg.mod.ExportedFunction("interactive_write") + pg.fnInteractiveRead = pg.mod.ExportedFunction("interactive_read") + pg.fnGetChannel = pg.mod.ExportedFunction("get_channel") + pg.fnUseWire = pg.mod.ExportedFunction("use_wire") + pg.fnClearError = pg.mod.ExportedFunction("clear_error") + pg.probeCMA(ctx) + fmt.Fprintf(cfg.Stderr, "# pglite: transport=%s\n", map[bool]string{true: "shared-memory", false: "files"}[pg.cmaOK]) + if pg.fnInteractiveOne == nil { + pg.Close() + return nil, errors.New("module missing 'interactive_one' export") + } + + if err := pg.startBridge(); err != nil { + pg.Close() + return nil, fmt.Errorf("starting socket bridge: %w", err) + } + return pg, nil +} + +func (pg *PGlite) guestPGData() string { + if pg.v2 { + return "/pgdata" + } + return "/tmp/pglite/base" +} + +func fileExists(path string) bool { + _, err := os.Stat(path) + return err == nil +} + +func appendSettings(confPath string, settings map[string]string) error { + f, err := os.OpenFile(confPath, os.O_APPEND|os.O_WRONLY, 0o644) + if err != nil { + return fmt.Errorf("opening postgresql.conf: %w", err) + } + defer f.Close() + for k, v := range settings { + if _, err := fmt.Fprintf(f, "\n%s = %s\n", k, v); err != nil { + return err + } + } + return nil +} + +// DSN returns a pgx connection string for the embedded instance. Simple protocol only: +// the WASI build does not cope with extended-protocol portals. +func (pg *PGlite) DSN() string { + return fmt.Sprintf("host=%s port=5432 dbname=%s user=%s sslmode=disable default_query_exec_mode=simple_protocol", + pg.socketDir, pg.cfg.Database, pg.cfg.User) +} + +// OpenDB returns a pool that reaches the backend over an in-process pipe, skipping the unix +// socket. The socket stays up for external tools. +func (pg *PGlite) OpenDB() (*sql.DB, error) { + cfg, err := pgx.ParseConfig(pg.DSN()) + if err != nil { + return nil, err + } + cfg.DialFunc = func(context.Context, string, string) (net.Conn, error) { + client, server := net.Pipe() + pg.wg.Add(1) + go func() { + defer pg.wg.Done() + pg.handleConn(server, pg.ioBase) + }() + return client, nil + } + return stdlib.OpenDB(*cfg), nil +} + +func (pg *PGlite) Close() error { + pg.cancel() + if pg.listener != nil { + _ = pg.listener.Close() + } + pg.wg.Wait() + if pg.socketDir != "" { + if pg.tempSocketDir { + _ = os.RemoveAll(pg.socketDir) + } else { + _ = os.Remove(pg.socketPath) + } + } + if pg.stdin != nil { + _ = pg.stdin.Close() + } + if pg.runtime != nil { + return pg.runtime.Close(context.Background()) + } + return nil +} + +func (pg *PGlite) startBridge() error { + sockDir := pg.cfg.SocketDir + if sockDir == "" { + var err error + if sockDir, err = os.MkdirTemp("", "pglite-sock-*"); err != nil { + return err + } + pg.tempSocketDir = true + } else { + if err := os.MkdirAll(sockDir, 0o700); err != nil { + return err + } + // pgx only treats a host as a unix socket when the path is absolute. + abs, err := filepath.Abs(sockDir) + if err != nil { + return err + } + sockDir = abs + _ = os.Remove(filepath.Join(sockDir, ".s.PGSQL.5432")) // a stale socket from a crash + } + pg.socketDir = sockDir + pg.socketPath = filepath.Join(sockDir, ".s.PGSQL.5432") + ln, err := net.Listen("unix", pg.socketPath) + if err != nil { + return err + } + pg.listener = ln + + ioBase := filepath.Join(pg.dataDir, "pglite", "base", ".s.PGSQL.5432") + if pg.v2 { + ioBase = filepath.Join(pg.dataDir, "pgdata", ".s.PGSQL.5432") + } + pg.ioBase = ioBase + for _, suffix := range []string{".in", ".out", ".lock.in", ".lock.out"} { + _ = os.Remove(ioBase + suffix) + } + + pg.wg.Add(1) + go func() { + defer pg.wg.Done() + for { + conn, err := ln.Accept() + if err != nil { + return + } + pg.wg.Add(1) + go func() { + defer pg.wg.Done() + pg.handleConn(conn, ioBase) + }() + } + }() + return nil +} + +// Connections reports how many client connections the bridge has accepted. +func (pg *PGlite) Connections() int64 { return pg.connections.Load() } + +func (pg *PGlite) handleConn(conn net.Conn, ioBase string) { + pg.connections.Add(1) + inTx, holding := false, false + defer func() { + if holding { + if inTx { + pg.rollback(ioBase) + } + pg.sessionMu.Unlock() + } + }() + handshakeDone := false + startupDone := false + var pending []byte + if pg.trace { + fmt.Fprintln(pg.cfg.Stderr, "# bridge: client connected") + defer fmt.Fprintln(pg.cfg.Stderr, "# bridge: client disconnected") + } + defer conn.Close() + outFile := ioBase + ".out" + buf := make([]byte, 65536) + + for { + select { + case <-pg.ctx.Done(): + return + default: + } + + _ = conn.SetReadDeadline(time.Now().Add(16 * time.Millisecond)) + n, readErr := conn.Read(buf) + if n > 0 { + pending = append(pending, buf[:n]...) + n = completeMessages(pending, &startupDone) + } + if n > 0 { + packet := pending[:n:n] + pending = append([]byte(nil), pending[n:]...) + if pg.trace { + fmt.Fprintf(pg.cfg.Stderr, "# bridge C> %s (%d bytes)\n", wireTags(packet), n) + } + // Take the session before touching .in: it is shared by every client. + if !holding { + pg.sessionMu.Lock() + holding = true + } + pg.wasmMu.Lock() + t0 := time.Now() + replies, trapErr := pg.forwardWire(packet, outFile) + if pg.trace { + fmt.Fprintf(pg.cfg.Stderr, "# bridge timing: total=%s files=%s wasm=%s ticks=%d\n", + time.Since(t0).Round(time.Microsecond), pg.tFiles.Round(time.Microsecond), pg.tWasm.Round(time.Microsecond), pg.ticks) + pg.tFiles, pg.tWasm, pg.ticks = 0, 0, 0 + } + pg.wasmMu.Unlock() + if !handshakeDone { + replies, handshakeDone = pg.fixHandshake(replies) + } else { + replies = ensureReadyForQuery(replies, trapErr) + } + if trapErr != nil && pg.trace { + fmt.Fprintf(pg.cfg.Stderr, "# bridge: trap recovered: %v\n", strings.SplitN(trapErr.Error(), "\n", 2)[0]) + } + // Keep the session across a multi-step handshake and across a transaction: both are + // stateful in the one backend, so another client must not interleave into them. + status := lastReadyStatus(replies) + inTx = status == 'T' || status == 'E' + if handshakeDone && !inTx { + pg.sessionMu.Unlock() + holding = false + } + // A PG ERROR surfaces as a trap here; the reply is already complete, so keep the client. + if !pg.sendReplies(conn, replies) { + return + } + } + if readErr != nil { + var netErr net.Error + if errors.As(readErr, &netErr) && netErr.Timeout() { + continue + } + return + } + } +} + +func (pg *PGlite) forwardWire(packet []byte, outFile string) ([][]byte, error) { + const maxTicks = 256 + ctx := pg.ctx + if len(packet) > 0 { + if err := pg.send(packet, strings.TrimSuffix(outFile, ".out")); err != nil { + return nil, err + } + } + if pg.fnUseWire != nil { + _, _ = pg.fnUseWire.Call(ctx, api.EncodeI32(1)) + } + var replies [][]byte + for range maxTicks { + producedBefore := pg.collectReply(outFile, &replies) + t0 := time.Now() + _, err := pg.fnInteractiveOne.Call(ctx) + pg.tWasm += time.Since(t0) + pg.ticks++ + if pg.trace { + fmt.Fprintf(pg.cfg.Stderr, "# bridge tick %d: %s\n", pg.ticks, time.Since(t0).Round(time.Microsecond)) + } + if err != nil && pg.v2 { + // v2 build: pgl_on_error sets a flag and traps; clear_error does the full cleanup. + pg.collectReply(outFile, &replies) + if fn := pg.mod.ExportedFunction("pgl_check_error"); fn != nil { + res, cerr := fn.Call(ctx) + if pg.trace { + fmt.Fprintf(pg.cfg.Stderr, "# bridge v2 trap: pgl_check_error=%v err=%v\n", res, cerr) + } + if cerr == nil && len(res) > 0 && api.DecodeI32(res[0]) != 0 { + if pg.fnClearError != nil { + _, _ = pg.fnClearError.Call(ctx) + } + if pg.fnInteractiveWrite != nil { + _, _ = pg.fnInteractiveWrite.Call(ctx, api.EncodeI32(-1)) + } + _, _ = pg.fnInteractiveOne.Call(ctx) + pg.collectReply(outFile, &replies) + } + } + return replies, err + } + if err != nil { + pg.collectReply(outFile, &replies) + // Traps bypass PG_CATCH; flag shmem_exit_inprogress so clear_error drops active portals. + const shmemExitAddr = 4895117 + pg.mod.Memory().WriteByte(shmemExitAddr, 1) + if pg.fnClearError != nil { + _, _ = pg.fnClearError.Call(ctx) + } + pg.mod.Memory().WriteByte(shmemExitAddr, 0) + _ = os.Remove(strings.TrimSuffix(outFile, ".out") + ".in") + if pg.fnInteractiveWrite != nil { + _, _ = pg.fnInteractiveWrite.Call(ctx, api.EncodeI32(-1)) + } + if pg.fnUseWire != nil { + _, _ = pg.fnUseWire.Call(ctx, api.EncodeI32(1)) + } + _, _ = pg.fnInteractiveOne.Call(ctx) + pg.collectReply(outFile, &replies) + return replies, err + } + producedAfter := pg.collectReply(outFile, &replies) + if !producedBefore && !producedAfter { + break + } + if endsWithReadyForQuery(replies) { + break // a complete response; skip the empty probe tick + } + } + return replies, nil +} + +func (pg *PGlite) collectReply(outFile string, replies *[][]byte) bool { + t0 := time.Now() + defer func() { pg.tFiles += time.Since(t0) }() + if pg.cmaOK { + // A negative channel means the C side put this reply in a file after all. + if res, err := pg.fnGetChannel.Call(pg.ctx); err == nil && len(res) > 0 && api.DecodeI32(res[0]) >= 0 { + return pg.collectFromMemory(replies) + } + } + data, err := os.ReadFile(outFile) + if err != nil || len(data) == 0 { + return false + } + _ = os.Remove(outFile) + *replies = append(*replies, data) + return true +} + +// collectFromMemory reads one reply out of the shared wire buffer. Must be called with wasmMu held. +func (pg *PGlite) collectFromMemory(replies *[][]byte) bool { + res, err := pg.fnInteractiveRead.Call(pg.ctx) + if err != nil || len(res) == 0 { + return false + } + n := api.DecodeI32(res[0]) + if n <= 0 { + return false + } + data, ok := pg.mod.Memory().Read(pg.cmaAddr+pg.pendingWireLen+1, uint32(n)) + if !ok { + return false + } + *replies = append(*replies, bytes.Clone(data)) + _, _ = pg.fnInteractiveWrite.Call(pg.ctx, api.EncodeI32(0)) + pg.pendingWireLen = 0 + return true +} + +func (pg *PGlite) sendReplies(conn net.Conn, replies [][]byte) bool { + for _, data := range replies { + if len(data) == 0 { + continue + } + if pg.trace { + fmt.Fprintf(pg.cfg.Stderr, "# bridge S> %s (%d bytes)\n", wireTags(data), len(data)) + if data[0] == 'E' || data[0] == 'N' { + fmt.Fprintf(pg.cfg.Stderr, "# bridge S> raw %q\n", data) + } + } + _ = conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + if _, err := conn.Write(data); err != nil { + return false + } + } + return true +} + +// completeMessages returns how many leading bytes of data form whole client messages. The first +// message of a connection (startup/cancel) is untagged; every later one is tag + int32 length. +func completeMessages(data []byte, startupDone *bool) int { + n := 0 + for { + rest := data[n:] + if !*startupDone { + if len(rest) < 4 { + return n + } + size := int(rest[0])<<24 | int(rest[1])<<16 | int(rest[2])<<8 | int(rest[3]) + if size < 4 || size > len(rest) { + return n + } + n += size + *startupDone = true + continue + } + if len(rest) < 5 { + return n + } + size := int(rest[1])<<24 | int(rest[2])<<16 | int(rest[3])<<8 | int(rest[4]) + if size < 4 || size+1 > len(rest) { + return n + } + n += size + 1 + } +} + +func endsWithReadyForQuery(replies [][]byte) bool { + if len(replies) == 0 { + return false + } + last := replies[len(replies)-1] + return len(last) >= 6 && last[len(last)-6] == 'Z' && last[len(last)-5] == 0 && last[len(last)-2] == 5 +} + +// lastReadyStatus returns the status byte of the final ReadyForQuery in replies, or 0 if there is none. +func lastReadyStatus(replies [][]byte) byte { + data := bytes.Join(replies, nil) + var status byte + for rest := data; len(rest) >= 5; { + n := int(rest[1])<<24 | int(rest[2])<<16 | int(rest[3])<<8 | int(rest[4]) + if n < 4 || n+1 > len(rest) { + break + } + if rest[0] == 'Z' && n == 5 { + status = rest[5] + } + rest = rest[n+1:] + } + return status +} + +// probeCMA asks the module for its shared-memory wire buffer. A negative channel means the +// module only speaks the file transport. +func (pg *PGlite) probeCMA(ctx context.Context) { + if os.Getenv("ND_PGLITE_NOCMA") != "" { + return + } + addrFn := pg.mod.ExportedFunction("get_buffer_addr") + if pg.fnInteractiveRead == nil || pg.fnInteractiveWrite == nil || pg.fnGetChannel == nil || addrFn == nil { + return + } + if _, err := pg.fnInteractiveWrite.Call(ctx, api.EncodeI32(0)); err != nil { + return + } + res, err := pg.fnGetChannel.Call(ctx) + if err != nil || len(res) == 0 || api.DecodeI32(res[0]) < 0 { + return + } + channel := api.DecodeI32(res[0]) + addr, err := addrFn.Call(ctx, api.EncodeI32(channel)) + if err != nil || len(addr) == 0 || api.DecodeI32(addr[0]) <= 0 { + return + } + pg.cmaAddr = uint32(api.DecodeI32(addr[0])) + pg.cmaOK = true +} + +// send hands one client packet to the backend. Must be called with wasmMu held. +func (pg *PGlite) send(packet []byte, ioBase string) error { + if pg.cmaOK { + if pg.fnUseWire != nil { + if _, err := pg.fnUseWire.Call(pg.ctx, api.EncodeI32(1)); err != nil { + return err + } + } + if !pg.mod.Memory().Write(pg.cmaAddr, packet) { + return fmt.Errorf("wire buffer too small for %d bytes", len(packet)) + } + if _, err := pg.fnInteractiveWrite.Call(pg.ctx, api.EncodeI32(int32(len(packet)))); err != nil { + return err + } + pg.pendingWireLen = uint32(len(packet)) + return nil + } + if err := os.WriteFile(ioBase+".lock.in", packet, 0o644); err != nil { + return err + } + return os.Rename(ioBase+".lock.in", ioBase+".in") +} + +// rollback aborts a transaction left open by a client that disconnected mid-transaction, so the +// next client does not inherit it. Must be called with sessionMu held. +func (pg *PGlite) rollback(ioBase string) { + sql := "ROLLBACK\x00" + msg := append([]byte{'Q'}, byte((len(sql)+4)>>24), byte((len(sql)+4)>>16), byte((len(sql)+4)>>8), byte(len(sql)+4)) + msg = append(msg, sql...) + pg.wasmMu.Lock() + _, _ = pg.forwardWire(msg, ioBase+".out") + pg.wasmMu.Unlock() +} + +// fixHandshake caches the ParameterStatus ('S') messages of the first handshake and injects them +// into later ones: PGlite only sends them once per process, but every new pgx connection needs them. +func (pg *PGlite) fixHandshake(replies [][]byte) ([][]byte, bool) { + data := bytes.Join(replies, nil) + var params []byte + ready := false + for rest := data; len(rest) >= 5; { + n := int(rest[1])<<24 | int(rest[2])<<16 | int(rest[3])<<8 | int(rest[4]) + if n < 4 || n+1 > len(rest) { + break + } + switch rest[0] { + case 'S': + params = append(params, rest[:n+1]...) + case 'Z': + ready = true + } + rest = rest[n+1:] + } + if !ready { + return replies, false + } + if len(params) > 0 { + pg.paramStatus = params + return replies, true + } + if pg.paramStatus == nil { + return replies, true + } + // No 'S' messages: insert the cached ones before the first 'K' or 'Z'. + for i := 0; i+5 <= len(data); { + n := int(data[i+1])<<24 | int(data[i+2])<<16 | int(data[i+3])<<8 | int(data[i+4]) + if data[i] == 'K' || data[i] == 'Z' { + patched := append(append(append([]byte{}, data[:i]...), pg.paramStatus...), data[i:]...) + return [][]byte{patched}, true + } + i += n + 1 + } + return replies, true +} + +// ensureReadyForQuery appends a ReadyForQuery when PGlite ends an error reply without one, +// which would otherwise leave the client waiting forever. +func ensureReadyForQuery(replies [][]byte, trapErr error) [][]byte { + data := bytes.Join(replies, nil) + if len(data) == 0 && trapErr == nil { + return replies + } + sawError, sawReady := false, false + for rest := data; len(rest) >= 5; { + n := int(rest[1])<<24 | int(rest[2])<<16 | int(rest[3])<<8 | int(rest[4]) + if n < 4 || n+1 > len(rest) { + break + } + switch rest[0] { + case 'E': + sawError = true + case 'Z': + sawReady = true + } + rest = rest[n+1:] + } + if trapErr != nil && !sawError && !sawReady { + replies = append(replies, errorResponse("XX000", "pglite trap: "+strings.SplitN(trapErr.Error(), "\n", 2)[0])) + sawError = true + } + if sawError && !sawReady { + return append(replies, []byte{'Z', 0, 0, 0, 5, 'I'}) + } + return replies +} + +func errorResponse(code, msg string) []byte { + body := "SERROR\x00VERROR\x00C" + code + "\x00M" + msg + "\x00\x00" + n := len(body) + 4 + return append([]byte{'E', byte(n >> 24), byte(n >> 16), byte(n >> 8), byte(n)}, body...) +} + +// wireTags summarizes wire-protocol messages as their type tags, for ND_PGLITE_TRACE. +func wireTags(data []byte) string { + var tags []string + for len(data) >= 5 { + tag := data[0] + if tag == 0 && len(data) >= 8 && data[4] == 0 && data[5] == 3 { // untagged startup message + tags = append(tags, "Startup") + n := int(data[0])<<24 | int(data[1])<<16 | int(data[2])<<8 | int(data[3]) + if n <= 0 || n > len(data) { + break + } + data = data[n:] + continue + } + n := int(data[1])<<24 | int(data[2])<<16 | int(data[3])<<8 | int(data[4]) + tags = append(tags, string(tag)) + if n < 4 || n+1 > len(data) { + tags = append(tags, "...") + break + } + data = data[n+1:] + } + return strings.Join(tags, " ") +} diff --git a/db/pglite/pglite_suite_test.go b/db/pglite/pglite_suite_test.go new file mode 100644 index 000000000..9f00c6a1a --- /dev/null +++ b/db/pglite/pglite_suite_test.go @@ -0,0 +1,17 @@ +package pglite_test + +import ( + "testing" + + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestPGlite(t *testing.T) { + tests.Init(t, false) + log.SetLevel(log.LevelFatal) + RegisterFailHandler(Fail) + RunSpecs(t, "PGlite Suite") +} diff --git a/db/pglite/pglite_test.go b/db/pglite/pglite_test.go new file mode 100644 index 000000000..d7511660e --- /dev/null +++ b/db/pglite/pglite_test.go @@ -0,0 +1,105 @@ +package pglite_test + +import ( + "context" + "database/sql" + "os" + "time" + + _ "github.com/jackc/pgx/v5/stdlib" + "github.com/navidrome/navidrome/db/pglite" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "github.com/tetratelabs/wazero/experimental" + "github.com/tetratelabs/wazero/experimental/logging" +) + +var _ = Describe("PGlite under wazero", func() { + var pg *pglite.PGlite + var db *sql.DB + + BeforeEach(func() { + tarball := os.Getenv("ND_PGLITE_TARBALL") + if tarball == "" { + tarball = "tmp/pglite-wasi-O2-fix.tar.gz" + } + if _, err := os.Stat(tarball); err != nil { + Skip("pglite tarball not found: " + tarball) + } + ctx := context.Background() + if os.Getenv("ND_PGLITE_TRACE") != "" { + ctx = experimental.WithFunctionListenerFactory(ctx, + logging.NewHostLoggingListenerFactory(os.Stderr, logging.LogScopeFilesystem)) + } + var err error + pg, err = pglite.New(ctx, pglite.Config{ + DataDir: GinkgoT().TempDir(), + Tarball: tarball, + Stderr: GinkgoWriter, + }) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(pg.Close) + + db, err = sql.Open("pgx", pg.DSN()) + Expect(err).ToNot(HaveOccurred()) + db.SetMaxOpenConns(1) + DeferCleanup(db.Close) + }) + + It("answers SELECT version()", func() { + var version string + Expect(db.QueryRow("SELECT version()").Scan(&version)).To(Succeed()) + GinkgoWriter.Println("version:", version) + Expect(version).To(ContainSubstring("PostgreSQL 17")) + }) + + It("keeps the connection after a SQL error and survives a reconnect", func() { + _, err := db.Exec("SELECT * FROM nope WHERE id = $1", 1) + Expect(err).To(MatchError(ContainSubstring("does not exist"))) + var one int + Expect(db.QueryRow("SELECT $1::int", 1).Scan(&one)).To(Succeed()) + GinkgoWriter.Println("connections after error:", pg.Connections()) + Expect(pg.Connections()).To(BeEquivalentTo(1)) + + // The session must still be usable for DDL, DML and transactions after the error. + _, err = db.Exec("CREATE TABLE t (id int PRIMARY KEY, name text)") + Expect(err).ToNot(HaveOccurred()) + tx, err := db.Begin() + Expect(err).ToNot(HaveOccurred()) + _, err = tx.Exec("INSERT INTO t VALUES ($1, $2)", 1, "one") + Expect(err).ToNot(HaveOccurred()) + Expect(tx.Commit()).To(Succeed()) + _, err = db.Exec("INSERT INTO t VALUES ($1, $2)", 1, "dup") + Expect(err).To(MatchError(ContainSubstring("duplicate key"))) + var name string + Expect(db.QueryRow("SELECT name FROM t WHERE id = $1", 1).Scan(&name)).To(Succeed()) + Expect(name).To(Equal("one")) + Expect(pg.Connections()).To(BeEquivalentTo(1)) + + // Force a new connection and make sure parameterized queries still work on it. + Expect(db.Close()).To(Succeed()) + db, err = sql.Open("pgx", pg.DSN()) + Expect(err).ToNot(HaveOccurred()) + db.SetMaxOpenConns(1) + Expect(db.QueryRow("SELECT $1::int", 2).Scan(&one)).To(Succeed()) + Expect(one).To(Equal(2)) + Expect(pg.Connections()).To(BeEquivalentTo(2)) + }) + + It("creates a table and reads rows back", func() { + _, err := db.Exec("CREATE TABLE property (id text PRIMARY KEY, value text)") + Expect(err).ToNot(HaveOccurred()) + _, err = db.Exec("INSERT INTO property (id, value) VALUES ($1, $2)", "JWTSecret", "abc") + Expect(err).ToNot(HaveOccurred()) + + start := time.Now() + var value string + Expect(db.QueryRow("SELECT value FROM property WHERE id = $1", "JWTSecret").Scan(&value)).To(Succeed()) + GinkgoWriter.Println("select took", time.Since(start)) + Expect(value).To(Equal("abc")) + + var count int + Expect(db.QueryRow("SELECT count(*) FROM property").Scan(&count)).To(Succeed()) + Expect(count).To(Equal(1)) + }) +}) diff --git a/db/pglite/proto_test.go b/db/pglite/proto_test.go new file mode 100644 index 000000000..2636c748d --- /dev/null +++ b/db/pglite/proto_test.go @@ -0,0 +1,76 @@ +package pglite_test + +import ( + "context" + "net" + "os" + "path/filepath" + + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/pgx/v5/pgproto3" + "github.com/navidrome/navidrome/db/pglite" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("PGlite wire protocol", func() { + var pg *pglite.PGlite + + BeforeEach(func() { + tarball := os.Getenv("ND_PGLITE_TARBALL") + if tarball == "" { + tarball = "tmp/pglite-wasi-O2-fix.tar.gz" + } + if _, err := os.Stat(tarball); err != nil { + Skip("pglite tarball not found: " + tarball) + } + var err error + pg, err = pglite.New(context.Background(), pglite.Config{DataDir: GinkgoT().TempDir(), Tarball: tarball, Stderr: GinkgoWriter}) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(pg.Close) + }) + + It("pgconn survives a SQL error", func() { + ctx := context.Background() + conn, err := pgconn.Connect(ctx, pg.DSN()) + Expect(err).ToNot(HaveOccurred()) + results, err := conn.Exec(ctx, "SELECT * FROM nope").ReadAll() + GinkgoWriter.Printf("results=%d err=%v closed=%v txStatus=%c\n", len(results), err, conn.IsClosed(), conn.TxStatus()) + Expect(err).To(MatchError(ContainSubstring("does not exist"))) + Expect(conn.IsClosed()).To(BeFalse()) + }) + + It("pgproto3 decodes every message of the error reply", func() { + cfg, err := pgconn.ParseConfig(pg.DSN()) + Expect(err).ToNot(HaveOccurred()) + sock, err := net.Dial("unix", filepath.Join(cfg.Host, ".s.PGSQL.5432")) + Expect(err).ToNot(HaveOccurred()) + defer sock.Close() + fe := pgproto3.NewFrontend(sock, sock) + fe.Send(&pgproto3.StartupMessage{ProtocolVersion: pgproto3.ProtocolVersionNumber, Parameters: map[string]string{"user": "postgres", "database": "postgres"}}) + Expect(fe.Flush()).To(Succeed()) + readUntilReady := func() { + for { + msg, err := fe.Receive() + Expect(err).ToNot(HaveOccurred()) + GinkgoWriter.Printf(" <- %T\n", msg) + switch m := msg.(type) { + case *pgproto3.AuthenticationMD5Password: + fe.Send(&pgproto3.PasswordMessage{Password: "md5" + "x"}) + Expect(fe.Flush()).To(Succeed()) + case *pgproto3.ErrorResponse: + GinkgoWriter.Printf(" error: %s\n", m.Message) + case *pgproto3.ReadyForQuery: + return + } + } + } + readUntilReady() + fe.Send(&pgproto3.Query{String: "SELECT * FROM nope"}) + Expect(fe.Flush()).To(Succeed()) + readUntilReady() + fe.Send(&pgproto3.Query{String: "SELECT 1"}) + Expect(fe.Flush()).To(Succeed()) + readUntilReady() + }) +}) diff --git a/db/pglite/setup.go b/db/pglite/setup.go new file mode 100644 index 000000000..df995cb82 --- /dev/null +++ b/db/pglite/setup.go @@ -0,0 +1,118 @@ +package pglite + +import ( + "archive/tar" + "bytes" + "compress/gzip" + "crypto/rand" + "fmt" + "io" + "os" + "path/filepath" + "strings" +) + +// WASIBinary holds the pglite-wasi tar.gz contents, read from Config.Tarball on first use. +var WASIBinary []byte + +// setupEnvironment extracts the archive on first run and returns the pglite.wasi bytes. +func setupEnvironment(dataDir string) (wasmBinary []byte, err error) { + pgBaseDir := filepath.Join(dataDir, "pglite") + + versionFile := filepath.Join(pgBaseDir, "base", "PG_VERSION") + wasmFile := filepath.Join(pgBaseDir, "bin", "pglite.wasi") + if _, err := os.Stat(versionFile); err != nil && !fileExists(wasmFile) { + if WASIBinary == nil { + return nil, fmt.Errorf("pglite WASI binary not available (WASIBinary is nil)") + } + if err := extractTarGz(dataDir, WASIBinary); err != nil { + return nil, fmt.Errorf("extracting pglite-wasi.tar.gz: %w", err) + } + } + + devDir := filepath.Join(dataDir, "dev") + if err := os.MkdirAll(devDir, 0o755); err != nil { + return nil, fmt.Errorf("creating dev dir: %w", err) + } + urandomPath := filepath.Join(devDir, "urandom") + if _, err := os.Stat(urandomPath); err != nil { + randomBytes := make([]byte, 256) + if _, err := rand.Read(randomBytes); err != nil { + return nil, fmt.Errorf("generating random bytes: %w", err) + } + if err := os.WriteFile(urandomPath, randomBytes, 0o644); err != nil { + return nil, fmt.Errorf("writing urandom: %w", err) + } + } + + wasmPath := filepath.Join(pgBaseDir, "bin", "pglite.wasi") + wasmBinary, err = os.ReadFile(wasmPath) + if err != nil { + return nil, fmt.Errorf("reading pglite.wasi: %w", err) + } + + return wasmBinary, nil +} + +func extractTarGz(destDir string, data []byte) error { + gzReader, err := gzip.NewReader(bytes.NewReader(data)) + if err != nil { + return fmt.Errorf("opening gzip: %w", err) + } + defer gzReader.Close() + + tarReader := tar.NewReader(gzReader) + + for { + header, err := tarReader.Next() + if err == io.EOF { + break + } + if err != nil { + return fmt.Errorf("reading tar entry: %w", err) + } + + // destDir is mounted as /tmp in the guest, so drop the archive's leading "tmp/". + name := header.Name + name = strings.TrimPrefix(name, "tmp/") + if name == "" { + continue + } + + target := filepath.Join(destDir, name) + + if !strings.HasPrefix(filepath.Clean(target), filepath.Clean(destDir)) { + return fmt.Errorf("tar entry %q escapes destination", header.Name) + } + + switch header.Typeflag { + case tar.TypeDir: + if err := os.MkdirAll(target, os.FileMode(header.Mode)); err != nil { + return err + } + case tar.TypeReg: + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return err + } + f, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, os.FileMode(header.Mode)) + if err != nil { + return err + } + if _, err := io.Copy(f, tarReader); err != nil { + f.Close() + return err + } + f.Close() + case tar.TypeSymlink: + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return err + } + os.Remove(target) + if err := os.Symlink(header.Linkname, target); err != nil { + return err + } + } + } + + return nil +} diff --git a/db/pglite/stdin.go b/db/pglite/stdin.go new file mode 100644 index 000000000..01493a79f --- /dev/null +++ b/db/pglite/stdin.go @@ -0,0 +1,125 @@ +package pglite + +import ( + "context" + "io" + "os" + "path/filepath" + "sync" + + "github.com/tetratelabs/wazero" + "github.com/tetratelabs/wazero/api" + "github.com/tetratelabs/wazero/experimental" + "github.com/tetratelabs/wazero/imports/wasi_snapshot_preview1" +) + +const errnoNotsup = 58 + +// switchableStdin is the guest's stdin. PGlite runs initdb by dup2()-ing a script file onto +// fd 0, which wazero refuses for preopens, so we emulate it by swapping the reader instead. +type switchableStdin struct { + mu sync.Mutex + r io.ReadCloser +} + +func (s *switchableStdin) Read(p []byte) (int, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.r == nil { + return 0, io.EOF + } + return s.r.Read(p) +} + +func (s *switchableStdin) set(path string) error { + f, err := os.Open(path) + if err != nil { + return err + } + s.mu.Lock() + defer s.mu.Unlock() + if s.r != nil { + _ = s.r.Close() + } + s.r = f + return nil +} + +func (s *switchableStdin) Close() error { + s.mu.Lock() + defer s.mu.Unlock() + if s.r != nil { + err := s.r.Close() + s.r = nil + return err + } + return nil +} + +// fdTracker listens to path_open so fd_renumber can resolve a guest fd back to a host path. +type fdTracker struct { + mounts map[int32]string // preopen fd -> host dir + opened map[int32]string // guest fd -> host path + pending string + outPtr uint32 + pathBad bool +} + +func newFDTracker(mounts ...string) *fdTracker { + t := &fdTracker{mounts: map[int32]string{}, opened: map[int32]string{}} + for i, m := range mounts { + t.mounts[int32(3+i)] = m + } + return t +} + +func (t *fdTracker) NewFunctionListener(def api.FunctionDefinition) experimental.FunctionListener { + if def.ModuleName() == wasi_snapshot_preview1.ModuleName && def.Name() == "path_open" { + return t + } + return nil +} + +func (t *fdTracker) Before(_ context.Context, mod api.Module, _ api.FunctionDefinition, params []uint64, _ experimental.StackIterator) { + dir, ok := t.mounts[int32(params[0])] + path, okMem := mod.Memory().Read(uint32(params[2]), uint32(params[3])) + t.pathBad = !ok || !okMem + if !t.pathBad { + t.pending = filepath.Join(dir, string(path)) + } + t.outPtr = uint32(params[8]) +} + +func (t *fdTracker) After(_ context.Context, mod api.Module, _ api.FunctionDefinition, results []uint64) { + if t.pathBad || results[0] != 0 { + return + } + if fd, ok := mod.Memory().ReadUint32Le(t.outPtr); ok { + t.opened[int32(fd)] = t.pending + } +} + +func (t *fdTracker) Abort(context.Context, api.Module, api.FunctionDefinition, error) {} + +// instantiateWASI installs wasi_snapshot_preview1 with fd_renumber(fd, 0) emulated via stdin. +func instantiateWASI(ctx context.Context, r wazero.Runtime, tracker *fdTracker, stdin *switchableStdin) error { + b := r.NewHostModuleBuilder(wasi_snapshot_preview1.ModuleName) + wasi_snapshot_preview1.NewFunctionExporter().ExportFunctions(b) + b.NewFunctionBuilder(). + WithFunc(func(_ context.Context, _ api.Module, from, to int32) int32 { + if to != 0 { + return errnoNotsup + } + path, ok := tracker.opened[from] + if !ok { + return errnoNotsup + } + if err := stdin.set(path); err != nil { + return errnoNotsup + } + return 0 + }). + Export("fd_renumber") + _, err := b.Instantiate(ctx) + return err +} diff --git a/go.mod b/go.mod index cdf8fc699..80c49ded0 100644 --- a/go.mod +++ b/go.mod @@ -32,6 +32,7 @@ require ( github.com/google/wire v0.7.0 github.com/gorilla/websocket v1.5.3 github.com/hashicorp/go-multierror v1.1.1 + github.com/jackc/pgx/v5 v5.10.0 github.com/jellydator/ttlcache/v3 v3.4.1 github.com/kardianos/service v1.3.0 github.com/kr/pretty v0.3.1 @@ -95,6 +96,9 @@ require ( github.com/hashicorp/errwrap v1.1.0 // indirect github.com/ianlancetaylor/demangle v0.0.0-20260724033716-83e58baca724 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51 // indirect github.com/klauspost/cpuid/v2 v2.4.0 // indirect github.com/kr/text v0.2.0 // indirect diff --git a/go.sum b/go.sum index d852285f1..14461daca 100644 --- a/go.sum +++ b/go.sum @@ -124,6 +124,14 @@ github.com/ianlancetaylor/demangle v0.0.0-20260724033716-83e58baca724 h1:QixF8Mc github.com/ianlancetaylor/demangle v0.0.0-20260724033716-83e58baca724/go.mod h1:gx7rwoVhcfuVKG5uya9Hs3Sxj7EIvldVofAWIUtGouw= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jellydator/ttlcache/v3 v3.4.1 h1:bOdXmXiycyK6E6Qjyuj5vl+/vU3SCOoDs8a86NbHjAQ= github.com/jellydator/ttlcache/v3 v3.4.1/go.mod h1:j7LO12PNghFg5+0v9budMAT4rDK4JY969jb9vOdOBBk= github.com/joshdk/go-junit v1.0.0 h1:S86cUKIdwBHWwA6xCmFlf3RTLfVXYQfvanM5Uh+K6GE= @@ -260,8 +268,10 @@ github.com/stretchr/objx v0.5.3 h1:jmXUvGomnU1o3W/V5h2VEradbpJDwGrzugQQvL0POH4= github.com/stretchr/objx v0.5.3/go.mod h1:rDQraq+vQZU7Fde9LOZLr8Tax6zZvy4kuNKF+QYS+U0= github.com/stretchr/testify v0.0.0-20161117074351-18a02ba4a312/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= @@ -340,8 +350,8 @@ google.golang.org/appengine v1.6.5/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCID google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127 h1:qIbj1fsPNlZgppZ+VLlY7N33q108Sa+fhmuc+sWQYwY= -gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/ini.v1 v1.67.3 h1:iM9Lhz5MRSGhHVGGwCuzG9KO8PoirCXj/m/qTmOJJQw= gopkg.in/ini.v1 v1.67.3/go.mod h1:x/cyOwCgZqOkJoDIJ3c1KNHMo10+nLGAhh+kn3Zizss= gopkg.in/natefinch/npipe.v2 v2.0.0-20160621034901-c1b8fa8bdcce h1:+JknDZhAj8YMt7GC73Ei8pv4MzjDUNPHgQWJdtMAaDU= diff --git a/persistence/sql_base_repository.go b/persistence/sql_base_repository.go index 5530d2568..23bddf7dc 100644 --- a/persistence/sql_base_repository.go +++ b/persistence/sql_base_repository.go @@ -537,9 +537,13 @@ func (r sqlRepository) classifyOwnedWriteMiss(id string) error { func (r sqlRepository) count(countQuery SelectBuilder, options ...model.QueryOptions) (int64, error) { countQuery = countQuery. RemoveColumns().Columns("count(distinct " + r.tableName + ".id) as count"). - RemoveOffset().RemoveLimit(). - OrderBy(r.tableName + ".id"). // To remove any ORDER BY clause that could slow down the query - From(r.tableName) + RemoveOffset().RemoveLimit() + if db.Dialect == "sqlite3" { + // To remove any ORDER BY clause that could slow down the query. Postgres rejects an + // ORDER BY on a column that is not grouped, so leave the clause off there. + countQuery = countQuery.OrderBy(r.tableName + ".id") + } + countQuery = countQuery.From(r.tableName) countQuery = r.applyFilters(countQuery, options...) var res struct{ Count int64 } err := r.queryOne(countQuery, &res)