fix(agent-assets): save rule markdown atomically
This commit is contained in:
207
server/tests/test_agent_asset_rule_markdown_atomicity.py
Normal file
207
server/tests/test_agent_asset_rule_markdown_atomicity.py
Normal file
@@ -0,0 +1,207 @@
|
||||
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
|
||||
Reference in New Issue
Block a user