from __future__ import annotations

"""Subscription manager responsible for tier resolution."""

import os
from dataclasses import dataclass, field
from typing import Optional

from .models import SubscriptionFeatures, SubscriptionTier, get_features_for_tier


@dataclass
class SubscriptionManager:
    """Represent a user's subscription and expose helper methods."""

    user_id: str
    tier: SubscriptionTier = field(default_factory=lambda: SubscriptionTier.FREE)
    _cached_features: Optional[SubscriptionFeatures] = field(default=None, init=False, repr=False)
    _override_features: Optional[SubscriptionFeatures] = field(default=None, init=False, repr=False)
    
    def __post_init__(self):
        """Initialise tier en respectant la valeur demandée tout en appliquant un upgrade depuis la base si nécessaire."""
        if not isinstance(self.tier, SubscriptionTier):
            try:
                self.tier = self._coerce_tier(self.tier)
            except ValueError:
                self.tier = SubscriptionTier.FREE

        db_tier = self._resolve_user_tier()
        if db_tier and db_tier != SubscriptionTier.FREE:
            self.tier = db_tier

    def set_tier(self, tier: str | SubscriptionTier) -> None:
        """Update the current tier."""
        tier_enum = self._coerce_tier(tier)
        if tier_enum != self.tier:
            self.tier = tier_enum
            self._cached_features = None
            self._override_features = None

    def get_features(self) -> SubscriptionFeatures:
        """Return the feature set for the current tier."""
        if self._override_features is not None:
            return self._override_features
        if self._cached_features is None:
            self._cached_features = get_features_for_tier(self.tier)
        return self._cached_features

    def get_domain_limit(self) -> int:
        return self.get_features().domain_limit

    def get_user_limit(self) -> int:
        return self.get_features().max_users

    def can_use_invasive_tests(self) -> bool:
        return self.get_features().allow_invasive_tests

    def can_export_html(self) -> bool:
        return self.get_features().allow_html_export

    def can_manage_users(self) -> bool:
        return self.get_features().allow_multi_user

    def can_use_api(self) -> bool:
        return self.get_features().allow_api_access
    
    def can_scan_source(self) -> bool:
        return self.get_features().allow_source_scan

    def apply_features(self, features: SubscriptionFeatures) -> None:
        """Override the current feature set using licence-derived values."""
        # Ne pas écraser le tier si l'utilisateur a un tier plus élevé dans SQLite
        db_tier = self._resolve_user_tier()
        if db_tier and db_tier != SubscriptionTier.FREE:
            # Garder le tier de la base SQLite et merger les features
            merged_features = self.get_features()
            self._override_features = merged_features
        else:
            # Utiliser les features de la licence
            self.tier = features.tier
            self._override_features = features
        self._cached_features = None

    def _resolve_user_tier(self) -> Optional[SubscriptionTier]:
        """Résoudre le tier de l'utilisateur depuis la base SQLite."""
        from pathlib import Path
        import sqlite3
        backend_mode = os.getenv("WEB_SENTINEL_SUBSCRIPTION_BACKEND", "auto").lower()
        disable_local = os.getenv("WEB_SENTINEL_DISABLE_LOCAL_DB", "").lower() in {"1", "true", "yes"}
        if disable_local or backend_mode in {"postgresql", "http"}:
            return None

        db_path = Path.home() / ".web-sentinel" / "auth.db"
        if not self.user_id or not db_path.exists():
            return None
        
        try:
            with sqlite3.connect(db_path) as conn:
                cursor = conn.execute(
                    "SELECT tier FROM user_subscriptions WHERE lower(email) = ?",
                    (self.user_id.lower(),),
                )
                row = cursor.fetchone()
        except sqlite3.Error:
            return None
        
        if not row or not row[0]:
            return None
        
        try:
            return SubscriptionTier(row[0].lower())
        except ValueError:
            return None

    @staticmethod
    def _coerce_tier(tier: str | SubscriptionTier) -> SubscriptionTier:
        if isinstance(tier, SubscriptionTier):
            return tier
        try:
            return SubscriptionTier(tier.lower())
        except ValueError as exc:
            raise ValueError(f"Unknown subscription tier: {tier}") from exc
