from __future__ import annotations import uuid from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest from sqlalchemy import select from app.api.deps import CurrentUserContext from app.models.ai_memory import MemoryEntry from app.models.organization import OrganizationUnit from app.schemas.expense_application_memory import ( ExpenseApplicationOrganizationMemoryCreate, ExpenseApplicationOrganizationMemoryRevoke, ExpenseApplicationOrganizationMemoryUpdate, ) from app.services.expense_application_memory import ExpenseApplicationMemoryService from app.services.expense_application_memory_admin import ( ExpenseApplicationOrganizationMemoryService, OrganizationMemoryConflictError, OrganizationMemoryNotFoundError, ) from app.services.organization_memory_locks import ( organization_memory_advisory_lock_id, organization_memory_operation_locks, ) from app.test_helpers.db import build_in_memory_session_factory def _user( *, tenant_id: str = "tenant-hierarchy", department_id: str = "department-delivery", is_admin: bool = False, ) -> CurrentUserContext: return CurrentUserContext( username="memory-admin@example.com" if is_admin else "memory-user@example.com", name="记忆管理员" if is_admin else "记忆员工", role_codes=["admin"] if is_admin else ["user"], is_admin=is_admin, tenant_id=tenant_id, employee_id="memory-admin" if is_admin else "memory-user", employee_no="E-MEMORY", department_id=department_id, department_name="交付部", ) def _active_organization_memory( *, tenant_id: str, scope_type: str, scope_id: str, value: str, generation: int, activated_at: datetime | None = None, expires_at: datetime | None = None, ) -> MemoryEntry: now = datetime.now(UTC) activated_at = activated_at or now expires_at = expires_at or now + timedelta(days=180) return MemoryEntry( id=str(uuid.uuid4()), tenant_id=tenant_id, scope_type=scope_type, scope_id=scope_id, origin_type="admin_managed", managed_by="memory-admin", managed_at=now, management_reason="测试组织记忆", policy_version="expense_application_transport_org_memory.v1", scene="travel_application", field_key="transport_mode", generation=generation, value_json={"value": value}, value_fingerprint=f"fingerprint-{scope_type}-{generation}-{value}", status="active", evidence_count=0, approved_evidence_count=0, confidence=Decimal("1.0000"), last_evidence_at=now, candidate_expires_at=expires_at, activated_at=activated_at, active_expires_at=expires_at, created_at=activated_at, updated_at=now, ) def test_enterprise_precedes_department_and_explicit_input_precedes_all() -> None: with build_in_memory_session_factory()() as db: current_user = _user() db.add_all( [ _active_organization_memory( tenant_id=current_user.tenant_id, scope_type="enterprise", scope_id=current_user.tenant_id, value="飞机", generation=1, ), _active_organization_memory( tenant_id=current_user.tenant_id, scope_type="department", scope_id=current_user.department_id, value="火车", generation=1, ), ] ) db.flush() facts: dict[str, object] = {} applications = ExpenseApplicationMemoryService(db).apply_active_transport_memory( facts, current_user ) assert facts == {"transport_mode": "飞机"} assert applications[0].scope_type == "enterprise" assert applications[0].priority == 300 assert applications[0].source == "enterprise_policy_memory" assert applications[0].can_revoke is False assert applications[0].conflicts[0].reason == "lower_priority_overridden" assert applications[0].conflicts[0].scope_type == "department" explicit_facts = {"transport_mode": "轮船"} assert ( ExpenseApplicationMemoryService(db).apply_active_transport_memory( explicit_facts, current_user ) == [] ) assert explicit_facts == {"transport_mode": "轮船"} def test_same_level_conflict_fails_closed_and_is_explainable() -> None: with build_in_memory_session_factory()() as db: current_user = _user() db.add_all( [ _active_organization_memory( tenant_id=current_user.tenant_id, scope_type="enterprise", scope_id=current_user.tenant_id, value="飞机", generation=1, ), _active_organization_memory( tenant_id=current_user.tenant_id, scope_type="enterprise", scope_id=current_user.tenant_id, value="火车", generation=2, ), ] ) db.flush() facts: dict[str, object] = {} applications = ExpenseApplicationMemoryService(db).apply_active_transport_memory( facts, current_user ) assert facts == {} assert len(applications) == 1 assert applications[0].status == "conflict" assert applications[0].memory_id == "" assert applications[0].scope_id == "" assert applications[0].value == "" assert len(applications[0].conflicts) == 2 assert { conflict.reason for conflict in applications[0].conflicts } == {"same_priority_conflict"} for conflict in applications[0].model_dump()["conflicts"]: assert "memory_id" not in conflict assert "scope_id" not in conflict assert "value" not in conflict def test_expiry_and_time_decay_never_apply_cross_tenant_or_expire_early() -> None: with build_in_memory_session_factory()() as db: current_user = _user() now = datetime.now(UTC) expired = _active_organization_memory( tenant_id=current_user.tenant_id, scope_type="department", scope_id=current_user.department_id, value="火车", generation=1, activated_at=now - timedelta(days=181), expires_at=now - timedelta(days=1), ) other_tenant = _active_organization_memory( tenant_id="tenant-other", scope_type="enterprise", scope_id="tenant-other", value="轮船", generation=1, ) decayed = _active_organization_memory( tenant_id=current_user.tenant_id, scope_type="enterprise", scope_id=current_user.tenant_id, value="飞机", generation=1, activated_at=now - timedelta(days=100), expires_at=now + timedelta(days=100), ) db.add_all([expired, other_tenant, decayed]) db.flush() facts: dict[str, object] = {} applications = ExpenseApplicationMemoryService(db).apply_active_transport_memory( facts, current_user ) db.refresh(expired) assert expired.status == "expired" assert facts == {"transport_mode": "飞机"} assert 0.49 <= applications[0].effective_confidence <= 0.51 assert applications[0].memory_id == decayed.id def test_organization_memory_lifecycle_is_audited_versioned_and_tenant_safe() -> None: with build_in_memory_session_factory()() as db: admin = _user(is_admin=True) service = ExpenseApplicationOrganizationMemoryService(db) created = service.create_organization_memory( ExpenseApplicationOrganizationMemoryCreate( scope_type="enterprise", value="火车", expires_in_days=90, reason="统一优先使用高铁", request_id="create-enterprise-memory-1", ), admin, ) assert created.scope_id == admin.tenant_id assert created.generation == 1 assert created.origin_type == "admin_managed" assert created.managed_by == admin.employee_id assert created.management_reason == "统一优先使用高铁" replayed_create = service.create_organization_memory( ExpenseApplicationOrganizationMemoryCreate( scope_type="enterprise", value="火车", expires_in_days=90, reason="统一优先使用高铁", request_id="create-enterprise-memory-1", ), admin, ) assert replayed_create.id == created.id with pytest.raises(OrganizationMemoryConflictError, match="已有生效记忆"): service.create_organization_memory( ExpenseApplicationOrganizationMemoryCreate( scope_type="enterprise", value="火车", expires_in_days=90, reason="重复创建企业基线", request_id="create-enterprise-memory-2", ), admin, ) updated = service.update_organization_memory( created.id, ExpenseApplicationOrganizationMemoryUpdate( value="飞机", expires_in_days=120, expected_generation=1, reason="适配跨区域差旅", request_id="update-enterprise-memory-1", ), admin, ) assert updated.id != created.id assert updated.generation == 2 old_entry = db.get(MemoryEntry, created.id) assert old_entry is not None assert old_entry.status == "suppressed" assert old_entry.superseded_by_id == updated.id replayed_update = service.update_organization_memory( created.id, ExpenseApplicationOrganizationMemoryUpdate( value="飞机", expires_in_days=120, expected_generation=1, reason="适配跨区域差旅", request_id="update-enterprise-memory-1", ), admin, ) assert replayed_update.id == updated.id with pytest.raises(OrganizationMemoryConflictError, match="刷新后重试"): service.update_organization_memory( created.id, ExpenseApplicationOrganizationMemoryUpdate( value="轮船", expected_generation=1, reason="使用旧记录地址更新", request_id="update-enterprise-memory-old-id", ), admin, ) with pytest.raises(OrganizationMemoryConflictError, match="刷新后重试"): service.update_organization_memory( updated.id, ExpenseApplicationOrganizationMemoryUpdate( expected_generation=1, reason="使用陈旧版本更新", request_id="update-enterprise-memory-stale", ), admin, ) with pytest.raises(OrganizationMemoryNotFoundError): service.revoke_organization_memory( updated.id, ExpenseApplicationOrganizationMemoryRevoke( expected_generation=2, reason="越权租户撤销", request_id="revoke-other-tenant-memory", ), _user(tenant_id="tenant-other", is_admin=True), ) revoked = service.revoke_organization_memory( updated.id, ExpenseApplicationOrganizationMemoryRevoke( expected_generation=2, reason="企业政策已调整", request_id="revoke-enterprise-memory-1", ), admin, ) assert revoked.memory_id == updated.id revoked_entry = db.get(MemoryEntry, updated.id) assert revoked_entry is not None assert revoked_entry.status == "revoked" assert revoked_entry.value_json == {"value": "飞机"} assert revoked_entry.revoked_reason == "企业政策已调整" replayed_revoke = service.revoke_organization_memory( updated.id, ExpenseApplicationOrganizationMemoryRevoke( expected_generation=2, reason="企业政策已调整", request_id="revoke-enterprise-memory-1", ), admin, ) assert replayed_revoke == revoked with pytest.raises(OrganizationMemoryConflictError, match="请求内容不一致"): service.revoke_organization_memory( updated.id, ExpenseApplicationOrganizationMemoryRevoke( expected_generation=2, reason="同一键却改变请求体", request_id="revoke-enterprise-memory-1", ), admin, ) assert ExpenseApplicationMemoryService(db).revoke_current_user_memory( updated.id, admin ) is None def test_department_memory_requires_existing_department_and_admin() -> None: with build_in_memory_session_factory()() as db: department = OrganizationUnit( id="department-delivery", unit_code="DELIVERY", name="交付部", unit_type="department", ) division = OrganizationUnit( id="division-east", unit_code="EAST", name="华东大区", unit_type="division", ) db.add_all([department, division]) db.flush() service = ExpenseApplicationOrganizationMemoryService(db) payload = ExpenseApplicationOrganizationMemoryCreate( scope_type="department", scope_id=department.id, value="火车", reason="部门短途差旅基线", request_id="create-department-memory-1", ) with pytest.raises(PermissionError, match="平台管理员"): service.create_organization_memory(payload, _user()) with pytest.raises(ValueError, match="department 类型"): service.create_organization_memory( payload.model_copy(update={"scope_id": division.id}), _user(is_admin=True), ) created = service.create_organization_memory(payload, _user(is_admin=True)) assert created.scope_id == department.id assert created.scope_label == "部门规则(交付部)" memories = service.list_organization_memories(_user(is_admin=True)) assert [item.id for item in memories.items] == [created.id] assert db.scalar( select(MemoryEntry).where(MemoryEntry.id == created.id) ) is not None def test_expired_organization_memory_revoke_returns_conflict_and_marks_expired() -> None: with build_in_memory_session_factory()() as db: admin = _user(is_admin=True) now = datetime.now(UTC) entry = _active_organization_memory( tenant_id=admin.tenant_id, scope_type="enterprise", scope_id=admin.tenant_id, value="火车", generation=1, activated_at=now - timedelta(days=100), expires_at=now - timedelta(days=1), ) db.add(entry) db.commit() with pytest.raises(OrganizationMemoryConflictError, match="已过期"): ExpenseApplicationOrganizationMemoryService(db).revoke_organization_memory( entry.id, ExpenseApplicationOrganizationMemoryRevoke( expected_generation=1, reason="撤销已过期规则", request_id="revoke-expired-memory-1", ), admin, ) db.refresh(entry) assert entry.status == "expired" assert entry.expired_at is not None def test_postgres_organization_mutations_use_signed_transaction_advisory_locks() -> None: executed: list[object] = [] class _Dialect: name = "postgresql" class _Bind: dialect = _Dialect() class _RecordingSession: @staticmethod def get_bind() -> _Bind: return _Bind() @staticmethod def execute(statement: object) -> None: executed.append(statement) session = _RecordingSession() with organization_memory_operation_locks( # type: ignore[arg-type] session, "scope-lock", "request-lock", ): assert len(executed) == 2 assert all("pg_advisory_xact_lock" in str(statement) for statement in executed) lock_id = organization_memory_advisory_lock_id("scope-lock") assert -(2**63) <= lock_id < 2**63 assert lock_id == organization_memory_advisory_lock_id("scope-lock") def test_new_generation_rejects_unexpected_active_set_before_suppression() -> None: with build_in_memory_session_factory()() as db: admin = _user(is_admin=True) active = _active_organization_memory( tenant_id=admin.tenant_id, scope_type="enterprise", scope_id=admin.tenant_id, value="火车", generation=1, ) db.add(active) db.commit() now = datetime.now(UTC) with pytest.raises(OrganizationMemoryConflictError, match="生效版本已变化"): ExpenseApplicationOrganizationMemoryService(db)._new_active_generation( entry_id=str(uuid.uuid4()), tenant_id=admin.tenant_id, scope_type="enterprise", scope_id=admin.tenant_id, value="飞机", expires_at=now + timedelta(days=180), reason="模拟并发创建", request_id="unexpected-active-generation", request_fingerprint="hmac-sha256:" + "0" * 64, expected_active_ids=set(), current_user=admin, now=now, ) db.refresh(active) assert active.status == "active" assert active.superseded_by_id is None