"""Unit tests for the billing scheduler (ticket 016). Drive scheduler.run_tick with an injected clock and a fake provider; assert state transitions without touching the network. """ from __future__ import annotations import asyncio import datetime as dt import uuid from dataclasses import dataclass from typing import Any import pytest from sqlalchemy import text from contract_check.core.billing.port import PaymentCreated, PaymentProvider, PaymentStatusInfo from contract_check.core.config import get_settings from contract_check.core.db.session import create_session_factory from contract_check.worker_billing.scheduler import run_tick pytestmark = [pytest.mark.unit, pytest.mark.integration] @dataclass class FakeProvider(PaymentProvider): payments: dict[str, dict[str, Any]] created: list[PaymentCreated] = None # type: ignore[assignment] status_overrides: dict[str, str] | None = None def __post_init__(self) -> None: self.created = [] self.status_overrides = {} async def aclose(self) -> None: pass async def create_payment( self, *, amount_kopecks: int, description: str, return_url: str, metadata: dict[str, str], idempotency_key: str, ) -> PaymentCreated: payment_id = f"pay-{metadata['invoice_id']}" self.payments[payment_id] = { "amount_kopecks": amount_kopecks, "status": "pending", "metadata": metadata, } created = PaymentCreated(payment_id=payment_id, confirmation_url="https://test/pay") self.created.append(created) return created async def get_payment(self, payment_id: str) -> PaymentStatusInfo: data = self.payments.get(payment_id, {}) status = self.status_overrides.get(payment_id, data.get("status", "pending")) return PaymentStatusInfo( payment_id=payment_id, status=status, paid=status == "succeeded", amount_kopecks=data.get("amount_kopecks", 0), metadata={str(k): str(v) for k, v in data.get("metadata", {}).items()}, ) async def refund(self, *, payment_id: str, amount_kopecks: int, idempotency_key: str) -> Any: raise NotImplementedError @pytest.fixture async def factory(): import os # Integration conftest (_set_env) already points DATABASE_URL at the test # infra (:25432); only fall back to the dev port when it is unset, and # restore the previous value afterwards. prev_url = os.environ.get("DATABASE_URL") os.environ.setdefault( "DATABASE_URL", "postgresql+asyncpg://contract_check:contract_check@localhost:15432/contract_check", ) os.environ["PLANS_ENABLED"] = "true" os.environ["YOOKASSA_ENABLED"] = "true" os.environ["YOOKASSA_SHOP_ID"] = "123456" os.environ["YOOKASSA_SECRET_KEY"] = "test-secret" os.environ["BILLING_RETURN_JWT_SECRET"] = "pay-secret" get_settings.cache_clear() f = create_session_factory() yield f if prev_url is None: os.environ.pop("DATABASE_URL", None) else: os.environ["DATABASE_URL"] = prev_url os.environ.pop("PLANS_ENABLED", None) os.environ.pop("YOOKASSA_ENABLED", None) os.environ.pop("YOOKASSA_SHOP_ID", None) os.environ.pop("YOOKASSA_SECRET_KEY", None) os.environ.pop("BILLING_RETURN_JWT_SECRET", None) get_settings.cache_clear() async def _seed_user(session, telegram_id: int = 123456789) -> uuid.UUID: result = await session.execute( text("SELECT id FROM users WHERE telegram_id = :t"), {"t": telegram_id}, ) row = result.first() if row is not None: return row[0] result = await session.execute( text( "INSERT INTO users (id, telegram_id, credits_left) " "VALUES (gen_random_uuid(), :t, 0) RETURNING id" ), {"t": telegram_id}, ) return result.scalar_one() async def _seed_subscription( session, user_id: uuid.UUID, *, auto_renew: bool = True, period_end_offset_days: int = 2, plan_code: str = "pro", ) -> uuid.UUID: now = dt.datetime.now(tz=dt.UTC) # Clean any existing active/past_due subscription for the user to keep the partial unique index happy. await session.execute( text( "UPDATE subscriptions SET status = 'expired' WHERE user_id = :u AND status IN ('active','past_due')" ), {"u": user_id}, ) result = await session.execute( text( "INSERT INTO subscriptions " "(id, user_id, plan_code, status, current_period_start, current_period_end, auto_renew) " "VALUES (gen_random_uuid(), :u, :plan, 'active', :start, :end, :renew) " "RETURNING id" ), { "u": user_id, "plan": plan_code, "start": now, "end": now + dt.timedelta(days=period_end_offset_days), "renew": auto_renew, }, ) return result.scalar_one() async def test_scheduler_creates_renewal_invoice(factory) -> None: provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session) sub_id = await _seed_subscription( session, user_id, auto_renew=True, period_end_offset_days=2 ) await session.commit() now = dt.datetime.now(tz=dt.UTC) settings = get_settings() await run_tick(session, provider, now, settings) invoice = ( await session.execute( text( "SELECT kind, subscription_id, external_id, status " "FROM invoices WHERE subscription_id = :s" ), {"s": sub_id}, ) ).first() assert invoice is not None assert invoice[0] == "renewal" assert invoice[1] == sub_id assert invoice[2].startswith("pay-") assert invoice[3] == "pending" async def test_scheduler_skips_non_autorenew(factory) -> None: provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session) sub_id = await _seed_subscription( session, user_id, auto_renew=False, period_end_offset_days=2 ) await session.commit() now = dt.datetime.now(tz=dt.UTC) settings = get_settings() await run_tick(session, provider, now, settings) count = ( await session.execute( text("SELECT count(*) FROM invoices WHERE subscription_id = :s"), {"s": sub_id}, ) ).scalar_one() assert int(count) == 0 async def test_scheduler_expires_without_renewal(factory) -> None: async with factory() as session: user_id = await _seed_user(session) sub_id = await _seed_subscription( session, user_id, auto_renew=False, period_end_offset_days=-1 ) await session.commit() now = dt.datetime.now(tz=dt.UTC) settings = get_settings() await run_tick(session, None, now, settings) status = ( await session.execute( text("SELECT status FROM subscriptions WHERE id = :s"), {"s": sub_id} ) ).scalar_one() assert status == "expired" async def test_scheduler_rolls_after_successful_renewal(factory) -> None: provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session) sub_id = await _seed_subscription( session, user_id, auto_renew=True, period_end_offset_days=-1 ) # Pre-create a succeeded renewal invoice. inv_id = uuid.uuid4() now = dt.datetime.now(tz=dt.UTC) await session.execute( text( "INSERT INTO invoices (id, user_id, amount, status, kind, subscription_id, paid_at) " "VALUES (:id, :u, 149000, 'succeeded', 'renewal', :s, :paid)" ), {"id": inv_id, "u": user_id, "s": sub_id, "paid": now - dt.timedelta(minutes=1)}, ) await session.commit() await run_tick(session, provider, now, settings=get_settings()) row = ( await session.execute( text( "SELECT status, current_period_start, current_period_end, origin_invoice_id " "FROM subscriptions WHERE id = :s" ), {"s": sub_id}, ) ).first() assert row[0] == "active" assert row[3] == inv_id assert row[2] > row[1] async def test_scheduler_reconciles_pending_invoice(factory) -> None: provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session) inv_id = uuid.uuid4() old = dt.datetime.now(tz=dt.UTC) - dt.timedelta(minutes=20) await session.execute( text( "INSERT INTO invoices (id, user_id, amount, status, kind, external_id, created_at) " "VALUES (:id, :u, 19900, 'pending', 'topup', :ext, :created)" ), {"id": inv_id, "u": user_id, "ext": "pay-old", "created": old}, ) provider.payments["pay-old"] = { "amount_kopecks": 19900, "status": "pending", "metadata": {"invoice_id": str(inv_id), "kind": "topup"}, } await session.commit() provider.status_overrides["pay-old"] = "succeeded" # type: ignore[index] await run_tick(session, provider, dt.datetime.now(tz=dt.UTC), settings=get_settings()) status = ( await session.execute( text("SELECT status FROM invoices WHERE id = :id"), {"id": inv_id} ) ).scalar_one() assert status == "succeeded" async def test_scheduler_tick_advisory_lock_prevents_double_run(factory) -> None: provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session) sub_id = await _seed_subscription( session, user_id, auto_renew=True, period_end_offset_days=2 ) await session.commit() now = dt.datetime.now(tz=dt.UTC) settings = get_settings() # First tick holds the advisory lock until it commits. task1 = asyncio.create_task(run_tick(session, provider, now, settings)) # Give task1 time to acquire the lock. await asyncio.sleep(0.1) # Second tick on a separate connection should skip. async with factory() as session2: await run_tick(session2, provider, now, settings) await task1 count = ( await session.execute( text("SELECT count(*) FROM invoices WHERE subscription_id = :s"), {"s": sub_id}, ) ).scalar_one() assert int(count) == 1 async def test_subscription_fulfillment_idempotent_by_invoice(factory) -> None: from contract_check.core.billing.fulfillment import apply_payment_status provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session, telegram_id=123_456_700) # Ensure no leftover subscription for this user from an interrupted run. await session.execute(text("DELETE FROM subscriptions WHERE user_id = :u"), {"u": user_id}) plan_code = "pro" await session.execute( text( "INSERT INTO plans (code, name, price_kopecks, monthly_quota, is_active, sort) " "VALUES (:code, :name, 149000, 10, TRUE, 1) " "ON CONFLICT (code) DO NOTHING" ), {"code": plan_code, "name": "Pro"}, ) inv_id = uuid.uuid4() await session.execute( text( "INSERT INTO invoices (id, user_id, amount, status, kind, external_id) " "VALUES (:id, :u, 149000, 'pending', 'subscription', 'pay-sub')" ), {"id": inv_id, "u": user_id}, ) provider.payments["pay-sub"] = { "amount_kopecks": 149000, "status": "succeeded", "metadata": {"invoice_id": str(inv_id), "kind": "subscription", "plan_code": plan_code}, } await session.commit() await apply_payment_status(session, provider, "pay-sub") await apply_payment_status(session, provider, "pay-sub") subs = await session.execute( text("SELECT count(*) FROM subscriptions WHERE user_id = :u"), {"u": user_id} ) assert int(subs.scalar_one()) == 1 inv_status = await session.execute( text("SELECT status FROM invoices WHERE id = :id"), {"id": inv_id} ) assert inv_status.scalar_one() == "succeeded" async def test_subscription_fulfillment_poison_on_conflicting_subscription(factory) -> None: from contract_check.core.billing.fulfillment import apply_payment_status provider = FakeProvider(payments={}) async with factory() as session: user_id = await _seed_user(session, telegram_id=123_456_701) await session.execute(text("DELETE FROM subscriptions WHERE user_id = :u"), {"u": user_id}) plan_code = "pro" await session.execute( text( "INSERT INTO plans (code, name, price_kopecks, monthly_quota, is_active, sort) " "VALUES (:code, :name, 149000, 10, TRUE, 1) " "ON CONFLICT (code) DO NOTHING" ), {"code": plan_code, "name": "Pro"}, ) # Pre-existing active subscription from a different invoice. other_inv_id = uuid.uuid4() await session.execute( text( "INSERT INTO invoices (id, user_id, amount, status, kind) " "VALUES (:id, :u, 149000, 'succeeded', 'subscription')" ), {"id": other_inv_id, "u": user_id}, ) await session.execute( text( "INSERT INTO subscriptions " "(id, user_id, plan_code, status, current_period_start, current_period_end, origin_invoice_id) " "VALUES (gen_random_uuid(), :u, :plan, 'active', now(), now() + interval '30 days', :inv)" ), {"u": user_id, "plan": plan_code, "inv": other_inv_id}, ) inv_id = uuid.uuid4() await session.execute( text( "INSERT INTO invoices (id, user_id, amount, status, kind, external_id) " "VALUES (:id, :u, 149000, 'pending', 'subscription', 'pay-sub2')" ), {"id": inv_id, "u": user_id}, ) provider.payments["pay-sub2"] = { "amount_kopecks": 149000, "status": "succeeded", "metadata": {"invoice_id": str(inv_id), "kind": "subscription", "plan_code": plan_code}, } await session.commit() await apply_payment_status(session, provider, "pay-sub2") inv_status = await session.execute( text("SELECT status FROM invoices WHERE id = :id"), {"id": inv_id} ) assert inv_status.scalar_one() == "cancelled" async def test_credit_refund_is_single_fire_under_concurrency(factory) -> None: from contract_check.core.credits import refund_credit async with factory() as session: user_id = await _seed_user(session, telegram_id=123_456_702) await session.execute( text("UPDATE users SET credits_left = 0 WHERE id = :u"), {"u": user_id} ) doc_id = uuid.uuid4() await session.execute( text( "INSERT INTO documents " "(id, user_id, s3_key, filename, mime, bytes, status, stage, refunded) " "VALUES (:id, :u, 's3', 'f.pdf', 'application/pdf', 100, 'failed', 'failed', FALSE)" ), {"id": doc_id, "u": user_id}, ) await session.commit() async def _refund() -> bool: async with factory() as s: result = await refund_credit(s, doc_id, "llm_quota", "all") await s.commit() return result results = await asyncio.gather(_refund(), _refund(), _refund()) assert sum(1 for r in results if r) == 1 async with factory() as session: balance = await session.execute( text("SELECT credits_left FROM users WHERE id = :u"), {"u": user_id} ) assert balance.scalar_one() == 1 events = await session.execute( text( "SELECT count(*) FROM credit_events " "WHERE user_id = :u AND kind = 'refund_auto' AND document_id = :d" ), {"u": user_id, "d": doc_id}, ) assert events.scalar_one() == 1