feat(db): spike an embedded PostgreSQL via PGlite under wazero

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://<dir>".

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.
This commit is contained in:
Deluan 2026-09-04 18:29:58 -04:00
commit 98555ca499
12 changed files with 1605 additions and 6 deletions

View file

@ -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() {

52
db/pglite/README.md Normal file
View file

@ -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://<dir>"`.
**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 |

View file

@ -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))
})
})

801
db/pglite/pglite.go Normal file
View file

@ -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, " ")
}

View file

@ -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")
}

105
db/pglite/pglite_test.go Normal file
View file

@ -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))
})
})

76
db/pglite/proto_test.go Normal file
View file

@ -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()
})
})

118
db/pglite/setup.go Normal file
View file

@ -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
}

125
db/pglite/stdin.go Normal file
View file

@ -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
}

4
go.mod
View file

@ -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

14
go.sum
View file

@ -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=

View file

@ -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)