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

691 lines
23 KiB
Python
Raw Normal View History

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"]