Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 17 additions & 3 deletions src/cli.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,13 @@
import yargs from "yargs"
import { migrate, reset, generate, createMigration, initPgstrap } from "./"
import { getProjectContext } from "./get-project-context"
;(yargs as any)

const cli =
typeof (yargs as any).command === "function"
? (yargs as any)
: (yargs as any)(process.argv.slice(2))

cli
.command("init", "initialize pgstrap", {}, async () => {
await initPgstrap({
cwd: process.cwd(),
Expand Down Expand Up @@ -35,10 +41,18 @@ import { getProjectContext } from "./get-project-context"
"generate",
"generate types and sql documentation from database",
(yargs) => {
yargs.option("pglite", { type: "boolean", default: false })
yargs.option("pglite", {
type: "boolean",
default: true,
describe:
"use pglite to run migrations and generate types (use --no-pglite to connect to external Postgres)",
})
},
async (argv) => {
generate({ ...(await getProjectContext()), pglite: !!argv.pglite })
await generate({
...(await getProjectContext()),
pglite: argv.pglite !== false,
})
},
)
.parse()
41 changes: 22 additions & 19 deletions src/generate.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ export const generate = async ({
schemas,
defaultDatabase,
dbDir,
pglite = false,
pglite = true,
migrationsDir,
}: Pick<Context, "schemas" | "defaultDatabase" | "dbDir"> & {
pglite?: boolean
Expand Down Expand Up @@ -67,25 +67,28 @@ export const generate = async ({
const prevDbUrl = process.env.DATABASE_URL
process.env.DATABASE_URL = connectionString

await zg.generate({
db: {
connectionString,
},
schemas: Object.fromEntries(
schemas.map((s) => [s, { include: "*", exclude: [] }]),
),
outDir: dbDir,
})

await dumpTree({
targetDir: path.join(dbDir, "structure"),
defaultDatabase: "postgres",
schemas,
})
try {
await zg.generate({
db: {
connectionString,
},
schemas: Object.fromEntries(
schemas.map((s) => [s, { include: "*", exclude: [] }]),
),
outDir: dbDir,
})

server.close()
if (prevDbUrl === undefined) delete process.env.DATABASE_URL
else process.env.DATABASE_URL = prevDbUrl
await dumpTree({
targetDir: path.join(dbDir, "structure"),
defaultDatabase: "postgres",
schemas,
})
} finally {
server.close()
if (prevDbUrl === undefined) delete process.env.DATABASE_URL
else process.env.DATABASE_URL = prevDbUrl
await (db as any).close?.()
}
return
}

Expand Down
130 changes: 129 additions & 1 deletion tests/generate.pglite.test.ts
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
import { test, expect } from "bun:test"
import { test, expect, spyOn } from "bun:test"
import fs from "fs"
import os from "os"
import path from "path"
import { generate } from "../src/generate"
import * as zg from "zapatos/generate"
import * as pgSchemaDump from "pg-schema-dump"

const migrationFile = `
exports.up = async (pgm) => {
Expand Down Expand Up @@ -45,3 +47,129 @@ test("generate with pglite runs migrations and dumps structure", async () => {

fs.rmSync(tmp, { recursive: true, force: true })
})

test("generate defaults to pglite when pglite option is omitted", async () => {
const tmp = fs.mkdtempSync(
path.join(os.tmpdir(), "pgstrap-generate-default-"),
)
const migrationsDir = path.join(tmp, "migrations")
fs.mkdirSync(migrationsDir, { recursive: true })
fs.writeFileSync(
path.join(migrationsDir, "001_create_table.js"),
migrationFile,
)

const originalDbUrl = process.env.DATABASE_URL

// pglite is omitted; should default to true and not require external postgres
await generate({
schemas: ["public"],
defaultDatabase: "postgres",
dbDir: path.join(tmp, "db"),
migrationsDir,
})

const zapatosFile = path.join(tmp, "db", "zapatos", "schema.d.ts")
const structureDir = path.join(
tmp,
"db",
"structure",
"public",
"tables",
"foo",
)

expect(fs.existsSync(zapatosFile)).toBe(true)
expect(fs.existsSync(path.join(structureDir, "table.sql"))).toBe(true)
expect(process.env.DATABASE_URL).toBe(originalDbUrl)

fs.rmSync(tmp, { recursive: true, force: true })
})

test("generate with pglite: false opts out and uses connectionString from env", async () => {
const tmp = fs.mkdtempSync(path.join(os.tmpdir(), "pgstrap-generate-optout-"))
const migrationsDir = path.join(tmp, "migrations")
fs.mkdirSync(migrationsDir, { recursive: true })
fs.writeFileSync(
path.join(migrationsDir, "001_create_table.js"),
migrationFile,
)

const originalDbUrl = process.env.DATABASE_URL
const customDbUrl =
"postgres://custom_user:custom_pass@127.0.0.1:5432/custom_db"
process.env.DATABASE_URL = customDbUrl

const zgSpy = spyOn(zg, "generate").mockImplementation(async () => {})
const dumpSpy = spyOn(pgSchemaDump, "dumpTree").mockImplementation(
async () => {},
)

try {
await generate({
schemas: ["public"],
defaultDatabase: "custom_db",
dbDir: path.join(tmp, "db"),
migrationsDir,
pglite: false,
})

expect(zgSpy).toHaveBeenCalled()
const lastCallArg = zgSpy.mock.calls[0][0] as any
expect(lastCallArg.db.connectionString).toContain(
"custom_user:custom_pass@127.0.0.1:5432/custom_db",
)
expect(process.env.DATABASE_URL).toBe(customDbUrl)
} finally {
zgSpy.mockRestore()
dumpSpy.mockRestore()
if (originalDbUrl === undefined) delete process.env.DATABASE_URL
else process.env.DATABASE_URL = originalDbUrl
fs.rmSync(tmp, { recursive: true, force: true })
}
})

test("DATABASE_URL is restored even if generate encounters an error", async () => {
const tmp = fs.mkdtempSync(path.join(os.tmpdir(), "pgstrap-generate-error-"))
const migrationsDir = path.join(tmp, "migrations")
fs.mkdirSync(migrationsDir, { recursive: true })
fs.writeFileSync(
path.join(migrationsDir, "001_create_table.js"),
migrationFile,
)

const originalDbUrl = "postgres://original:original@localhost:5432/original"
process.env.DATABASE_URL = originalDbUrl

const zgSpy = spyOn(zg, "generate").mockImplementation(async () => {
// Verify DATABASE_URL was set to the temporary PGlite socket connection during generation
expect(process.env.DATABASE_URL).toContain(
"postgres://postgres:postgres@127.0.0.1:",
)
throw new Error("Simulated zg.generate failure")
})

try {
let thrownError: any = null
try {
await generate({
schemas: ["public"],
defaultDatabase: "postgres",
dbDir: path.join(tmp, "db"),
migrationsDir,
pglite: true,
})
} catch (err) {
thrownError = err
}

expect(thrownError).not.toBeNull()
expect(thrownError.message).toBe("Simulated zg.generate failure")
// Verify DATABASE_URL is cleanly restored to original value
expect(process.env.DATABASE_URL).toBe(originalDbUrl)
} finally {
zgSpy.mockRestore()
delete process.env.DATABASE_URL
fs.rmSync(tmp, { recursive: true, force: true })
}
})