diff --git a/src/lock.ts b/src/lock.ts index 2098600..c65a146 100644 --- a/src/lock.ts +++ b/src/lock.ts @@ -16,16 +16,20 @@ export async function withLock(key: string, fn: () => Promise): Promise const prev = locks.get(key) ?? Promise.resolve() let release!: () => void const next = new Promise(r => (release = r)) - locks.set( - key, - prev.then(() => next), - ) + const chained = prev.then(() => next) + locks.set(key, chained) await prev.catch(() => {}) // wait our turn; ignore prior errors try { return await fn() } finally { release() - // Clean up if we're the tail of the chain to avoid unbounded growth. - if (locks.get(key) === next) locks.delete(key) + // Clean up if we're the tail of the chain to avoid unbounded growth. The + // map holds the chained promise, so the tail check must compare against it. + if (locks.get(key) === chained) locks.delete(key) } } + +/** Number of keys currently held in the lock table. For tests/observability. */ +export function activeLockCount(): number { + return locks.size +} diff --git a/tests/lock.test.ts b/tests/lock.test.ts new file mode 100644 index 0000000..b8bd7b4 --- /dev/null +++ b/tests/lock.test.ts @@ -0,0 +1,40 @@ +import { describe, expect, test } from 'bun:test' +import { activeLockCount, withLock } from '../src/lock.ts' + +describe('withLock', () => { + test('serializes critical sections for the same key', async () => { + const order: number[] = [] + let running = 0 + let overlap = false + const task = (n: number) => + withLock('k', async () => { + running++ + if (running > 1) overlap = true + await Promise.resolve() + order.push(n) + running-- + }) + await Promise.all([task(1), task(2), task(3)]) + expect(overlap).toBe(false) + expect(order).toEqual([1, 2, 3]) + }) + + test('does not leak lock-table entries once sections complete', async () => { + // Distinct keys, each fully awaited: the table must return to empty. + for (let i = 0; i < 100; i++) { + await withLock(`ns:${i}`, async () => i) + } + // Concurrent contention on one key must also drain. + await Promise.all(Array.from({ length: 20 }, () => withLock('hot', async () => 1))) + expect(activeLockCount()).toBe(0) + }) + + test('cleans up even when the critical section throws', async () => { + await expect( + withLock('boom', async () => { + throw new Error('fail') + }), + ).rejects.toThrow('fail') + expect(activeLockCount()).toBe(0) + }) +})