# -*- coding: utf-8 -*-
"""License storage backed by PostgreSQL."""

from __future__ import annotations

import json
import secrets
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, Optional

try:
    import psycopg2
    from psycopg2 import errors
    from psycopg2.extras import RealDictCursor, Json
except ImportError as exc:
    psycopg2 = None  # type: ignore[assignment]
    RealDictCursor = Json = None
    _IMPORT_ERROR = exc
else:
    _IMPORT_ERROR = None

from ..subscription.models import SubscriptionTier, get_features_for_tier
from ..payment.database import LicenseDatabase  # reuse helpers (generate key signature)


class PostgresLicenseStore:
    """Persist licences in PostgreSQL instead of the local SQLite file."""

    def __init__(self, dsn: str):
        if psycopg2 is None:  # pragma: no cover - dépend si psycopg2 installé
            raise ImportError("psycopg2 est requis pour PostgresLicenseStore") from _IMPORT_ERROR
        self.dsn = dsn
        self._ensure_schema()

    # ------------------------------------------------------------------
    # Helpers
    # ------------------------------------------------------------------
    def _connect(self):
        return psycopg2.connect(self.dsn)

    @staticmethod
    def _generate_api_key() -> str:
        return f"ws_live_{secrets.token_urlsafe(32)}"

    @staticmethod
    def _sast_defaults_for_tier(tier: str) -> Dict[str, int | bool]:
        try:
            tier_enum = SubscriptionTier(tier.lower())
        except ValueError:
            tier_enum = SubscriptionTier.FREE
        features = get_features_for_tier(tier_enum)
        return {
            "allow_source_scan": features.allow_source_scan,
            "max_source_files": features.max_source_files,
            "max_source_size_mb": features.max_source_size_mb,
            "advanced_rules": features.advanced_rules,
        }

    def _ensure_schema(self) -> None:
        with self._connect() as conn:
            with conn.cursor() as cur:
                try:
                    cur.execute(
                        """
                        CREATE TABLE IF NOT EXISTS api_key (
                            id SERIAL PRIMARY KEY,
                            email VARCHAR(255) NOT NULL,
                            api_key VARCHAR(255) UNIQUE NOT NULL,
                            tier VARCHAR(32) NOT NULL DEFAULT 'FREE',
                            status VARCHAR(32) NOT NULL DEFAULT 'active',
                            max_domains INTEGER DEFAULT 10,
                            max_users INTEGER DEFAULT 1,
                            allow_source_scan BOOLEAN DEFAULT FALSE,
                            max_source_files INTEGER DEFAULT 0,
                            max_source_size_mb INTEGER DEFAULT 0,
                            advanced_rules BOOLEAN DEFAULT FALSE,
                            stripe_customer_id VARCHAR(255),
                            stripe_subscription_id VARCHAR(255),
                            created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
                            updated_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
                            expires_at TIMESTAMP WITH TIME ZONE,
                            usage_count INTEGER DEFAULT 0,
                            last_used_at TIMESTAMP WITH TIME ZONE,
                            events JSONB DEFAULT '[]'::jsonb
                        );
                        CREATE INDEX IF NOT EXISTS idx_api_key_lookup ON api_key(api_key);
                        CREATE INDEX IF NOT EXISTS idx_subscription_lookup ON api_key(stripe_subscription_id);
                        """
                    )
                except (errors.DuplicateObject, errors.DuplicateTable):
                    conn.rollback()

    # ------------------------------------------------------------------
    # CRUD
    # ------------------------------------------------------------------
    def create_license(
        self,
        email: str,
        tier: str,
        stripe_customer_id: Optional[str] = None,
        stripe_subscription_id: Optional[str] = None,
        max_domains: int = 10,
        max_users: int = 1,
        features: Optional[str] = None,
        expires_at: Optional[datetime] = None,
        allow_source_scan: Optional[bool] = None,
        max_source_files: Optional[int] = None,
        max_source_size_mb: Optional[int] = None,
        advanced_rules: Optional[bool] = None,
    ) -> Dict[str, Any]:
        api_key = self._generate_api_key()
        events: Iterable[Dict[str, Any]] = [
            {"event_type": "license_created", "created_at": datetime.now(timezone.utc).isoformat()}
        ]
        tier_label = tier.upper()
        sast_defaults = self._sast_defaults_for_tier(tier_label)
        sast_settings = {
            "allow_source_scan": sast_defaults["allow_source_scan"] if allow_source_scan is None else allow_source_scan,
            "max_source_files": sast_defaults["max_source_files"] if max_source_files is None else max_source_files,
            "max_source_size_mb": sast_defaults["max_source_size_mb"] if max_source_size_mb is None else max_source_size_mb,
            "advanced_rules": sast_defaults["advanced_rules"] if advanced_rules is None else advanced_rules,
        }

        with self._connect() as conn:
            with conn.cursor(cursor_factory=RealDictCursor) as cur:
                cur.execute(
                    """
                    INSERT INTO api_key (
                        email, api_key, tier, status, max_domains, max_users,
                        allow_source_scan, max_source_files, max_source_size_mb, advanced_rules,
                        stripe_customer_id, stripe_subscription_id, expires_at,
                        events
                    )
                    VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
                    RETURNING email, api_key, tier, status, max_domains, max_users,
                              allow_source_scan, max_source_files, max_source_size_mb, advanced_rules,
                              stripe_customer_id, stripe_subscription_id, created_at, expires_at;
                    """,
                    (
                        email,
                        api_key,
                        tier_label,
                        "active",
                        max_domains,
                        max_users,
                        sast_settings["allow_source_scan"],
                        sast_settings["max_source_files"],
                        sast_settings["max_source_size_mb"],
                        sast_settings["advanced_rules"],
                        stripe_customer_id,
                        stripe_subscription_id,
                        expires_at,
                        Json(list(events)),
                    ),
                )
                row = cur.fetchone()

        return self._map_row(row)

    def get_license_by_email(self, email: str) -> Optional[Dict[str, Any]]:
        with self._connect() as conn:
            with conn.cursor(cursor_factory=RealDictCursor) as cur:
                cur.execute(
                    """
                    SELECT * FROM api_key
                    WHERE email = %s
                    ORDER BY created_at DESC
                    LIMIT 1
                    """,
                    (email,),
                )
                row = cur.fetchone()
        if not row:
            return None
        return self._map_row(row)

    def get_license_by_api_key(self, api_key: str) -> Optional[Dict[str, Any]]:
        with self._connect() as conn:
            with conn.cursor(cursor_factory=RealDictCursor) as cur:
                cur.execute(
                    """
                    SELECT * FROM api_key WHERE api_key = %s
                    """,
                    (api_key,),
                )
                row = cur.fetchone()
        if not row:
            return None
        return self._map_row(row)

    def get_license_by_stripe_subscription(self, subscription_id: str) -> Optional[Dict[str, Any]]:
        if not subscription_id:
            return None
        with self._connect() as conn:
            with conn.cursor(cursor_factory=RealDictCursor) as cur:
                cur.execute(
                    """
                    SELECT * FROM api_key WHERE stripe_subscription_id = %s
                    """,
                    (subscription_id,),
                )
                row = cur.fetchone()
        if not row:
            return None
        return self._map_row(row)

    def update_license_usage(self, api_key: str) -> None:
        with self._connect() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    """
                    UPDATE api_key
                    SET usage_count = COALESCE(usage_count, 0) + 1,
                        last_used_at = NOW(),
                        updated_at = NOW()
                    WHERE api_key = %s
                    """,
                    (api_key,),
                )

    def update_license_status(
        self,
        api_key: str,
        status: str,
        event_type: Optional[str] = None,
        event_data: Optional[str] = None,
    ) -> None:
        with self._connect() as conn:
            with conn.cursor(cursor_factory=RealDictCursor) as cur:
                cur.execute("SELECT events FROM api_key WHERE api_key = %s", (api_key,))
                row = cur.fetchone()
                events = row["events"] if row else []
                if event_type:
                    events.append(
                        {
                            "event_type": event_type,
                            "event_data": event_data,
                            "created_at": datetime.now(timezone.utc).isoformat(),
                        }
                    )
                cur.execute(
                    """
                    UPDATE api_key
                    SET status = %s,
                        updated_at = NOW(),
                        events = %s
                    WHERE api_key = %s
                    """,
                    (status, Json(events), api_key),
                )

    # ------------------------------------------------------------------
    # Validation
    # ------------------------------------------------------------------
    def validate_license(self, api_key: str) -> Dict[str, Any]:
        license_data = self.get_license_by_api_key(api_key)
        if not license_data:
            return {"valid": False, "error": "API key invalide", "status": "not_found"}

        if license_data["status"] != "active":
            return {
                "valid": False,
                "error": f"Licence {license_data['status']}",
                "status": license_data["status"],
                "tier": license_data["tier"],
            }

        expires_at = license_data.get("expires_at")
        if expires_at:
            if isinstance(expires_at, str):
                expires_dt = datetime.fromisoformat(expires_at)
            else:
                expires_dt = expires_at
            if expires_dt < datetime.now(timezone.utc):
                self.update_license_status(api_key, "expired", "license_expired")
                return {
                    "valid": False,
                    "error": "Licence expirée",
                    "status": "expired",
                    "tier": license_data["tier"],
                }

        self.update_license_usage(api_key)
        refreshed = self.get_license_by_api_key(api_key) or license_data
        return {
            "valid": True,
            "status": refreshed["status"],
            "tier": refreshed["tier"],
            "email": refreshed["email"],
            "max_domains": refreshed["max_domains"],
            "max_users": refreshed["max_users"],
            "allow_source_scan": refreshed["allow_source_scan"],
            "max_source_files": refreshed["max_source_files"],
            "max_source_size_mb": refreshed["max_source_size_mb"],
            "advanced_rules": refreshed["advanced_rules"],
            "usage_count": refreshed.get("usage_count", 0),
            "created_at": refreshed.get("created_at"),
        }

    # ------------------------------------------------------------------
    # Utilities
    # ------------------------------------------------------------------
    def _map_row(self, row: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]:
        if not row:
            return None
        data = dict(row)
        for key in ("created_at", "updated_at", "expires_at", "last_used_at"):
            value = data.get(key)
            if isinstance(value, datetime):
                data[key] = value.astimezone(timezone.utc).isoformat()
        return data
