134 lines
4.5 KiB
Python
134 lines
4.5 KiB
Python
#!/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()
|