233 lines
8.5 KiB
Python
233 lines
8.5 KiB
Python
"""Unit tests for the repository layer.
|
|
|
|
Run against the compose Postgres stack; skipped automatically if it is down.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
|
|
import pytest
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from contract_check.core.db.repositories import (
|
|
ApiKeyRepository,
|
|
CreditsRepository,
|
|
DocumentRepository,
|
|
JobRepository,
|
|
ReportRepository,
|
|
UserRepository,
|
|
)
|
|
|
|
pytestmark = pytest.mark.usefixtures("db_session")
|
|
|
|
|
|
class TestUserRepository:
|
|
async def test_create_and_get_email_user(self, db_session: AsyncSession) -> None:
|
|
repo = UserRepository(db_session)
|
|
user = await repo.create_email_user(
|
|
email="repo-test@example.com",
|
|
name="Repo Test",
|
|
password_hash="hash",
|
|
)
|
|
assert user.email == "repo-test@example.com"
|
|
loaded = await repo.get_by_id(user.id)
|
|
assert loaded is not None
|
|
assert loaded.email == user.email
|
|
|
|
async def test_telegram_user_unique(self, db_session: AsyncSession) -> None:
|
|
repo = UserRepository(db_session)
|
|
first = await repo.create_telegram_user(telegram_id=987654321)
|
|
second = await repo.create_telegram_user(telegram_id=987654322)
|
|
assert first.telegram_id != second.telegram_id
|
|
assert await repo.get_by_telegram_id(987654321) is not None
|
|
|
|
async def test_password_and_credits_adjustments(self, db_session: AsyncSession) -> None:
|
|
repo = UserRepository(db_session)
|
|
user = await repo.create_email_user(email="adjust@example.com", name="A", password_hash="h")
|
|
await repo.set_password(user.id, "new_hash")
|
|
await repo.adjust_credits(user.id, 5)
|
|
loaded = await repo.get_by_id(user.id)
|
|
assert loaded is not None
|
|
assert loaded.password_hash == "new_hash"
|
|
assert loaded.credits_left == 5
|
|
|
|
|
|
class TestCreditsRepository:
|
|
async def test_reserve_refund_cycle(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
credits = CreditsRepository(db_session)
|
|
user = await users.create_email_user(
|
|
email="credits@example.com", name="C", password_hash="h"
|
|
)
|
|
await users.adjust_credits(user.id, 10)
|
|
await db_session.commit()
|
|
|
|
assert await credits.reserve(user.id) is True
|
|
await db_session.commit()
|
|
assert await credits.get_balance(user.id) == 9
|
|
|
|
# Refund requires a document, so create one.
|
|
docs = DocumentRepository(db_session)
|
|
doc = await docs.create(
|
|
document_id=uuid.uuid4(),
|
|
user_id=user.id,
|
|
s3_key="s3://x",
|
|
filename="x.pdf",
|
|
mime="application/pdf",
|
|
bytes_=123,
|
|
)
|
|
await db_session.commit()
|
|
|
|
refunded = await credits.refund(doc.id, "infra", "infra_only")
|
|
await db_session.commit()
|
|
assert refunded is True
|
|
assert await credits.get_balance(user.id) == 10
|
|
|
|
# Idempotent second refund.
|
|
assert await credits.refund(doc.id, "infra", "infra_only") is False
|
|
|
|
|
|
class TestDocumentRepository:
|
|
async def test_create_update_status(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
repo = DocumentRepository(db_session)
|
|
user = await users.create_email_user(email="doc@example.com", name="D", password_hash="h")
|
|
doc_id = uuid.uuid4()
|
|
doc = await repo.create(
|
|
document_id=doc_id,
|
|
user_id=user.id,
|
|
s3_key="k",
|
|
filename="f.pdf",
|
|
mime="application/pdf",
|
|
bytes_=100,
|
|
)
|
|
await repo.update_status(doc.id, status="extracting", stage="downloading")
|
|
status = await repo.get_status(doc.id)
|
|
assert status == "extracting"
|
|
|
|
async def test_for_update_returns_status(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
repo = DocumentRepository(db_session)
|
|
user = await users.create_email_user(email="lock@example.com", name="L", password_hash="h")
|
|
doc_id = uuid.uuid4()
|
|
await repo.create(
|
|
document_id=doc_id,
|
|
user_id=user.id,
|
|
s3_key="k",
|
|
filename="f.pdf",
|
|
mime="application/pdf",
|
|
bytes_=100,
|
|
)
|
|
status, filename = await repo.get_status_and_filename_for_update(doc_id)
|
|
assert status == "queued"
|
|
assert filename == "f.pdf"
|
|
|
|
|
|
class TestJobRepository:
|
|
async def test_create_claim_done(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
docs = DocumentRepository(db_session)
|
|
repo = JobRepository(db_session)
|
|
user = await users.create_email_user(email="job@example.com", name="J", password_hash="h")
|
|
doc = await docs.create(
|
|
document_id=uuid.uuid4(),
|
|
user_id=user.id,
|
|
s3_key="k",
|
|
filename="f.pdf",
|
|
mime="application/pdf",
|
|
bytes_=1,
|
|
)
|
|
cid = uuid.uuid4()
|
|
job = await repo.create(document_id=doc.id, correlation_id=cid, queue="extract")
|
|
assert job.status == "pending"
|
|
await repo.claim_start(doc.id, "extract")
|
|
loaded = await repo.get_by_document_id_and_queue(doc.id, "extract")
|
|
assert loaded is not None
|
|
assert loaded.status == "running"
|
|
assert loaded.attempts == 1
|
|
await repo.mark_done(doc.id, "extract")
|
|
loaded = await repo.get_by_document_id_and_queue(doc.id, "extract")
|
|
assert loaded is not None
|
|
assert loaded.status == "done"
|
|
|
|
|
|
class TestApiKeyRepository:
|
|
async def test_key_hash_lookup_and_quota(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
repo = ApiKeyRepository(db_session)
|
|
user = await users.create_email_user(email="key@example.com", name="K", password_hash="h")
|
|
key = await repo.create(
|
|
user_id=user.id, name="test", key_hash="sha256-deadbeef", monthly_quota=10
|
|
)
|
|
found = await repo.get_by_hash("sha256-deadbeef")
|
|
assert found is not None
|
|
assert found.id == key.id
|
|
assert not found.revoked
|
|
quota, used = await repo.get_monthly_quota_state(key.id)
|
|
assert quota == 10
|
|
assert used == 0
|
|
await repo.bump_monthly_used(key.id)
|
|
quota, used = await repo.get_monthly_quota_state(key.id)
|
|
assert used == 1
|
|
|
|
|
|
class TestReportRepository:
|
|
async def test_upsert_analyze_report(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
docs = DocumentRepository(db_session)
|
|
repo = ReportRepository(db_session)
|
|
user = await users.create_email_user(
|
|
email="report@example.com", name="R", password_hash="h"
|
|
)
|
|
doc = await docs.create(
|
|
document_id=uuid.uuid4(),
|
|
user_id=user.id,
|
|
s3_key="k",
|
|
filename="f.pdf",
|
|
mime="application/pdf",
|
|
bytes_=1,
|
|
)
|
|
report_id = await repo.upsert_analyze_report(
|
|
document_id=doc.id,
|
|
content_json={"findings": []},
|
|
markdown="# Report",
|
|
model_used="gpt-4",
|
|
prompt_tokens=10,
|
|
eval_tokens=20,
|
|
latency_ms=100,
|
|
)
|
|
loaded = await repo.get_by_document_id(doc.id)
|
|
assert loaded is not None
|
|
assert loaded.id == report_id
|
|
assert loaded.markdown == "# Report"
|
|
|
|
# Upsert is idempotent and updates.
|
|
await repo.upsert_analyze_report(
|
|
document_id=doc.id,
|
|
content_json={"findings": [{"x": 1}]},
|
|
markdown="# Updated",
|
|
model_used="gpt-4",
|
|
prompt_tokens=11,
|
|
eval_tokens=21,
|
|
latency_ms=101,
|
|
)
|
|
await db_session.refresh(loaded)
|
|
assert loaded.markdown == "# Updated"
|
|
|
|
|
|
class TestMagicLink:
|
|
async def test_round_trip(self, db_session: AsyncSession) -> None:
|
|
users = UserRepository(db_session)
|
|
user = await users.create_email_user(email="magic@example.com", name="M", password_hash="h")
|
|
expires = datetime.now(UTC) + timedelta(hours=1)
|
|
await users.set_magic_link_token(user.id, "sha256-abc", expires)
|
|
loaded = await users.get_by_magic_link_token_hash("sha256-abc")
|
|
assert loaded is not None
|
|
assert loaded.id == user.id
|
|
await users.consume_magic_link_token(user.id)
|
|
consumed = await users.get_by_id(user.id)
|
|
assert consumed is not None
|
|
assert consumed.magic_link_token_hash is None
|
|
assert consumed.magic_link_expires_at is None
|