"""Standalone persistent quota checks."""

from __future__ import annotations

import sys
import unittest
from pathlib import Path

PROJECT_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(PROJECT_ROOT))
loaded_platform = sys.modules.get("platform")
if loaded_platform is not None and not hasattr(loaded_platform, "__path__"):
    del sys.modules["platform"]

from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool

from config.db import Base
from platform.quotas import service
from platform.tenants.models import App, Tenant, TenantLimit


class QuotaChecks(unittest.TestCase):
    def setUp(self):
        self.engine = create_engine(
            "sqlite://",
            connect_args={"check_same_thread": False},
            poolclass=StaticPool,
            future=True,
        )
        Base.metadata.create_all(self.engine)
        self.original_factory = service.get_control_plane_session
        service.get_control_plane_session = lambda: Session(self.engine, future=True)
        with Session(self.engine, future=True) as session:
            session.add(Tenant(
                tenant_id="tenant-a", slug="tenant-a", name="Tenant A",
                status="active", is_active=True,
            ))
            session.add(App(
                app_id="app-a", display_name="App A",
                module_path="apps.app_a", route_prefix="/app-a", is_active=True,
            ))
            session.add(TenantLimit(
                tenant_id="tenant-a", app_id=None,
                limit_name="requests_per_minute", limit_value=2,
                window_seconds=60, overage_policy="block",
            ))
            session.commit()

    def tearDown(self):
        service.get_control_plane_session = self.original_factory
        self.engine.dispose()

    def test_usage_is_recorded_and_remaining_is_reported(self):
        first = service.consume(
            tenant_id="tenant-a", app_id="app-a", metric="requests",
        )
        self.assertTrue(first.allowed)
        self.assertEqual(1, first.remaining)
        snapshot = service.usage_snapshot("tenant-a")
        self.assertEqual(1, snapshot["requests_per_minute"]["used"])
        self.assertEqual(1, snapshot["requests_per_minute"]["remaining"])

    def test_request_is_blocked_without_recording_overage(self):
        service.consume(tenant_id="tenant-a", app_id="app-a", metric="requests")
        service.consume(tenant_id="tenant-a", app_id="app-a", metric="requests")
        with self.assertRaises(service.QuotaExceeded) as caught:
            service.consume(tenant_id="tenant-a", app_id="app-a", metric="requests")
        self.assertEqual(2, caught.exception.decision.used)
        self.assertEqual(0, caught.exception.decision.remaining)
        self.assertEqual(2, service.usage_snapshot("tenant-a")["requests_per_minute"]["used"])


if __name__ == "__main__":
    unittest.main(verbosity=2)
