Files
NianAIGC/lib/server/database.ts

166 lines
6.8 KiB
TypeScript

import "server-only";
import { Pool, type PoolClient, type PoolConfig, type QueryResult, type QueryResultRow } from "pg";
export type DataBackend = "local" | "postgres";
let pool: Pool | undefined;
export function getDataBackend(): DataBackend {
const configured = process.env.ZHINIAN_DATA_BACKEND?.trim().toLowerCase();
if (configured === "local" || configured === "postgres") return configured;
if (!configured && process.env.NODE_ENV !== "production") return "local";
throw new Error("ZHINIAN_DATA_BACKEND must be explicitly set to 'local' or 'postgres'");
}
export function isPostgresBackend(): boolean {
return getDataBackend() === "postgres";
}
export async function queryDatabase<T extends QueryResultRow>(
text: string,
values: readonly unknown[] = []
): Promise<QueryResult<T>> {
return getPool().query<T>(text, [...values]);
}
export async function withDatabaseTransaction<T>(fn: (client: PoolClient) => Promise<T>): Promise<T> {
const client = await getPool().connect();
try {
await client.query("BEGIN");
const result = await fn(client);
await client.query("COMMIT");
return result;
} catch (error) {
await client.query("ROLLBACK");
throw error;
} finally {
client.release();
}
}
export async function checkDatabaseReadiness(): Promise<void> {
if (!isPostgresBackend()) return;
const { rows } = await queryDatabase<{ ready: boolean }>(`
WITH required_table_privileges(table_name, privilege_name) AS (
VALUES
('assets', 'SELECT'), ('assets', 'INSERT'), ('assets', 'DELETE'),
('generation_jobs', 'SELECT'), ('generation_jobs', 'INSERT'), ('generation_jobs', 'UPDATE'), ('generation_jobs', 'DELETE'),
('usage_events', 'SELECT'), ('usage_events', 'INSERT'), ('usage_events', 'UPDATE'),
('projects', 'SELECT'), ('projects', 'UPDATE'),
('image_templates', 'SELECT'), ('image_templates', 'INSERT'), ('image_templates', 'UPDATE'), ('image_templates', 'DELETE'),
('platform_organizations', 'SELECT'), ('platform_organizations', 'INSERT'), ('platform_organizations', 'UPDATE'), ('platform_organizations', 'DELETE'),
('platform_users', 'SELECT'), ('platform_users', 'INSERT'), ('platform_users', 'UPDATE'), ('platform_users', 'DELETE'),
('platform_account_migrations', 'SELECT'), ('platform_account_migrations', 'INSERT'), ('platform_account_migrations', 'UPDATE'),
('billing_price_rules', 'SELECT'), ('billing_price_rules', 'INSERT'), ('billing_price_rules', 'UPDATE'),
('billing_wallets', 'SELECT'), ('billing_wallets', 'INSERT'), ('billing_wallets', 'UPDATE'),
('billing_ledger', 'SELECT'), ('billing_ledger', 'INSERT')
)
SELECT
NOT EXISTS (
SELECT 1
FROM required_table_privileges
WHERE to_regclass('public.' || table_name) IS NULL
OR NOT has_table_privilege(current_user, 'public.' || table_name, privilege_name)
)
AND has_function_privilege(
current_user,
'public.claim_generation_jobs(text,integer,integer)',
'EXECUTE'
)
AND has_function_privilege(
current_user,
'public.billing_post_wallet_entry(text,text,text,text,text,bigint,text,text,text,jsonb)',
'EXECUTE'
) AS ready
`);
if (!rows[0]?.ready) {
throw new Error("PostgreSQL schema or application privileges are not ready");
}
}
export function getDatabaseStatus(): { backend: DataBackend; configured: boolean } {
const backend = getDataBackend();
return { backend, configured: backend === "local" || Boolean(process.env.DATABASE_URL?.trim()) };
}
export async function closeDatabasePool(): Promise<void> {
const current = pool;
pool = undefined;
if (current) await current.end();
}
function getPool(): Pool {
if (!isPostgresBackend()) throw new Error("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=local");
if (pool) return pool;
const connectionString = process.env.DATABASE_URL?.trim();
if (!connectionString) throw new Error("DATABASE_URL is required when ZHINIAN_DATA_BACKEND=postgres");
const config: PoolConfig = {
...buildPlaintextPoolConfig(connectionString),
max: positiveInteger("DATABASE_POOL_MAX", 10),
idleTimeoutMillis: nonNegativeInteger("DATABASE_IDLE_TIMEOUT_MS", 30_000),
connectionTimeoutMillis: positiveInteger("DATABASE_CONNECTION_TIMEOUT_MS", 10_000),
statement_timeout: positiveInteger("DATABASE_STATEMENT_TIMEOUT_MS", 30_000),
application_name: process.env.DATABASE_APPLICATION_NAME?.trim() || "zhinian-web"
};
pool = new Pool(config);
pool.on("error", (error) => console.error("Unexpected PostgreSQL pool error", error));
return pool;
}
export function buildPlaintextPoolConfig(
connectionString: string
): Pick<PoolConfig, "connectionString" | "ssl"> {
const parsed = assertConnectionStringContract(connectionString);
for (const key of [...parsed.searchParams.keys()]) {
if (["sslmode", "sslrootcert"].includes(key.toLowerCase())) parsed.searchParams.delete(key);
}
parsed.searchParams.set("sslmode", "disable");
return { connectionString: parsed.toString(), ssl: false };
}
function assertConnectionStringContract(connectionString: string): URL {
let parsed: URL;
try {
parsed = new URL(connectionString);
} catch {
throw new Error("DATABASE_URL must be a valid PostgreSQL connection URI");
}
if (parsed.protocol !== "postgres:" && parsed.protocol !== "postgresql:") {
throw new Error("DATABASE_URL must use the postgres:// or postgresql:// scheme");
}
const sslParameters = [...parsed.searchParams.keys()].filter((key) => key.toLowerCase().startsWith("ssl"));
const unsupportedSSLParameters = sslParameters.filter((key) => !["sslmode", "sslrootcert"].includes(key.toLowerCase()));
if (unsupportedSSLParameters.length > 0) {
throw new Error(
`DATABASE_URL contains unsupported SSL query parameters (${unsupportedSSLParameters.join(", ")}); PostgreSQL transport is forced to sslmode=disable`
);
}
const sslMode = parsed.searchParams.get("sslmode")?.trim().toLowerCase();
if (sslMode && sslMode !== "disable" && sslMode !== "verify-full") {
throw new Error("DATABASE_URL sslmode must be 'disable' or 'verify-full'");
}
return parsed;
}
function positiveInteger(name: string, fallback: number): number {
const value = integer(name, fallback);
if (value <= 0) throw new Error(`${name} must be a positive integer`);
return value;
}
function nonNegativeInteger(name: string, fallback: number): number {
const value = integer(name, fallback);
if (value < 0) throw new Error(`${name} must be a non-negative integer`);
return value;
}
function integer(name: string, fallback: number): number {
const raw = process.env[name]?.trim();
if (!raw) return fallback;
const value = Number(raw);
if (!Number.isSafeInteger(value)) throw new Error(`${name} must be an integer`);
return value;
}