Files
X-Financial/server/tests/test_agent_asset_rule_markdown_atomicity.py

208 lines
7.0 KiB
Python
Raw Permalink Normal View History

from __future__ import annotations
from collections.abc import Generator
from copy import deepcopy
import pytest
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.api.deps import get_db
from app.core.agent_enums import AgentAssetType
from app.db.base import Base
from app.main import create_app
from app.models.agent_asset import AgentAsset, AgentAssetVersion
from app.models.audit_log import AuditLog
from app.schemas.agent_asset import RuleMarkdownUpdate
from app.services.agent_assets import AgentAssetService
def build_client() -> tuple[TestClient, sessionmaker[Session]]:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()
try:
yield db
finally:
db.close()
app.dependency_overrides[get_db] = override_db
return TestClient(app), session_factory
def _finance_headers() -> dict[str, str]:
return {
"x-auth-username": "finance",
"x-auth-name": "Finance",
"x-auth-role-codes": "finance",
"x-auth-is-admin": "true",
}
def _markdown_rule(db: Session) -> AgentAsset:
service = AgentAssetService(db)
service.list_assets(asset_type=AgentAssetType.RULE.value)
return db.scalar(
select(AgentAsset)
.where(AgentAsset.asset_type == AgentAssetType.RULE.value)
.where(AgentAsset.config_json["detail_mode"].as_string().is_(None))
.order_by(AgentAsset.code)
)
def _payload(version: str) -> RuleMarkdownUpdate:
runtime_rule = {
"kind": "policy_rule_draft",
"version": 2,
"template_key": "general_policy_v1",
"rule_name": "原子保存验证规则",
"scenario": "expense",
"review_required": True,
}
return RuleMarkdownUpdate(
version=version,
content=(
"# 原子保存验证规则\n\n"
"```expense-rule\n"
'{"kind":"policy_rule_draft","version":2,'
'"template_key":"general_policy_v1",'
'"rule_name":"原子保存验证规则",'
'"scenario":"expense","review_required":true}\n'
"```"
),
config_json={
"runtime_kind": "policy_rule_draft",
"runtime_rule": runtime_rule,
"rule_template_key": "general_policy_v1",
"rule_template_label": "通用制度模板",
"unrelated_client_field": "must-not-overwrite-server-config",
},
change_note="验证 Markdown 与运行配置原子保存。",
created_by="untrusted-client-actor",
)
def test_save_rule_markdown_endpoint_commits_version_and_runtime_config_together() -> None:
client, session_factory = build_client()
with session_factory() as db:
rule = _markdown_rule(db)
assert rule is not None
rule_id = rule.id
published_version = rule.published_version
response = client.post(
f"/api/v1/agent-assets/{rule_id}/rule-markdown",
headers=_finance_headers(),
json=_payload("v9.0.1").model_dump(),
)
assert response.status_code == 201, response.text
assert response.json()["version"] == "v9.0.1"
with session_factory() as db:
stored = db.get(AgentAsset, rule_id)
assert stored is not None
assert stored.current_version == "v9.0.1"
assert stored.working_version == "v9.0.1"
assert stored.published_version == published_version
assert stored.config_json["runtime_rule"]["version"] == 2
assert "unrelated_client_field" not in stored.config_json
version = db.scalar(
select(AgentAssetVersion).where(
AgentAssetVersion.asset_id == rule_id,
AgentAssetVersion.version == "v9.0.1",
)
)
assert version is not None
assert version.created_by == "username:finance"
def test_save_rule_markdown_rolls_back_version_and_config_when_audit_fails(
monkeypatch,
) -> None:
_, session_factory = build_client()
with session_factory() as db:
rule = _markdown_rule(db)
assert rule is not None
original_current_version = rule.current_version
original_working_version = rule.working_version
original_config = deepcopy(rule.config_json)
service = AgentAssetService(db)
def fail_audit(**_kwargs) -> None:
raise RuntimeError("injected audit failure")
monkeypatch.setattr(service.audit_service, "log_action", fail_audit)
with pytest.raises(RuntimeError, match="injected audit failure"):
service.save_rule_markdown(
rule.id,
_payload("v9.0.2"),
actor="username:finance",
)
db.expire_all()
stored = db.get(AgentAsset, rule.id)
assert stored is not None
assert stored.current_version == original_current_version
assert stored.working_version == original_working_version
assert stored.config_json == original_config
assert db.scalar(
select(AgentAssetVersion).where(
AgentAssetVersion.asset_id == rule.id,
AgentAssetVersion.version == "v9.0.2",
)
) is None
def test_save_rule_markdown_rolls_back_when_response_build_fails(monkeypatch) -> None:
_, session_factory = build_client()
with session_factory() as db:
rule = _markdown_rule(db)
assert rule is not None
original_current_version = rule.current_version
original_config = deepcopy(rule.config_json)
service = AgentAssetService(db)
def fail_response(*_args) -> None:
raise RuntimeError("injected response failure")
monkeypatch.setattr(service, "_serialize_version", fail_response)
with pytest.raises(RuntimeError, match="injected response failure"):
service.save_rule_markdown(
rule.id,
_payload("v9.0.3"),
actor="username:finance",
)
db.expire_all()
stored = db.get(AgentAsset, rule.id)
assert stored is not None
assert stored.current_version == original_current_version
assert stored.config_json == original_config
assert db.scalar(
select(AgentAssetVersion).where(
AgentAssetVersion.asset_id == rule.id,
AgentAssetVersion.version == "v9.0.3",
)
) is None
assert db.scalar(
select(AuditLog).where(
AuditLog.resource_id == rule.id,
AuditLog.action == "save_rule_markdown",
)
) is None