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