muse-agent-example/tests/skills/抽取作品知识/test_upgrade_work_lock_offline.py

134 lines
4.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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()