From 98555ca499e03509b7a548a4a4aa4eafe6182e71 Mon Sep 17 00:00:00 2001 From: Deluan Date: Fri, 4 Sep 2026 18:29:58 -0400 Subject: [PATCH] feat(db): spike an embedded PostgreSQL via PGlite under wazero MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds db/pglite, a bridge that runs the PGlite WASI build of PostgreSQL 17.5 inside the Navidrome process with wazero and exposes it to database/sql through pgx over the real wire protocol. Selected with DbPath = "pglite://". The bridge handles what the WASI build cannot do on its own: it emulates dup2 onto stdin (wazero refuses fd_renumber on preopens), replays ParameterStatus for connections after the first, synthesizes a missing ReadyForQuery after an error, recovers from the trap every PG ERROR raises, and passes wire bytes through shared wasm memory instead of files (round-trip floor 377 µs to 39 µs). Several connections are accepted and serialized onto the single backend, holding the session across a whole transaction and a whole handshake; a client that disconnects mid-transaction gets a ROLLBACK. db.go opens the pglite:// scheme, skips the SQLite-only migrations, and can apply a translated schema from ND_PGLITE_SCHEMA. The shared count() helper now adds its ORDER BY only for SQLite, since Postgres rejects it next to count(distinct ...). The wasm archive is not committed. The README explains how to build it and lists the limits found: one session, one process (DevExternalScanner must be off), shared session state, simple protocol only. --- db/db.go | 95 +++- db/pglite/README.md | 52 ++ db/pglite/concurrency_test.go | 194 +++++++ db/pglite/pglite.go | 801 +++++++++++++++++++++++++++++ db/pglite/pglite_suite_test.go | 17 + db/pglite/pglite_test.go | 105 ++++ db/pglite/proto_test.go | 76 +++ db/pglite/setup.go | 118 +++++ db/pglite/stdin.go | 125 +++++ go.mod | 4 + go.sum | 14 +- persistence/sql_base_repository.go | 10 +- 12 files changed, 1605 insertions(+), 6 deletions(-) create mode 100644 db/pglite/README.md create mode 100644 db/pglite/concurrency_test.go create mode 100644 db/pglite/pglite.go create mode 100644 db/pglite/pglite_suite_test.go create mode 100644 db/pglite/pglite_test.go create mode 100644 db/pglite/proto_test.go create mode 100644 db/pglite/setup.go create mode 100644 db/pglite/stdin.go 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)