import asyncio, sqlite3
from reserve import SCHEMA, reserve

class Gate:
    def __init__(self, n): self.n=n; self.count=0; self.event=asyncio.Event()
    async def __call__(self):
        self.count += 1
        if self.count == self.n: self.event.set()
        await self.event.wait()

def db_with(stock):
    db=sqlite3.connect(':memory:'); db.executescript(SCHEMA)
    db.executemany('INSERT INTO stock VALUES(?,?,?)', stock); db.commit(); return db

async def main():
    db=db_with([('t','s',1)]); gate=Gate(2)
    results=await asyncio.gather(reserve(db,'t','r1','s',1,gate), reserve(db,'t','r2','s',1,gate), return_exceptions=True)
    print('oversell:', results, 'passed_SELECT=',gate.count, 'stock=',db.execute("select qty from stock").fetchone()[0], 'reservations=',db.execute('select request_id from reservations order by request_id').fetchall())
    async def no_wait(): pass
    db=db_with([('t','s',5)])
    print('payload-replay:', await reserve(db,'t','same','s',1,no_wait), await reserve(db,'t','same','s',2,no_wait), 'stock=',db.execute('select qty from stock').fetchone()[0], 'stored=',db.execute('select tenant,sku,qty from reservations').fetchall())
    db=db_with([('t1','s',3),('t2','s',3)])
    print('cross-tenant:', await reserve(db,'t1','shared','s',1,no_wait), await reserve(db,'t2','shared','s',1,no_wait), 'stock=',db.execute('select tenant,qty from stock order by tenant').fetchall(), 'rows=',db.execute('select tenant,request_id from reservations').fetchall())
    db=db_with([('t','s',4)])
    db.execute("CREATE TRIGGER fail_insert BEFORE INSERT ON reservations BEGIN SELECT RAISE(ABORT,'injected insert fault'); END")
    reached=[]
    try: await reserve(db,'t','retry','s',2,no_wait)
    except sqlite3.IntegrityError as e: reached.append(str(e))
    print('insert-fault-control:', 'reached=',reached, 'stock_after_failure=',db.execute('select qty from stock').fetchone()[0], 'rows_after_failure=',db.execute('select count(*) from reservations').fetchone()[0])
    db.execute('DROP TRIGGER fail_insert')
    print('retry-after-fault-removed:', await reserve(db,'t','retry','s',2,no_wait), 'stock=',db.execute('select qty from stock').fetchone()[0], 'rows=',db.execute('select count(*) from reservations').fetchone()[0])

asyncio.run(main())
