#services/auth_service.py
import ast
import logging
from typing import Optional
from fastapi import Header, HTTPException, status
from langchain_community.utilities import SQLDatabase
from config.settings import Config

logger = logging.getLogger(__name__)


class AuthService:
    """Enhanced authentication service with better error handling"""

    def __init__(self, db: SQLDatabase):
        self.db = db

    def get_current_company_id(self, api_key: str = Header(..., alias="X-API-Key")) -> Optional[str]:
        """
        Validate API key and return company ID.
        Returns None for super admin, company ID for regular users.

        Args:
            api_key: API key from request header

        Returns:
            Company ID for regular users, None for super admin

        Raises:
            HTTPException: If authentication fails
        """
        if not api_key:
            raise HTTPException(
                status_code=status.HTTP_401_UNAUTHORIZED,
                detail="API key is required"
            )

        # Check for super admin access
        if api_key == Config.SUPER_ADMIN_SECRET_KEY:
            logger.info("Super admin access granted")
            return None

        # Validate company API key
        try:
            # Use parameterized query to prevent SQL injection
            # Format the query directly into the string
            result_str = self.db.run(f"SELECT id FROM company WHERE api_key = '{api_key}'")

            if result_str and result_str != "[]":
                # Parse the result
                parsed_result = ast.literal_eval(result_str)
                if parsed_result and len(parsed_result) > 0 and len(parsed_result[0]) > 0:
                    company_id = str(parsed_result[0][0])
                    logger.info(f"Company access granted for ID: {company_id}")
                    return company_id

            raise HTTPException(
                status_code=status.HTTP_401_UNAUTHORIZED,
                detail="Invalid API key"
            )

        except HTTPException:
            raise
        except Exception as e:
            logger.error(f"Authentication error: {e}")
            raise HTTPException(
                status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
                detail="Authentication service unavailable"
            )
