from __future__ import annotations import json import uuid from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import Any import pytest from commercial_runtime_testkit import seed_meter as _seed_meter from commercial_runtime_testkit import seed_run as _seed_run from commercial_runtime_testkit import seed_tool_call as _seed_tool_call from sqlalchemy import create_engine, select from sqlalchemy.orm import Session, sessionmaker from sqlalchemy.pool import StaticPool import app.models # noqa: F401 - 注册完整 metadata from app.db.base_class import Base from app.models.agent_run import AgentToolCall from app.models.commercial import CommercialCostEvent, UsageMeterEvent from app.models.commercial_runtime import CommercialRuntimeReservation from app.services.agent_runs import AgentRunService from app.services.commercial_access_policy import CommercialConflictError from app.services.commercial_entitlements import CommercialEntitlementService from app.services.commercial_metering import CommercialMeteringService from app.services.commercial_runtime_bridge import CommercialRuntimeBridge from app.services.commercial_runtime_metering import CommercialRuntimeMeteringService from app.services.commercial_runtime_reservations import ( CommercialRuntimeReservationService, ) from app.services.orchestrator_execution import OrchestratorExecutionEngine @pytest.fixture() def db() -> Session: engine = create_engine( "sqlite+pysqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) Base.metadata.create_all(engine) factory = sessionmaker(bind=engine, expire_on_commit=False) with factory() as session: yield session Base.metadata.drop_all(engine) engine.dispose() def test_real_tokens_create_redacted_usage_and_linked_cost_idempotently( db: Session, ) -> None: now = datetime.now(UTC) _, entitlement = _seed_meter( db, "tenant-a", now, basis="total_tokens", internal_cost={ "enabled": True, "cost_category": "ai_inference", "unit": "token", "unit_cost": "0.002", "original_currency": "CNY", "reporting_currency": "CNY", "fx_rate": "1", "provider": "provider-a", "model_name": "model-a", }, ) _, tool_call = _seed_tool_call( db, now, route_json={"tenant_id": "tenant-a"}, request_json={ "usage": {"input_tokens": 12}, "prompt": "REQUEST-TOP-SECRET", "tenant_id": "tenant-spoofed", "user_id": "raw-user-in-request", }, response_json={ "usage": {"output_tokens": 8}, "answer": "RESPONSE-TOP-SECRET", }, user_id="raw-run-user", ) result = CommercialRuntimeMeteringService(db).sync_tool_call(tool_call.id) replay = CommercialRuntimeMeteringService(db).sync_tool_call(tool_call.id) usage = db.scalars(select(UsageMeterEvent)).one() cost = db.scalars(select(CommercialCostEvent)).one() assert result.status == "created" assert result.quantity == Decimal("20") assert replay.status == "replayed" assert replay.usage_event_id == usage.id assert replay.cost_event_id == cost.id assert usage.tenant_id == "tenant-a" assert usage.entitlement_id == entitlement.id assert Decimal(usage.quantity) == Decimal("20") assert cost.usage_event_id == usage.id assert Decimal(cost.quantity) == Decimal("20") assert Decimal(cost.cost_amount) == Decimal("0.0400") assert db.query(UsageMeterEvent).count() == 1 assert db.query(CommercialCostEvent).count() == 1 serialized_metadata = json.dumps( {"usage": usage.metadata_json, "cost": cost.metadata_json}, ensure_ascii=False, ) for forbidden in ( "REQUEST-TOP-SECRET", "RESPONSE-TOP-SECRET", "raw-run-user", "raw-user-in-request", "tenant-spoofed", "request_json", "response_json", ): assert forbidden not in serialized_metadata @pytest.mark.parametrize( ("basis", "request_json", "response_json", "duration_ms", "expected"), [ ("input_tokens", {"token_usage": {"prompt_tokens": 11}}, {}, 0, Decimal("11")), ("output_tokens", {}, {"metrics": {"completion_tokens": 7}}, 0, Decimal("7")), ("duration_ms", {}, {}, 321, Decimal("321")), ], ) def test_supported_measured_quantity_bases( db: Session, basis: str, request_json: dict[str, Any], response_json: dict[str, Any], duration_ms: int, expected: Decimal, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis=basis) _, tool_call = _seed_tool_call( db, now, route_json={"tenant_id": "tenant-a"}, request_json=request_json, response_json=response_json, duration_ms=duration_ms, ) result = CommercialRuntimeMeteringService(db).sync_tool_call(tool_call.id) assert result.status == "created" assert result.quantity_basis == basis assert result.quantity == expected assert Decimal(db.scalars(select(UsageMeterEvent)).one().quantity) == expected def test_call_basis_uses_route_tenant_and_isolated_entitlement(db: Session) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="call") _, tenant_b_entitlement = _seed_meter( db, "tenant-b", now, basis="call", subscription_status="trialing", ) _, tool_call = _seed_tool_call( db, now, route_json={"tenant_id": "tenant-b"}, request_json={"tenant_id": "tenant-a", "user_id": "tenant-a-user"}, response_json={"tenant_id": "tenant-a"}, ) result = CommercialRuntimeMeteringService(db).sync_tool_call(tool_call.id) usage = db.scalars(select(UsageMeterEvent)).one() assert result.status == "created" assert result.quantity == Decimal("1") assert usage.tenant_id == "tenant-b" assert usage.entitlement_id == tenant_b_entitlement.id assert ( db.scalar(select(UsageMeterEvent.id).where(UsageMeterEvent.tenant_id == "tenant-a")) is None ) def test_missing_route_tenant_and_missing_tokens_are_collecting_not_estimated( db: Session, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="total_tokens") _, unowned_call = _seed_tool_call( db, now, route_json={}, request_json={"tenant_id": "tenant-a", "usage": {"input_tokens": 9}}, response_json={"usage": {"output_tokens": 3}}, ) _, unmeasured_call = _seed_tool_call( db, now + timedelta(microseconds=1), route_json={"tenant_id": "tenant-a"}, request_json={"prompt": "x" * 20_000}, response_json={"answer": "y" * 20_000}, ) service = CommercialRuntimeMeteringService(db) unowned = service.sync_tool_call(unowned_call.id) unmeasured = service.sync_tool_call(unmeasured_call.id) assert unowned.status == "skipped" assert unowned.reason_code == "missing_tenant" assert unowned.business_call_occurred is True assert unowned.requires_reconciliation is True assert unmeasured.status == "skipped" assert unmeasured.reason_code == "collecting_missing_tokens" assert unmeasured.collecting is True assert "正文长度" in unmeasured.reason assert db.query(UsageMeterEvent).count() == 0 def test_only_active_exact_runtime_meter_is_eligible(db: Session) -> None: now = datetime.now(UTC) subscription, entitlement = _seed_meter(db, "tenant-a", now, basis="call") service = CommercialRuntimeMeteringService(db) _, wrong_name = _seed_tool_call( db, now, route_json={"tenant_id": "tenant-a"}, tool_name="llm.responses", ) assert service.sync_tool_call(wrong_name.id).reason_code == "no_matching_runtime_meter" entitlement.status = "suspended" db.flush() _, suspended_entitlement = _seed_tool_call( db, now + timedelta(microseconds=1), route_json={"tenant_id": "tenant-a"}, ) result = service.sync_tool_call(suspended_entitlement.id) assert result.reason_code == "no_matching_runtime_meter" entitlement.status = "active" subscription.status = "past_due" db.flush() _, inactive_subscription = _seed_tool_call( db, now + timedelta(microseconds=2), route_json={"tenant_id": "tenant-a"}, ) result = service.sync_tool_call(inactive_subscription.id) assert result.reason_code == "no_active_subscription" assert db.query(UsageMeterEvent).count() == 0 def test_preflight_blocks_quota_and_post_call_error_requires_reconciliation( db: Session, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="call", hard_limit=Decimal("2")) run, first = _seed_tool_call(db, now, route_json={"tenant_id": "tenant-a"}) _, second = _seed_tool_call( db, now + timedelta(microseconds=1), route_json={"tenant_id": "tenant-a"}, run=run, ) _, third = _seed_tool_call( db, now + timedelta(microseconds=2), route_json={"tenant_id": "tenant-a"}, run=run, ) service = CommercialRuntimeMeteringService(db) allowed = service.preflight_run_tool( run.run_id, tool_type="llm", tool_name="chat.completions", ) assert allowed.allowed is True assert ( service.assert_run_tool_allowed( run.run_id, tool_type="llm", tool_name="chat.completions", ).allowed is True ) oversized = service.preflight_run_tool( run.run_id, tool_type="llm", tool_name="chat.completions", requested_quantity=Decimal("3"), ) assert oversized.allowed is False with pytest.raises(CommercialConflictError): service.assert_run_tool_allowed( run.run_id, tool_type="llm", tool_name="chat.completions", requested_quantity=Decimal("3"), ) assert service.sync_tool_call(first.id).status == "created" assert service.sync_tool_call(second.id).status == "created" exhausted = service.preflight_run_tool( run.run_id, tool_type="llm", tool_name="chat.completions", ) assert exhausted.allowed is False failed_fact = service.sync_tool_call(third.id) assert failed_fact.status == "error" assert failed_fact.reason_code == "metering_failed_after_business_call" assert failed_fact.business_call_occurred is True assert failed_fact.requires_reconciliation is True assert db.query(UsageMeterEvent).count() == 2 def test_zero_internal_cost_never_writes_fake_usage_or_cost(db: Session) -> None: now = datetime.now(UTC) _seed_meter( db, "tenant-a", now, basis="call", internal_cost={ "enabled": True, "cost_category": "ai_inference", "unit_cost": "0", "fx_rate": "1", "original_currency": "CNY", "reporting_currency": "CNY", }, ) _, tool_call = _seed_tool_call(db, now, route_json={"tenant_id": "tenant-a"}) result = CommercialRuntimeMeteringService(db).sync_tool_call(tool_call.id) assert result.status == "error" assert result.business_call_occurred is True assert result.requires_reconciliation is True assert db.query(UsageMeterEvent).count() == 0 assert db.query(CommercialCostEvent).count() == 0 def test_batch_is_bounded_cursor_based_and_keeps_error_semantics(db: Session) -> None: now = datetime.now(UTC) - timedelta(seconds=1) _seed_meter(db, "tenant-a", now, basis="input_tokens") _seed_tool_call( db, now, route_json={"tenant_id": "tenant-a"}, request_json={"usage": {"input_tokens": 3}}, ) _seed_tool_call( db, now + timedelta(microseconds=1), route_json={"tenant_id": "tenant-a"}, request_json={"prompt": "not-a-token-count"}, ) _seed_tool_call( db, now + timedelta(microseconds=2), route_json={"tenant_id": "tenant-a"}, request_json={"usage": {"input_tokens": "3"}}, ) service = CommercialRuntimeMeteringService(db) first_page = service.sync_batch(limit=2) assert first_page.created == 1 assert first_page.skipped == 1 assert first_page.errors == 0 assert first_page.has_more is True assert first_page.next_cursor is not None second_page = service.sync_batch(limit=2, cursor=first_page.next_cursor) assert second_page.created == 0 assert second_page.skipped == 0 assert second_page.errors == 1 assert second_page.items[0].business_call_occurred is True assert second_page.items[0].requires_reconciliation is True assert second_page.has_more is False assert db.query(UsageMeterEvent).count() == 1 with pytest.raises(ValueError, match="200"): service.sync_batch(limit=201) def test_non_successful_tool_calls_never_enter_commercial_ledger(db: Session) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="call") calls = [ _seed_tool_call( db, now + timedelta(microseconds=index), route_json={"tenant_id": "tenant-a"}, status=status, )[1] for index, status in enumerate(("running", "blocked", "failed")) ] service = CommercialRuntimeMeteringService(db) batch = service.sync_batch(limit=10) assert [item.reason_code for item in batch.items] == [ "tool_call_not_terminal", "tool_call_not_billable", "tool_call_not_billable", ] assert batch.items[0].collecting is True assert batch.items[0].requires_reconciliation is True assert batch.items[1].business_call_occurred is False assert batch.items[2].business_call_occurred is True assert db.query(UsageMeterEvent).count() == 0 calls[0].status = "succeeded" db.flush() completed = service.sync_tool_call(calls[0].id) assert completed.status == "created" assert db.query(UsageMeterEvent).count() == 1 def test_runtime_bridge_bypasses_unconfigured_tool_without_fake_fact(db: Session) -> None: now = datetime.now(UTC) run = _seed_run(db, now, route_json={"tenant_id": "tenant-a"}) bridge = CommercialRuntimeBridge(db) gate = bridge.preflight(run.run_id, tool_type="llm", tool_name="chat.completions") _, call = _seed_tool_call( db, now, route_json={"tenant_id": "tenant-a"}, run=run, ) result = bridge.sync_tool_call(call.id) assert gate.enforced is False assert gate.allowed is True assert gate.reason_code == "runtime_meter_not_configured" assert result.status == "skipped" assert result.requires_reconciliation is False assert db.query(UsageMeterEvent).count() == 0 def test_agent_run_service_without_preflight_creates_durable_backlog( db: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="call") run = _seed_run(db, now, route_json={"tenant_id": "tenant-a"}) run_service = AgentRunService(db) monkeypatch.setattr(run_service, "_ensure_ready", lambda: None) call = run_service.record_tool_call( run_id=run.run_id, tool_type="llm", tool_name="chat.completions", request_json={"prompt": "not persisted to commercial metadata"}, response_json={"answer": "ok"}, status="succeeded", duration_ms=12, ) replay = CommercialRuntimeBridge(db).sync_tool_call(call.id) backlog = db.scalars(select(CommercialRuntimeReservation)).one() assert replay.status == "error" assert replay.reason_code == "reservation_not_settleable" assert backlog.status == "reconciliation_required" assert backlog.tool_call_id == call.id assert db.query(UsageMeterEvent).count() == 0 def test_running_direct_tool_call_requires_pre_execution_reservation_on_update( db: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="duration_ms") run = _seed_run(db, now, route_json={"tenant_id": "tenant-a"}) run_service = AgentRunService(db) monkeypatch.setattr(run_service, "_ensure_ready", lambda: None) call = run_service.record_tool_call( run_id=run.run_id, tool_type="llm", tool_name="chat.completions", status="running", ) assert db.query(UsageMeterEvent).count() == 0 run_service.update_tool_call( call.id, response_json={"answer": "done"}, status="succeeded", duration_ms=321, ) backlog = db.scalars(select(CommercialRuntimeReservation)).one() assert backlog.status == "reconciliation_required" assert Decimal(backlog.actual_quantity or 0) == Decimal("321") assert db.query(UsageMeterEvent).count() == 0 def test_metering_failure_keeps_tool_call_and_supports_idempotent_reconciliation( db: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="call") run = _seed_run(db, now, route_json={"tenant_id": "tenant-a"}) run_service = AgentRunService(db) monkeypatch.setattr(run_service, "_ensure_ready", lambda: None) call_id = str(uuid.uuid4()) permit = CommercialRuntimeBridge(db).reserve_tool( run.run_id, tool_call_id=call_id, tool_type="llm", tool_name="chat.completions", ) assert permit.reservation_id is not None original_sync = CommercialRuntimeMeteringService.sync_reserved_tool_call def fail_metering(self, tool_call_id: str, reservation): # noqa: ANN001 del self, tool_call_id, reservation raise RuntimeError("simulated metering outage") monkeypatch.setattr( CommercialRuntimeMeteringService, "sync_reserved_tool_call", fail_metering, ) call = run_service.record_tool_call( run_id=run.run_id, tool_call_id=call_id, tool_type="llm", tool_name="chat.completions", status="succeeded", ) assert db.get(AgentToolCall, call.id) is not None assert db.query(UsageMeterEvent).count() == 0 monkeypatch.setattr( CommercialRuntimeMeteringService, "sync_reserved_tool_call", original_sync, ) reconciled = CommercialRuntimeBridge(db).sync_tool_call(call.id) replay = CommercialRuntimeBridge(db).sync_tool_call(call.id) assert reconciled.status == "created" assert replay.status == "replayed" assert db.query(UsageMeterEvent).count() == 1 def test_cost_failure_commits_usage_then_reconciliation_appends_missing_cost( db: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: now = datetime.now(UTC) _seed_meter( db, "tenant-a", now, basis="call", internal_cost={ "enabled": True, "cost_category": "ai_inference", "unit": "call", "unit_cost": "0.5", "original_currency": "CNY", "reporting_currency": "CNY", "fx_rate": "1", }, ) run = _seed_run(db, now, route_json={"tenant_id": "tenant-a"}) run_service = AgentRunService(db) monkeypatch.setattr(run_service, "_ensure_ready", lambda: None) call_id = str(uuid.uuid4()) permit = CommercialRuntimeBridge(db).reserve_tool( run.run_id, tool_call_id=call_id, tool_type="llm", tool_name="chat.completions", ) assert permit.reservation_id is not None original_record_cost = CommercialMeteringService.record_cost def fail_cost(self, tenant_id, payload): # noqa: ANN001 del self, tenant_id, payload raise RuntimeError("simulated cost ledger outage") monkeypatch.setattr(CommercialMeteringService, "record_cost", fail_cost) call = run_service.record_tool_call( run_id=run.run_id, tool_call_id=call_id, tool_type="llm", tool_name="chat.completions", status="succeeded", ) assert db.query(UsageMeterEvent).count() == 1 assert db.query(CommercialCostEvent).count() == 0 reservation = db.get(CommercialRuntimeReservation, permit.reservation_id) assert reservation is not None assert reservation.status == "committed_reconciliation_required" assert reservation.resolution_code == "cost_metering_failed" candidates = CommercialRuntimeReservationService(db).reconciliation_candidates() quota = CommercialEntitlementService(db).get_account_for_tenant("tenant-a").quotas[0] assert [row.id for row in candidates] == [reservation.id] assert quota.used_quantity == Decimal("1") assert quota.reserved_quantity == Decimal("0") monkeypatch.setattr(CommercialMeteringService, "record_cost", original_record_cost) reconciled = CommercialRuntimeBridge(db).sync_tool_call(call.id) db.refresh(reservation) assert reconciled.status == "created" assert reconciled.usage_created is False assert reconciled.cost_created is True assert db.query(UsageMeterEvent).count() == 1 assert db.query(CommercialCostEvent).count() == 1 assert reservation.status == "committed" assert reservation.resolution_code is None def test_orchestrator_preflight_blocks_executor_after_quota_exhaustion( db: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: now = datetime.now(UTC) _seed_meter(db, "tenant-a", now, basis="call", hard_limit=Decimal("1")) run = _seed_run(db, now, route_json={"tenant_id": "tenant-a"}) run_service = AgentRunService(db) monkeypatch.setattr(run_service, "_ensure_ready", lambda: None) first_call_id = str(uuid.uuid4()) permit = CommercialRuntimeBridge(db).reserve_tool( run.run_id, tool_call_id=first_call_id, tool_type="llm", tool_name="chat.completions", ) assert permit.reservation_id is not None run_service.record_tool_call( run_id=run.run_id, tool_call_id=first_call_id, tool_type="llm", tool_name="chat.completions", status="succeeded", ) engine = OrchestratorExecutionEngine( db=db, run_service=run_service, expense_claim_service=None, knowledge_service=None, user_agent_service=None, database_query_builder=None, ) executed = False def executor() -> dict[str, Any]: nonlocal executed executed = True return {"answer": "should not run"} response, degraded = engine._invoke_tool( run_id=run.run_id, tool_type="llm", tool_name="chat.completions", request_json={"prompt": "quota check"}, context_json={}, executor=executor, fallback_factory=lambda error: {"error": str(error)}, ) assert executed is False assert degraded is True assert "配额" in response["error"] assert db.query(UsageMeterEvent).count() == 1 statuses = sorted( db.scalars( select(AgentToolCall.status) .where(AgentToolCall.run_id == run.run_id) ).all() ) assert statuses == ["blocked", "succeeded"]