feat(platform): close AI expense value loop
Add tenant-safe value, telemetry, connector, commercial, and production-readiness foundations.
This commit is contained in:
387
server/tests/test_commercial_concurrency_postgres.py
Normal file
387
server/tests/test_commercial_concurrency_postgres.py
Normal file
@@ -0,0 +1,387 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
|
||||
from savings_postgres_testkit import ( # noqa: F401 - 注册 pg_factory fixture
|
||||
_pg_factory_fixture,
|
||||
)
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from app.models.commercial import CommercialCostEvent, UsageMeterEvent
|
||||
from app.models.commercial_billing import CommercialBillingPeriod
|
||||
from app.models.commercial_runtime import CommercialRuntimeReservation
|
||||
from app.schemas.commercial import (
|
||||
CommercialCostEventCreate,
|
||||
CommercialEntitlementUpsert,
|
||||
CommercialPlanCreate,
|
||||
CommercialSubscriptionCreate,
|
||||
UsageMeterEventCreate,
|
||||
)
|
||||
from app.services.commercial_access_policy import CommercialConflictError
|
||||
from app.services.commercial_admin import CommercialAdminService
|
||||
from app.services.commercial_metering import CommercialMeteringService
|
||||
from app.services.commercial_runtime_reservations import (
|
||||
CommercialRuntimeReservationService,
|
||||
)
|
||||
from app.services.commercial_subscription_rollover import (
|
||||
CommercialSubscriptionRolloverService,
|
||||
)
|
||||
|
||||
|
||||
def test_concurrent_usage_replay_creates_one_event(
|
||||
pg_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
tenant_id, subscription_id, entitlement_id = _seed_account(pg_factory)
|
||||
payload = UsageMeterEventCreate(
|
||||
subscription_id=subscription_id,
|
||||
entitlement_id=entitlement_id,
|
||||
quantity=Decimal("2"),
|
||||
occurred_at=datetime.now(UTC),
|
||||
source_system="postgres-runtime",
|
||||
idempotency_key=f"usage-{uuid.uuid4().hex}",
|
||||
)
|
||||
ready = threading.Barrier(2)
|
||||
|
||||
def record_once() -> bool:
|
||||
with pg_factory() as db:
|
||||
ready.wait(timeout=5)
|
||||
_, created = CommercialMeteringService(db).record_usage(
|
||||
tenant_id,
|
||||
payload,
|
||||
actor_type="system",
|
||||
actor_id="postgres-runtime",
|
||||
)
|
||||
db.commit()
|
||||
return created
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
outcomes = [
|
||||
future.result(timeout=10)
|
||||
for future in (pool.submit(record_once), pool.submit(record_once))
|
||||
]
|
||||
assert sorted(outcomes) == [False, True]
|
||||
with pg_factory() as db:
|
||||
assert (
|
||||
db.scalar(
|
||||
select(func.count())
|
||||
.select_from(UsageMeterEvent)
|
||||
.where(
|
||||
UsageMeterEvent.tenant_id == tenant_id,
|
||||
UsageMeterEvent.idempotency_key == payload.idempotency_key,
|
||||
)
|
||||
)
|
||||
== 1
|
||||
)
|
||||
|
||||
|
||||
def test_concurrent_usage_cannot_bypass_hard_limit(
|
||||
pg_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
tenant_id, subscription_id, entitlement_id = _seed_account(pg_factory)
|
||||
ready = threading.Barrier(2)
|
||||
|
||||
def consume_once(suffix: str) -> str:
|
||||
payload = UsageMeterEventCreate(
|
||||
subscription_id=subscription_id,
|
||||
entitlement_id=entitlement_id,
|
||||
quantity=Decimal("4"),
|
||||
occurred_at=datetime.now(UTC),
|
||||
source_system="postgres-runtime",
|
||||
idempotency_key=f"quota-{suffix}-{uuid.uuid4().hex}",
|
||||
)
|
||||
with pg_factory() as db:
|
||||
ready.wait(timeout=5)
|
||||
try:
|
||||
CommercialMeteringService(db).record_usage(
|
||||
tenant_id,
|
||||
payload,
|
||||
actor_type="system",
|
||||
actor_id="postgres-runtime",
|
||||
)
|
||||
db.commit()
|
||||
return "created"
|
||||
except CommercialConflictError:
|
||||
db.rollback()
|
||||
return "blocked"
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
outcomes = [
|
||||
future.result(timeout=10)
|
||||
for future in (pool.submit(consume_once, "a"), pool.submit(consume_once, "b"))
|
||||
]
|
||||
assert sorted(outcomes) == ["blocked", "created"]
|
||||
with pg_factory() as db:
|
||||
used = db.scalar(
|
||||
select(func.sum(UsageMeterEvent.quantity)).where(
|
||||
UsageMeterEvent.tenant_id == tenant_id,
|
||||
UsageMeterEvent.entitlement_id == entitlement_id,
|
||||
)
|
||||
)
|
||||
assert Decimal(used or 0) == Decimal("4.000000")
|
||||
|
||||
|
||||
def test_concurrent_cost_reversal_has_one_winner(
|
||||
pg_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
tenant_id, subscription_id, _ = _seed_account(pg_factory)
|
||||
original_payload = _cost_payload(
|
||||
subscription_id,
|
||||
idempotency_key=f"cost-{uuid.uuid4().hex}",
|
||||
)
|
||||
with pg_factory() as db:
|
||||
original, _ = CommercialMeteringService(db).record_cost(tenant_id, original_payload)
|
||||
db.commit()
|
||||
original_id = original.id
|
||||
ready = threading.Barrier(2)
|
||||
|
||||
def reverse_once(suffix: str) -> str:
|
||||
payload = original_payload.model_copy(
|
||||
update={
|
||||
"event_type": "reversal",
|
||||
"idempotency_key": f"reversal-{suffix}-{uuid.uuid4().hex}",
|
||||
"reversal_of_cost_event_id": original_id,
|
||||
"occurred_at": datetime.now(UTC),
|
||||
}
|
||||
)
|
||||
with pg_factory() as db:
|
||||
ready.wait(timeout=5)
|
||||
try:
|
||||
CommercialMeteringService(db).record_cost(tenant_id, payload)
|
||||
db.commit()
|
||||
return "created"
|
||||
except CommercialConflictError:
|
||||
db.rollback()
|
||||
return "blocked"
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
outcomes = [
|
||||
future.result(timeout=10)
|
||||
for future in (pool.submit(reverse_once, "a"), pool.submit(reverse_once, "b"))
|
||||
]
|
||||
assert sorted(outcomes) == ["blocked", "created"]
|
||||
with pg_factory() as db:
|
||||
assert (
|
||||
db.scalar(
|
||||
select(func.count())
|
||||
.select_from(CommercialCostEvent)
|
||||
.where(
|
||||
CommercialCostEvent.tenant_id == tenant_id,
|
||||
CommercialCostEvent.reversal_of_cost_event_id == original_id,
|
||||
)
|
||||
)
|
||||
== 1
|
||||
)
|
||||
|
||||
|
||||
def test_concurrent_runtime_reservations_cannot_oversell_hard_quota(
|
||||
pg_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
runtime_meter = {
|
||||
"enabled": True,
|
||||
"tool_type": "llm",
|
||||
"tool_name": "chat.completions",
|
||||
"quantity_basis": "call",
|
||||
}
|
||||
tenant_id, subscription_id, entitlement_id = _seed_account(
|
||||
pg_factory,
|
||||
hard_limit=Decimal("1"),
|
||||
runtime_meter=runtime_meter,
|
||||
)
|
||||
ready = threading.Barrier(2)
|
||||
|
||||
def reserve_once(suffix: str) -> str:
|
||||
with pg_factory() as db:
|
||||
ready.wait(timeout=5)
|
||||
try:
|
||||
CommercialRuntimeReservationService(db).reserve(
|
||||
tenant_id=tenant_id,
|
||||
entitlement_id=entitlement_id,
|
||||
run_id=f"run-{suffix}-{uuid.uuid4().hex}",
|
||||
tool_call_id=str(uuid.uuid4()),
|
||||
tool_type="llm",
|
||||
tool_name="chat.completions",
|
||||
quantity_basis="call",
|
||||
reserved_quantity=Decimal("1"),
|
||||
meter_config=runtime_meter,
|
||||
)
|
||||
db.commit()
|
||||
return "reserved"
|
||||
except CommercialConflictError:
|
||||
db.rollback()
|
||||
return "blocked"
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
outcomes = [
|
||||
future.result(timeout=10)
|
||||
for future in (pool.submit(reserve_once, "a"), pool.submit(reserve_once, "b"))
|
||||
]
|
||||
assert sorted(outcomes) == ["blocked", "reserved"]
|
||||
with pg_factory() as db:
|
||||
held = db.scalar(
|
||||
select(func.sum(CommercialRuntimeReservation.reserved_quantity)).where(
|
||||
CommercialRuntimeReservation.tenant_id == tenant_id,
|
||||
CommercialRuntimeReservation.subscription_id == subscription_id,
|
||||
CommercialRuntimeReservation.entitlement_id == entitlement_id,
|
||||
CommercialRuntimeReservation.status == "reserved",
|
||||
)
|
||||
)
|
||||
assert Decimal(held or 0) == Decimal("1.000000")
|
||||
assert (
|
||||
db.scalar(
|
||||
select(func.count())
|
||||
.select_from(UsageMeterEvent)
|
||||
.where(UsageMeterEvent.tenant_id == tenant_id)
|
||||
)
|
||||
== 0
|
||||
)
|
||||
|
||||
|
||||
def test_concurrent_subscription_rollover_creates_one_period(
|
||||
pg_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
tenant_id = f"tenant-rollover-{uuid.uuid4().hex}"
|
||||
period_start = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
period_end = datetime(2026, 2, 1, tzinfo=UTC)
|
||||
with pg_factory() as db:
|
||||
admin = CommercialAdminService(db)
|
||||
plan = admin.create_plan(
|
||||
tenant_id,
|
||||
CommercialPlanCreate(
|
||||
plan_code="rollover",
|
||||
name="并发续期版",
|
||||
pricing_model="subscription",
|
||||
billing_interval="monthly",
|
||||
currency="CNY",
|
||||
base_fee=Decimal("100"),
|
||||
effective_from=datetime(2025, 12, 1, tzinfo=UTC),
|
||||
),
|
||||
actor_id="postgres-test",
|
||||
)
|
||||
admin.activate_plan(tenant_id, plan.id, expected_version=plan.version)
|
||||
subscription = admin.create_subscription(
|
||||
tenant_id,
|
||||
CommercialSubscriptionCreate(
|
||||
subscription_key=f"rollover-{uuid.uuid4().hex}",
|
||||
plan_id=plan.id,
|
||||
starts_at=period_start,
|
||||
current_period_start=period_start,
|
||||
current_period_end=period_end,
|
||||
seats=1,
|
||||
auto_renew=True,
|
||||
),
|
||||
actor_id="postgres-test",
|
||||
)
|
||||
db.commit()
|
||||
subscription_id = subscription.id
|
||||
ready = threading.Barrier(2)
|
||||
|
||||
def rollover_once() -> str:
|
||||
with pg_factory() as db:
|
||||
ready.wait(timeout=5)
|
||||
result = CommercialSubscriptionRolloverService(db).rollover_due(
|
||||
tenant_id,
|
||||
subscription_id,
|
||||
as_of=period_end,
|
||||
)
|
||||
db.commit()
|
||||
return result.status
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
outcomes = [
|
||||
future.result(timeout=10)
|
||||
for future in (pool.submit(rollover_once), pool.submit(rollover_once))
|
||||
]
|
||||
assert sorted(outcomes) == ["not_due", "rolled_over"]
|
||||
with pg_factory() as db:
|
||||
count = db.scalar(
|
||||
select(func.count())
|
||||
.select_from(CommercialBillingPeriod)
|
||||
.where(
|
||||
CommercialBillingPeriod.tenant_id == tenant_id,
|
||||
CommercialBillingPeriod.subscription_id == subscription_id,
|
||||
)
|
||||
)
|
||||
assert count == 2
|
||||
|
||||
|
||||
def _seed_account(
|
||||
factory: sessionmaker[Session],
|
||||
*,
|
||||
hard_limit: Decimal = Decimal("6"),
|
||||
runtime_meter: dict[str, object] | None = None,
|
||||
) -> tuple[str, str, str]:
|
||||
tenant_id = f"tenant-commercial-{uuid.uuid4().hex}"
|
||||
now = datetime.now(UTC)
|
||||
with factory() as db:
|
||||
service = CommercialAdminService(db)
|
||||
plan = service.create_plan(
|
||||
tenant_id,
|
||||
CommercialPlanCreate(
|
||||
plan_code="concurrency",
|
||||
name="并发验证版",
|
||||
pricing_model="usage",
|
||||
billing_interval="monthly",
|
||||
currency="CNY",
|
||||
base_fee=Decimal("0"),
|
||||
overage_enabled=False,
|
||||
effective_from=now - timedelta(days=30),
|
||||
),
|
||||
actor_id="postgres-test",
|
||||
)
|
||||
service.activate_plan(tenant_id, plan.id, expected_version=plan.version)
|
||||
subscription = service.create_subscription(
|
||||
tenant_id,
|
||||
CommercialSubscriptionCreate(
|
||||
subscription_key=f"subscription-{uuid.uuid4().hex}",
|
||||
plan_id=plan.id,
|
||||
starts_at=now - timedelta(days=5),
|
||||
current_period_start=now - timedelta(days=1),
|
||||
current_period_end=now + timedelta(days=29),
|
||||
seats=1,
|
||||
),
|
||||
actor_id="postgres-test",
|
||||
)
|
||||
entitlement = service.upsert_entitlement(
|
||||
tenant_id,
|
||||
CommercialEntitlementUpsert(
|
||||
subscription_id=subscription.id,
|
||||
entitlement_key="concurrency_usage",
|
||||
metric_key="concurrency_usage",
|
||||
entitlement_type="metered",
|
||||
unit="run",
|
||||
included_quantity=hard_limit,
|
||||
hard_limit_quantity=hard_limit,
|
||||
reset_interval="monthly",
|
||||
overage_policy="block",
|
||||
effective_from=now - timedelta(days=5),
|
||||
config_json={"runtime_meter": runtime_meter} if runtime_meter else {},
|
||||
),
|
||||
)
|
||||
db.commit()
|
||||
return tenant_id, subscription.id, entitlement.id
|
||||
|
||||
|
||||
def _cost_payload(
|
||||
subscription_id: str,
|
||||
*,
|
||||
idempotency_key: str,
|
||||
) -> CommercialCostEventCreate:
|
||||
return CommercialCostEventCreate(
|
||||
subscription_id=subscription_id,
|
||||
cost_category="ai_inference",
|
||||
quantity=Decimal("10"),
|
||||
unit="1k_tokens",
|
||||
unit_cost=Decimal("1"),
|
||||
original_currency="CNY",
|
||||
reporting_currency="CNY",
|
||||
fx_rate=Decimal("1"),
|
||||
allocation_key="postgres-concurrency",
|
||||
occurred_at=datetime.now(UTC),
|
||||
source_system="postgres-provider",
|
||||
idempotency_key=idempotency_key,
|
||||
)
|
||||
Reference in New Issue
Block a user