#!/usr/bin/env python3 """升格作品会话锁的纯离线测试:不连接数据库,不调用模型。""" import pathlib import sys import unittest PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3] SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "entity" / "skills" / "ingest" / "抽取作品知识" / "scripts" sys.path.insert(0, str(SCRIPT_DIR)) from upgrade_work_lock import ( # noqa: E402 UpgradeWorkLockUnavailable, advisory_lock_keys, upgrade_work_lock, ) class _Result: """模拟 psycopg 查询结果。""" def __init__(self, value): self.value = value def fetchone(self): return (self.value,) class _LockServer: """在内存中模拟 PostgreSQL 会话级 advisory lock。""" def __init__(self): self.held = {} self.connections = [] def connect(self, dsn, *, autocommit): connection = _LockConnection(self, dsn, autocommit) self.connections.append(connection) return connection class _LockConnection: """每个实例代表一条独立数据库会话,只接受锁相关 SQL。""" def __init__(self, server, dsn, autocommit): self.server = server self.dsn = dsn self.autocommit = autocommit self.owned = set() self.statements = [] self.closed = False def execute(self, sql, params): normalized = " ".join(sql.split()) key = tuple(params) self.statements.append((normalized, key)) if normalized == "SELECT pg_try_advisory_lock(%s, %s)": owner = self.server.held.get(key) if owner is not None and owner is not self: return _Result(False) self.server.held[key] = self self.owned.add(key) return _Result(True) if normalized == "SELECT pg_advisory_unlock(%s, %s)": owned = key in self.owned if owned: self.owned.remove(key) self.server.held.pop(key, None) return _Result(owned) raise AssertionError(f"锁连接禁止执行业务 SQL:{normalized}") def close(self): for key in list(self.owned): self.server.held.pop(key, None) self.owned.clear() self.closed = True class UpgradeWorkLockOfflineTest(unittest.TestCase): """覆盖锁键稳定性、互斥语义和 finally 释放。""" def test_keys_are_stable_signed_int32_pair(self): first = advisory_lock_keys(tenant_id=1, work_id=8) second = advisory_lock_keys(tenant_id=1, work_id=8) self.assertEqual(first, second) self.assertEqual(len(first), 2) self.assertTrue(all(-(2 ** 31) <= value < 2 ** 31 for value in first)) self.assertNotEqual(first, advisory_lock_keys(tenant_id=1, work_id=9)) self.assertNotEqual(first, advisory_lock_keys(tenant_id=2, work_id=8)) def test_same_work_concurrent_session_is_rejected(self): server = _LockServer() with upgrade_work_lock("postgresql://offline", 1, 8, connect=server.connect): with self.assertRaises(UpgradeWorkLockUnavailable): with upgrade_work_lock("postgresql://offline", 1, 8, connect=server.connect): self.fail("同书第二条会话不得进入受保护区") self.assertEqual(server.held, {}) self.assertTrue(all(connection.closed for connection in server.connections)) def test_exception_releases_lock_and_closes_autocommit_session(self): server = _LockServer() with self.assertRaisesRegex(RuntimeError, "业务失败"): with upgrade_work_lock("postgresql://offline", 1, 8, connect=server.connect): raise RuntimeError("业务失败") self.assertEqual(server.held, {}) self.assertEqual(len(server.connections), 1) connection = server.connections[0] self.assertTrue(connection.autocommit) self.assertTrue(connection.closed) self.assertEqual( [sql for sql, _ in connection.statements], ["SELECT pg_try_advisory_lock(%s, %s)", "SELECT pg_advisory_unlock(%s, %s)"], ) def test_different_work_does_not_conflict(self): server = _LockServer() entered = [] with upgrade_work_lock("postgresql://offline", 1, 8, connect=server.connect): with upgrade_work_lock("postgresql://offline", 1, 9, connect=server.connect): entered.append(True) self.assertEqual(entered, [True]) self.assertEqual(server.held, {}) if __name__ == "__main__": unittest.main()