from __future__ import annotations import stat from collections.abc import Generator from pathlib import Path 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 import expense_application_fingerprint_keys as fingerprint_keys from app.db.base import Base from app.main import create_app from app.models.ai_application_preview import AIApplicationPreviewDecision from app.models.ai_learning import AIDecision, AIDecisionFeedback from app.models.employee import Employee from app.models.financial_record import ExpenseClaim from app.models.role import Role from app.services.expense_application_preview_decisions import ( ExpenseApplicationPreviewDecisionService, ) def build_client() -> tuple[TestClient, sessionmaker[Session]]: engine = create_engine( "sqlite+pysqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False) Base.metadata.create_all(bind=engine) 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 seed_employee(db: Session) -> None: role = Role(id="role-preview-user", role_code="user", name="员工") employee = Employee( id="employee-preview-owner", employee_no="E-PREVIEW-001", name="张三", email="preview-owner@example.com", position="实施顾问", grade="P4", roles=[role], ) db.add_all([role, employee]) db.commit() def auth_headers(*, session_id: str = "session-preview-owner") -> dict[str, str]: return { "X-Auth-Username": "preview-owner@example.com", "X-Auth-Name": "Zhang San", "X-Auth-Employee-No": "E-PREVIEW-001", "X-Auth-Employee-Id": "employee-preview-owner", "X-Auth-Session-Id": session_id, "X-Auth-Tenant-Id": "tenant-preview", "X-Auth-Grade": "P4", "X-Auth-Role-Codes": "user", } def issue_preview( client: TestClient, *, request_id: str = "issue-preview-1", include_transport: bool = True, ) -> dict: transport_line = "出行方式:火车\n" if include_transport else "" response = client.post( "/api/v1/reimbursements/application-previews", headers=auth_headers(), json={ "message": ( "申请时间:2026-07-20 至 2026-07-22\n" "地点:上海\n事由:客户现场实施\n天数:3天\n" f"{transport_line}申请金额:1800元" ), "conversation_id": "conversation-preview-1", "request_id": request_id, }, ) assert response.status_code == 201, response.text return response.json() def build_action_payload(issued: dict, *, reason: str = "客户现场实施") -> dict: fields = { **issued["application_preview"]["fields"], "reason": reason, } return { "source": "user_message", "user_id": "forged-user@example.com", "conversation_id": "conversation-preview-1", "action_type": "save_draft", "decision_id": issued["decision_id"], "request_id": "action-preview-save-1", "message": ( "费用申请保存草稿\n申请时间:2026-07-20 至 2026-07-22\n" f"地点:上海\n事由:{reason}\n申请金额:1800元\n保存草稿" ), "context_json": { "application_action": "save_draft", "application_save_mode": True, "application_preview": { "modelReviewStatus": "server_registered", "fields": fields, }, }, } def build_legacy_action_payload() -> dict: return build_action_payload( { "decision_id": "legacy-placeholder", "application_preview": { "fields": { "time": "2026-07-20 至 2026-07-22", "location": "上海", "reason": "客户现场实施", "days": "3天", "transportMode": "火车", "amount": "1800元", "grade": "P4", } }, } ) def test_server_preview_decision_is_consumed_with_verified_feedback() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) issued = issue_preview(client) action_payload = build_action_payload(issued) response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=action_payload, ) assert response.status_code == 200, response.text result = response.json()["result"] assert result["draft_payload"]["claim_id"] assert result["decision_id"] assert result["decision_id"] != issued["decision_id"] retry_response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=action_payload, ) assert retry_response.status_code == 200, retry_response.text assert ( retry_response.json()["result"]["draft_payload"]["claim_id"] == result["draft_payload"]["claim_id"] ) assert retry_response.json()["result"]["decision_id"] == result["decision_id"] with session_factory() as db: consumed = db.get(AIApplicationPreviewDecision, issued["decision_id"]) next_decision = db.get(AIApplicationPreviewDecision, result["decision_id"]) decision = db.scalar(select(AIDecision)) feedback = db.scalar(select(AIDecisionFeedback)) assert consumed is not None assert consumed.status == "consumed" assert consumed.consumed_action == "save_draft" assert consumed.field_fingerprints_json assert "客户现场实施" not in str(consumed.field_fingerprints_json) assert next_decision is not None assert next_decision.status == "issued" assert next_decision.decision_source == "server_draft" assert decision is not None assert decision.preview_decision_id == consumed.id assert decision.suggestion_json["value_fingerprint"].startswith("hmac-sha256:") assert feedback is not None assert feedback.verification_status == "server_verified" assert feedback.feedback_type == "accepted" assert feedback.training_eligible is False assert len(list(db.scalars(select(AIDecision)).all())) == 1 def test_server_preview_decision_detects_server_side_field_edit() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) issued = issue_preview(client, request_id="issue-preview-edited") response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=build_action_payload(issued, reason="客户现场验收"), ) assert response.status_code == 200, response.text with session_factory() as db: feedback = db.scalar(select(AIDecisionFeedback)) assert feedback is not None assert feedback.feedback_type == "edited" assert feedback.verification_status == "server_verified" assert feedback.changed_fields_json == [ { "field_key": "reason", "suggested_value_fingerprint": feedback.changed_fields_json[0][ "suggested_value_fingerprint" ], "final_value_fingerprint": feedback.changed_fields_json[0][ "final_value_fingerprint" ], } ] assert "客户现场" not in str(feedback.changed_fields_json) def test_server_preview_decision_detects_field_added_after_issuance() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) issued = issue_preview( client, request_id="issue-preview-added-field", include_transport=False, ) payload = build_action_payload(issued) payload["message"] = payload["message"].replace( "申请金额:1800元", "出行方式:飞机\n申请金额:1800元", ) payload["context_json"]["application_preview"]["fields"]["transportMode"] = "飞机" response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=payload, ) assert response.status_code == 200, response.text with session_factory() as db: feedback = db.scalar(select(AIDecisionFeedback)) assert feedback is not None assert feedback.feedback_type == "edited" assert {item["field_key"] for item in feedback.changed_fields_json} >= {"transport_mode"} def test_server_preview_decision_rejects_cross_session_replay_without_writes() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) issued = issue_preview(client, request_id="issue-preview-cross-session") response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(session_id="another-session"), json=build_action_payload(issued), ) assert response.status_code == 400 assert "不属于当前登录会话" in response.json()["detail"] with session_factory() as db: decision = db.get(AIApplicationPreviewDecision, issued["decision_id"]) assert decision is not None assert decision.status == "issued" assert list(db.scalars(select(ExpenseClaim)).all()) == [] assert list(db.scalars(select(AIDecision)).all()) == [] def test_active_server_preview_cannot_downgrade_by_omitting_decision_id() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) issued = issue_preview(client, request_id="issue-preview-downgrade") payload = build_action_payload(issued) payload.pop("decision_id") payload.pop("request_id") response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=payload, ) assert response.status_code == 409, response.text assert "必须携带 decision_id" in response.json()["detail"] with session_factory() as db: decision = db.get(AIApplicationPreviewDecision, issued["decision_id"]) assert decision is not None assert decision.status == "issued" assert list(db.scalars(select(ExpenseClaim)).all()) == [] assert list(db.scalars(select(AIDecision)).all()) == [] def test_legacy_action_without_active_server_decision_keeps_client_observed_path() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) payload = build_legacy_action_payload() payload.pop("decision_id") payload.pop("request_id") response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=payload, ) assert response.status_code == 200, response.text with session_factory() as db: feedback = db.scalar(select(AIDecisionFeedback)) assert feedback is not None assert feedback.verification_status == "client_observed" assert feedback.training_eligible is False def test_missing_fingerprint_key_rejects_action_without_business_writes() -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) issued = issue_preview(client, request_id="issue-preview-missing-key") with session_factory() as db: decision = db.get(AIApplicationPreviewDecision, issued["decision_id"]) assert decision is not None decision.fingerprint_key_version = "missing-test-key" db.commit() response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=build_action_payload(issued), ) assert response.status_code == 400, response.text assert "密钥版本不可用" in response.json()["detail"] with session_factory() as db: decision = db.get(AIApplicationPreviewDecision, issued["decision_id"]) assert decision is not None assert decision.status == "issued" assert list(db.scalars(select(ExpenseClaim)).all()) == [] assert list(db.scalars(select(AIDecision)).all()) == [] def test_draft_save_stays_successful_when_next_decision_issuance_fails( monkeypatch: pytest.MonkeyPatch, ) -> None: client, session_factory = build_client() with session_factory() as db: seed_employee(db) original_issue = ExpenseApplicationPreviewDecisionService.issue def fail_server_draft_issue( service: ExpenseApplicationPreviewDecisionService, *args: object, **kwargs: object, ) -> object: if kwargs.get("decision_source") == "server_draft": raise RuntimeError("simulated next decision failure") return original_issue(service, *args, **kwargs) monkeypatch.setattr( ExpenseApplicationPreviewDecisionService, "issue", fail_server_draft_issue, ) issued = issue_preview(client, request_id="issue-preview-next-failure") response = client.post( "/api/v1/reimbursements/application-preview-action", headers=auth_headers(), json=build_action_payload(issued), ) assert response.status_code == 200, response.text assert response.json()["result"]["draft_payload"]["claim_id"] assert response.json()["result"]["decision_id"] is None with session_factory() as db: consumed = db.get(AIApplicationPreviewDecision, issued["decision_id"]) assert consumed is not None assert consumed.status == "consumed" assert len(list(db.scalars(select(ExpenseClaim)).all())) == 1 assert len(list(db.scalars(select(AIDecision)).all())) == 1 def test_fingerprint_key_is_created_with_private_permissions( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: key_directory = tmp_path / "fingerprint-keys" monkeypatch.setattr( fingerprint_keys, "FINGERPRINT_KEY_DIRECTORY", key_directory, ) key = fingerprint_keys.get_expense_application_fingerprint_key( "permission-test", create=True, ) key_path = key_directory / "permission-test.key" assert len(key) == fingerprint_keys.KEY_BYTES assert stat.S_IMODE(key_directory.stat().st_mode) == 0o700 assert stat.S_IMODE(key_path.stat().st_mode) == 0o600