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

436 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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_same_snapshot_with_different_issue_request_reuses_active_decision() -> None:
client, session_factory = build_client()
with session_factory() as db:
seed_employee(db)
first = issue_preview(client, request_id="issue-preview-same-snapshot-1")
second = issue_preview(client, request_id="issue-preview-same-snapshot-2")
assert second["decision_id"] == first["decision_id"]
with session_factory() as db:
decisions = list(db.scalars(select(AIApplicationPreviewDecision)).all())
assert len(decisions) == 1
assert decisions[0].status == "issued"
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