diff --git a/packages/adapter-pg/package.json b/packages/adapter-pg/package.json index 18c99696018f..2440e6b6b993 100644 --- a/packages/adapter-pg/package.json +++ b/packages/adapter-pg/package.json @@ -37,6 +37,7 @@ "sideEffects": false, "dependencies": { "@prisma/driver-adapter-utils": "workspace:*", + "async-mutex": "0.5.0", "pg": "^8.16.3", "postgres-array": "3.0.4", "@types/pg": "^8.16.0" diff --git a/packages/adapter-pg/src/__tests__/pg.test.ts b/packages/adapter-pg/src/__tests__/pg.test.ts index 451005039372..7c02f75ec2e0 100644 --- a/packages/adapter-pg/src/__tests__/pg.test.ts +++ b/packages/adapter-pg/src/__tests__/pg.test.ts @@ -134,6 +134,64 @@ describe('PrismaPgAdapterFactory', () => { await adapter.dispose() }) + it('should serialize concurrent queries within a transaction', async () => { + const config: pg.PoolConfig = { user: 'test', password: 'test', database: 'test', port: 5432, host: 'localhost' } + const factory = new PrismaPgAdapterFactory(config) + const adapter = await factory.connect() + + let inFlight = 0 + let maxInFlight = 0 + const mockConnection = { + on: vi.fn(), + removeListener: vi.fn(), + query: vi.fn(async () => { + inFlight++ + maxInFlight = Math.max(maxInFlight, inFlight) + await new Promise((resolve) => setTimeout(resolve, 10)) + inFlight-- + return { rows: [], fields: [], rowCount: 0 } + }), + release: vi.fn(), + listenerCount: vi.fn().mockReturnValue(0), + } + adapter['client'].connect = vi.fn().mockResolvedValue(mockConnection) + + const transaction = await adapter.startTransaction() + const query: SqlQuery = { sql: 'SELECT 1', args: [], argTypes: [] } + await Promise.all([transaction.queryRaw(query), transaction.queryRaw(query), transaction.queryRaw(query)]) + + expect(maxInFlight).toBe(1) + expect(mockConnection.query).toHaveBeenCalledTimes(4) // BEGIN + 3 queries + + await transaction.commit() + await adapter.dispose() + }) + + it('should release the transaction mutex when a query fails', async () => { + const config: pg.PoolConfig = { user: 'test', password: 'test', database: 'test', port: 5432, host: 'localhost' } + const factory = new PrismaPgAdapterFactory(config) + const adapter = await factory.connect() + + const mockConnection = { + on: vi.fn(), + removeListener: vi.fn(), + query: vi.fn().mockResolvedValue({ rows: [], fields: [], rowCount: 0 }), + release: vi.fn(), + listenerCount: vi.fn().mockReturnValue(0), + } + adapter['client'].connect = vi.fn().mockResolvedValue(mockConnection) + + const transaction = await adapter.startTransaction() + mockConnection.query.mockRejectedValueOnce(new Error('boom')) + + const query: SqlQuery = { sql: 'SELECT 1', args: [], argTypes: [] } + await expect(transaction.queryRaw(query)).rejects.toThrow() + await expect(transaction.queryRaw(query)).resolves.toBeDefined() + + await transaction.rollback() + await adapter.dispose() + }) + it('should not pass name when statement name generator is not provided', async () => { const factory = new PrismaPgAdapterFactory('postgresql://test:test@localhost/test') const adapter = await factory.connect() diff --git a/packages/adapter-pg/src/pg.ts b/packages/adapter-pg/src/pg.ts index 6a677bdad614..5a0775dfe604 100644 --- a/packages/adapter-pg/src/pg.ts +++ b/packages/adapter-pg/src/pg.ts @@ -13,6 +13,7 @@ import type { TransactionOptions, } from '@prisma/driver-adapter-utils' import { Debug, DriverAdapterError } from '@prisma/driver-adapter-utils' +import { Mutex } from 'async-mutex' // @ts-ignore: this is used to avoid the `Module '"/node_modules/@types/pg/index"' has no default export.` error. import pg from 'pg' @@ -98,7 +99,7 @@ class PgQueryable implements SqlQ * Should the query fail due to a connection error, the connection is * marked as unhealthy. */ - private async performIO(query: SqlQuery): Promise> { + protected async performIO(query: SqlQuery): Promise> { const { sql, args } = query const values = args.map((arg, i) => mapArg(arg, query.argTypes[i])) @@ -132,6 +133,10 @@ class PgQueryable implements SqlQ } class PgTransaction extends PgQueryable implements Transaction { + // pg.PoolClient does not support concurrent queries on the same connection, + // so we serialize all performIO calls with a mutex. + #mutex = new Mutex() + constructor( client: pg.PoolClient, readonly options: TransactionOptions, @@ -141,6 +146,10 @@ class PgTransaction extends PgQueryable implements Transactio super(client, pgOptions) } + protected async performIO(query: SqlQuery): Promise> { + return this.#mutex.runExclusive(() => super.performIO(query)) + } + async commit(): Promise { debug(`[js::commit]`) diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 329dc31e3b16..7f46d3c62985 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -315,6 +315,9 @@ importers: '@types/pg': specifier: ^8.16.0 version: 8.20.0 + async-mutex: + specifier: 0.5.0 + version: 0.5.0 pg: specifier: ^8.16.3 version: 8.16.3