208 lines
7.0 KiB
Python
208 lines
7.0 KiB
Python
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
|