from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal import pytest from sqlalchemy import create_engine, select from sqlalchemy.orm import Session, sessionmaker from sqlalchemy.pool import StaticPool from app.db.base import Base from app.models.ai_learning import AIDecision, AIDecisionFeedback, WorkflowOutcome from app.models.employee import Employee from app.models.expense_case import BusinessEvent, ExpenseCaseLink from app.models.financial_record import ExpenseClaim from app.schemas.ontology import OntologyParseResult from app.schemas.user_agent import UserAgentRequest from app.services.expense_cases import ExpenseCaseService from app.services.user_agent import UserAgentService def build_session_factory() -> sessionmaker[Session]: engine = create_engine( "sqlite+pysqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) Base.metadata.create_all(bind=engine) return sessionmaker(bind=engine, autoflush=False, autocommit=False) def build_request( *, run_id: str, tenant_id: str = "tenant-a", changed_fields: list[dict[str, str]] | None = None, ) -> UserAgentRequest: return UserAgentRequest( run_id=run_id, user_id="owner@example.com", message="保存申请草稿", ontology=OntologyParseResult(run_id=run_id), context_json={ "session_type": "application", "tenant_id": tenant_id, "name": "张三", "employee_no": "E001", "department_name": "市场部", "application_preview": { "modelReviewStatus": "completed", "fields": { "applicationType": "差旅费用申请", "time": "2026-07-13 至 2026-07-14", "location": "上海", "reason": "客户现场差旅", "amount": "880元", }, "aiDecisionFeedback": { "schemaVersion": 1, "changedFields": changed_fields or [], }, }, }, ) def build_facts(*, reason: str = "客户现场差旅") -> dict[str, str]: return { "application_type": "差旅费用申请", "time": "2026-07-13 至 2026-07-14", "location": "上海", "reason": reason, "amount": "880元", } def test_ai_application_draft_update_writes_tenant_scoped_event() -> None: session_factory = build_session_factory() with session_factory() as db: owner = Employee( id="owner-1", employee_no="E001", name="张三", email="owner@example.com", ) claim = ExpenseClaim( id="application-1", claim_no="AP-DRAFT-001", employee_id=owner.id, employee_name=owner.name, department_name="市场部", expense_type="travel_application", reason="原申请事由", location="北京", amount=Decimal("500.00"), currency="CNY", invoice_count=0, occurred_at=datetime(2026, 7, 10, tzinfo=UTC), submitted_at=None, status="returned", approval_stage="退回补充", risk_flags_json=[], ) db.add_all([owner, claim]) db.commit() service = UserAgentService(db) request = build_request(run_id="application-draft-update") updated = service._update_expense_application_record( request, build_facts(reason="更新后的申请事由"), claim, submit=False, learning_current_user=service._build_application_current_user(request), ) assert updated.status == "draft" event = db.scalar(select(BusinessEvent).where(BusinessEvent.aggregate_id == claim.id)) assert event is not None assert event.event_type == "claim_draft_updated" assert event.tenant_id == "tenant-a" assert event.actor_id == "owner@example.com" assert event.correlation_id == "application-draft-update" assert event.payload_json["previous_status"] == "returned" assert event.payload_json["next_status"] == "draft" link = db.scalar(select(ExpenseCaseLink).where(ExpenseCaseLink.resource_id == claim.id)) assert link is not None assert link.tenant_id == "tenant-a" decision = db.scalar(select(AIDecision).where(AIDecision.expense_claim_id == claim.id)) feedback = db.scalar( select(AIDecisionFeedback).where(AIDecisionFeedback.expense_claim_id == claim.id) ) outcome = db.scalar( select(WorkflowOutcome).where(WorkflowOutcome.expense_claim_id == claim.id) ) assert decision is not None assert decision.business_event_id == event.id assert feedback is not None assert feedback.feedback_type == "accepted" assert outcome is not None assert outcome.outcome_type == "draft_saved" def test_untrusted_application_path_does_not_write_learning_ledger() -> None: session_factory = build_session_factory() with session_factory() as db: claim = UserAgentService(db)._create_expense_application_record( build_request(run_id="application-untrusted-path"), build_facts(), submit=False, ) assert db.get(ExpenseClaim, claim.id) is not None assert db.scalar( select(BusinessEvent).where(BusinessEvent.aggregate_id == claim.id) ) is not None assert list(db.scalars(select(AIDecision)).all()) == [] assert list(db.scalars(select(AIDecisionFeedback)).all()) == [] assert list(db.scalars(select(WorkflowOutcome)).all()) == [] @pytest.mark.parametrize( "preview_patch", [ {"modelReviewStatus": "template"}, {"applicationEditMode": True}, ], ) def test_non_ai_or_detail_preview_is_not_recorded_as_ai_feedback( preview_patch: dict[str, object], ) -> None: session_factory = build_session_factory() with session_factory() as db: request = build_request(run_id="application-ineligible-preview") request.context_json["application_preview"].update(preview_patch) service = UserAgentService(db) claim = service._create_expense_application_record( request, build_facts(), submit=False, learning_current_user=service._build_application_current_user(request), ) assert db.get(ExpenseClaim, claim.id) is not None assert list(db.scalars(select(AIDecision)).all()) == [] assert list(db.scalars(select(AIDecisionFeedback)).all()) == [] assert list(db.scalars(select(WorkflowOutcome)).all()) == [] def test_ai_application_draft_creation_rolls_back_when_event_write_fails( monkeypatch: pytest.MonkeyPatch, ) -> None: session_factory = build_session_factory() def fail_event(*_args, **_kwargs): raise RuntimeError("模拟草稿事件写入失败") monkeypatch.setattr(ExpenseCaseService, "record_claim_event", fail_event) with session_factory() as db: with pytest.raises(RuntimeError, match="模拟草稿事件写入失败"): UserAgentService(db)._create_expense_application_record( build_request(run_id="application-draft-rollback"), build_facts(), submit=False, ) db.commit() assert list(db.scalars(select(ExpenseClaim)).all()) == [] assert list(db.scalars(select(ExpenseCaseLink)).all()) == [] assert list(db.scalars(select(BusinessEvent)).all()) == [] def test_ai_application_draft_update_rolls_back_when_event_write_fails( monkeypatch: pytest.MonkeyPatch, ) -> None: session_factory = build_session_factory() with session_factory() as db: owner = Employee( id="owner-1", employee_no="E001", name="张三", email="owner@example.com", ) claim = ExpenseClaim( id="application-1", claim_no="AP-DRAFT-ROLLBACK-001", employee_id=owner.id, employee_name=owner.name, department_name="市场部", expense_type="travel_application", reason="原申请事由", location="北京", amount=Decimal("500.00"), currency="CNY", invoice_count=0, occurred_at=datetime(2026, 7, 10, tzinfo=UTC), submitted_at=None, status="returned", approval_stage="退回补充", risk_flags_json=[], ) db.add_all([owner, claim]) db.commit() def fail_event(*_args, **_kwargs): raise RuntimeError("模拟草稿事件写入失败") monkeypatch.setattr(ExpenseCaseService, "record_claim_event", fail_event) with pytest.raises(RuntimeError, match="模拟草稿事件写入失败"): UserAgentService(db)._update_expense_application_record( build_request(run_id="application-draft-update-rollback"), build_facts(reason="不应落库的新事由"), claim, submit=False, ) db.commit() persisted = db.get(ExpenseClaim, claim.id) assert persisted is not None assert persisted.reason == "原申请事由" assert persisted.location == "北京" assert persisted.amount == Decimal("500.00") assert persisted.status == "returned" assert persisted.approval_stage == "退回补充" assert list(db.scalars(select(ExpenseCaseLink)).all()) == [] assert list(db.scalars(select(BusinessEvent)).all()) == [] def test_ai_application_draft_update_deduplicates_identical_snapshot() -> None: session_factory = build_session_factory() with session_factory() as db: owner = Employee( id="owner-1", employee_no="E001", name="张三", email="owner@example.com", ) claim = ExpenseClaim( id="application-1", claim_no="AP-DRAFT-IDEMPOTENT-001", employee_id=owner.id, employee_name=owner.name, department_name="市场部", expense_type="travel_application", reason="原申请事由", location="北京", amount=Decimal("500.00"), currency="CNY", invoice_count=0, occurred_at=datetime(2026, 7, 10, tzinfo=UTC), submitted_at=None, status="returned", approval_stage="退回补充", risk_flags_json=[], ) db.add_all([owner, claim]) db.commit() request = build_request(run_id="application-draft-idempotent") facts = build_facts(reason="同一份草稿") service = UserAgentService(db) learning_current_user = service._build_application_current_user(request) service._update_expense_application_record( request, facts, claim, submit=False, learning_current_user=learning_current_user, ) service._update_expense_application_record( request, facts, claim, submit=False, learning_current_user=learning_current_user, ) events = list( db.scalars( select(BusinessEvent).where(BusinessEvent.aggregate_id == claim.id) ).all() ) assert len(events) == 1 assert events[0].event_type == "claim_draft_updated" assert events[0].payload_json["previous_status"] == "returned" assert len(list(db.scalars(select(AIDecision)).all())) == 1 assert len(list(db.scalars(select(AIDecisionFeedback)).all())) == 1 assert len(list(db.scalars(select(WorkflowOutcome)).all())) == 1 def test_ai_application_draft_update_keeps_distinct_snapshots_in_same_run() -> None: session_factory = build_session_factory() with session_factory() as db: owner = Employee( id="owner-1", employee_no="E001", name="张三", email="owner@example.com", ) claim = ExpenseClaim( id="application-1", claim_no="AP-DRAFT-VERSIONS-001", employee_id=owner.id, employee_name=owner.name, department_name="市场部", expense_type="travel_application", reason="原申请事由", location="北京", amount=Decimal("500.00"), currency="CNY", invoice_count=0, occurred_at=datetime(2026, 7, 10, tzinfo=UTC), submitted_at=None, status="returned", approval_stage="退回补充", risk_flags_json=[], ) db.add_all([owner, claim]) db.commit() request = build_request(run_id="application-draft-multi-version") service = UserAgentService(db) learning_current_user = service._build_application_current_user(request) service._update_expense_application_record( request, build_facts(reason="第一版草稿"), claim, submit=False, learning_current_user=learning_current_user, ) service._update_expense_application_record( request, build_facts(reason="第二版草稿"), claim, submit=False, learning_current_user=learning_current_user, ) events = list( db.scalars( select(BusinessEvent).where(BusinessEvent.aggregate_id == claim.id) ).all() ) assert len(events) == 2 assert {event.event_type for event in events} == {"claim_draft_updated"} assert len({event.idempotency_key for event in events}) == 2 assert len(list(db.scalars(select(AIDecision)).all())) == 2 def test_ai_application_draft_records_explicit_field_correction() -> None: session_factory = build_session_factory() with session_factory() as db: owner = Employee( id="owner-1", employee_no="E001", name="张三", email="owner@example.com", ) claim = ExpenseClaim( id="application-1", claim_no="AP-DRAFT-FEEDBACK-001", employee_id=owner.id, employee_name=owner.name, department_name="市场部", expense_type="travel_application", reason="原申请事由", location="北京", amount=Decimal("500.00"), currency="CNY", invoice_count=0, occurred_at=datetime(2026, 7, 10, tzinfo=UTC), submitted_at=None, status="returned", approval_stage="退回补充", risk_flags_json=[], ) db.add_all([owner, claim]) db.commit() service = UserAgentService(db) request = build_request( run_id="application-draft-feedback", changed_fields=[ { "fieldKey": "reason", "suggestedValue": "客户拜访", "finalValue": "客户端伪造值不会直接入账", } ], ) service._update_expense_application_record( request, build_facts(reason="客户现场实施"), claim, submit=False, learning_current_user=service._build_application_current_user(request), ) decision = db.scalar(select(AIDecision)) feedback = db.scalar(select(AIDecisionFeedback)) assert decision is not None assert decision.training_eligible is False assert decision.suggestion_json["field_keys"] assert decision.suggestion_json["value_fingerprint"].startswith("sha256:") assert feedback is not None assert feedback.feedback_type == "edited" assert feedback.verification_status == "client_observed" assert feedback.training_eligible is False assert feedback.changed_fields_json[0]["field_key"] == "reason" assert feedback.changed_fields_json[0]["suggested_value_fingerprint"].startswith( "sha256:" ) assert "客户拜访" not in str(feedback.changed_fields_json) assert "客户现场实施" not in str(feedback.final_value_json)