"""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 bridge_platform.quotas import service
from bridge_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"])

    def test_reset_starts_counter_at_zero_without_deleting_history(self):
        service.consume(tenant_id="tenant-a", app_id="app-a", metric="requests")
        reset = service.reset_usage_counter(
            tenant_id="tenant-a", metric="requests", reason="Separate production baseline",
        )
        self.assertEqual(1, reset["previous_usage"])
        self.assertEqual(0, service.usage_snapshot("tenant-a")["requests_per_minute"]["used"])
        with Session(self.engine, future=True) as session:
            records = session.query(service.TenantUsage).filter_by(tenant_id="tenant-a").all()
        self.assertEqual(2, len(records))
        self.assertEqual({"requests", "requests_reset"}, {record.metric_name for record in records})

    def test_capacity_check_blocks_before_recording(self):
        service.consume(tenant_id="tenant-a", app_id="app-a", metric="requests", amount=1)
        with self.assertRaises(service.QuotaExceeded):
            service.ensure_capacity(tenant_id="tenant-a", app_id="app-a", metric="requests", amount=2)
        self.assertEqual(1, service.usage_snapshot("tenant-a")["requests_per_minute"]["used"])

    def test_observed_usage_can_reconcile_a_reservation_downward(self):
        service.consume(tenant_id="tenant-a", app_id="app-a", metric="requests", amount=2)
        service.record_observed_usage(
            tenant_id="tenant-a", app_id="app-a", metric="requests", amount=-1,
            metadata={"kind": "provider_reservation_reconciliation"},
        )
        self.assertEqual(1, service.usage_snapshot("tenant-a")["requests_per_minute"]["used"])


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