DealDocumentScreening/tests/unit/test_compensation.py

142 lines
4.7 KiB
Python

"""Unit tests for Document Slot compensation (ticket 001)."""
from __future__ import annotations
import uuid
import pytest
from sqlalchemy import text
from contract_check.core.billing.quota import compensate_document_slot, reserve_document_slot
pytestmark = pytest.mark.unit
async def _seed_user(session, *, telegram_id: int, credits: int = 0) -> uuid.UUID:
result = await session.execute(
text("INSERT INTO users (telegram_id, credits_left) VALUES (:t, :c) RETURNING id"),
{"t": telegram_id, "c": credits},
)
await session.flush()
return result.scalar_one()
async def _seed_plan_and_subscription(session, user_id: uuid.UUID) -> uuid.UUID:
result = await session.execute(
text(
"INSERT INTO plans (code, name, price_kopecks, monthly_quota, is_active, sort) "
"VALUES ('basic', 'Basic', 9900, 5, TRUE, 1) RETURNING code"
)
)
plan_code = result.scalar_one()
result = await session.execute(
text(
"INSERT INTO subscriptions "
"(id, user_id, plan_code, status, current_period_start, current_period_end) "
"VALUES (gen_random_uuid(), :u, :p, 'active', now() - interval '1 day', now() + interval '30 days') "
"RETURNING id"
),
{"u": user_id, "p": plan_code},
)
await session.flush()
return result.scalar_one()
async def _seed_document(session, user_id: uuid.UUID) -> uuid.UUID:
result = await session.execute(
text(
"INSERT INTO documents (user_id, s3_key, filename, mime, bytes, status) "
"VALUES (:u, 's3/key', 'file.pdf', 'application/pdf', 100, 'queued') "
"RETURNING id"
),
{"u": user_id},
)
await session.flush()
return result.scalar_one()
async def test_compensate_quota_paid_restores_quota_not_credits(db_session) -> None:
user_id = await _seed_user(db_session, telegram_id=1001, credits=3)
sub_id = await _seed_plan_and_subscription(db_session, user_id)
document_id = await _seed_document(db_session, user_id)
source = await reserve_document_slot(db_session, user_id, document_id)
assert source == "quota"
quota_released, credit_refunded = await compensate_document_slot(
db_session, document_id, "llm_quota", "all"
)
assert quota_released is True
assert credit_refunded is False
row = await db_session.execute(
text("SELECT refunded FROM documents WHERE id = :d"), {"d": document_id}
)
assert row.scalar_one() is True
credits = await db_session.execute(
text("SELECT credits_left FROM users WHERE id = :u"), {"u": user_id}
)
assert credits.scalar_one() == 3
quota = await db_session.execute(
text("SELECT count(*) FROM quota_usage WHERE subscription_id = :s"),
{"s": sub_id},
)
assert quota.scalar_one() == 0
async def test_compensate_credits_paid_refunds_once(db_session) -> None:
user_id = await _seed_user(db_session, telegram_id=1002, credits=3)
document_id = await _seed_document(db_session, user_id)
source = await reserve_document_slot(db_session, user_id, document_id)
assert source == "credits"
quota_released, credit_refunded = await compensate_document_slot(
db_session, document_id, "llm_quota", "all"
)
assert quota_released is False
assert credit_refunded is True
row = await db_session.execute(
text(
"SELECT credits_left, refunded FROM users u JOIN documents d ON d.user_id = u.id "
"WHERE d.id = :d"
),
{"d": document_id},
)
credits_left, refunded = row.one()
assert credits_left == 3
assert refunded is True
# Second execution must be a no-op.
quota_released2, credit_refunded2 = await compensate_document_slot(
db_session, document_id, "llm_quota", "all"
)
assert quota_released2 is False
assert credit_refunded2 is False
credits2 = await db_session.execute(
text("SELECT credits_left FROM users WHERE id = :u"), {"u": user_id}
)
assert credits2.scalar_one() == 3
async def test_compensate_infra_only_policy_skips_extraction_failed(db_session) -> None:
user_id = await _seed_user(db_session, telegram_id=1003, credits=3)
document_id = await _seed_document(db_session, user_id)
source = await reserve_document_slot(db_session, user_id, document_id)
assert source == "credits"
quota_released, credit_refunded = await compensate_document_slot(
db_session, document_id, "extraction_failed", "infra_only"
)
assert quota_released is False
assert credit_refunded is False
row = await db_session.execute(
text("SELECT refunded FROM documents WHERE id = :d"), {"d": document_id}
)
assert row.scalar_one() is False