diff --git a/src/cli.ts b/src/cli.ts index 9a9bdec..980498d 100644 --- a/src/cli.ts +++ b/src/cli.ts @@ -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(), @@ -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() diff --git a/src/generate.ts b/src/generate.ts index f337094..2fb3ddd 100644 --- a/src/generate.ts +++ b/src/generate.ts @@ -12,7 +12,7 @@ export const generate = async ({ schemas, defaultDatabase, dbDir, - pglite = false, + pglite = true, migrationsDir, }: Pick & { pglite?: boolean @@ -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 } diff --git a/tests/generate.pglite.test.ts b/tests/generate.pglite.test.ts index 56dcd53..70180d9 100644 --- a/tests/generate.pglite.test.ts +++ b/tests/generate.pglite.test.ts @@ -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) => { @@ -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 }) + } +})