import json
from typing import Optional

import boto3 
import redis


def get_environment() -> str:
    """Detect environment based on AWS account ID."""
    sts = boto3.client("sts")
    account_id = sts.get_caller_identity()["Account"]
    PROD_ACCOUNT_ID = "734185074900"
    STAGING_ACCOUNT_ID = "590183763515"
    if account_id == PROD_ACCOUNT_ID:
        return "PROD"
    elif account_id == STAGING_ACCOUNT_ID:
        return "STAGING"
    else:
        raise Exception("Invalid account ID.") 


def get_secret(secret_name: str, region_name: str) -> str:
    session = boto3.session.Session() 
    client = session.client(service_name='secretsmanager', region_name=region_name)
    response = client.get_secret_value(SecretId=secret_name)
    return response["SecretString"]


def get_redis_secret_name(environment: Optional[str] = None) -> str:
    """Get the appropriate Redis secret name based on environment."""
    if environment is None:
        environment = get_environment()
    
    if environment == "PROD":
        return "studio-api-prod-envs"
    elif environment == "STAGING":
        return "studio-api-service-only-secret"
    else:
        raise ValueError(f"Unknown environment: {environment}")


def get_redis_client(secret_name: Optional[str] = None, region_name: str = "us-east-2") -> redis.Redis:
    """
    Get a Redis client configured with credentials from AWS Secrets Manager.
    
    Args:
        secret_name: Optional secret name. If not provided, will auto-detect based on environment.
        region_name: AWS region name. Defaults to us-east-2.
    
    Returns:
        Configured Redis client.
    """
    if secret_name is None:
        secret_name = get_redis_secret_name()
    
    secret = json.loads(get_secret(secret_name, region_name))
    # Handle different secret key names between environments
    # Staging uses REDIS_RECS_URL, Prod uses RECS_REDIS_URL
    host: str = secret.get("REDIS_RECS_URL") or secret.get("RECS_REDIS_URL")
    if not host:
        raise KeyError("Neither REDIS_RECS_URL nor RECS_REDIS_URL found in secret")
    return redis.Redis.from_url(host, decode_responses=True)


