import datetime
import os

import redis.asyncio as redis
import structlog
from fastapi import FastAPI, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from localization_utils import _t, resolve_preferred_language_from_request
from pydantic import BaseModel

app = FastAPI()

origins = [
    "https://app.suno.ai",
    "https://suno.com",
    "https://www.suno.com",
    "https://staging.suno.ai",
    "http://localhost:3003",
    "http://localhost:3000",
    "https://replica2.suno.com",
    "https://b.suno.fm",
]

app.add_middleware(
    CORSMiddleware,
    allow_origins=origins,
    allow_credentials=True,
    allow_methods=["GET", "POST"],
    allow_headers=["*"],
)

BANNER_UPDATE_SECRET = "Upd4te7h3Sun0Bann3rW3r3D0wn-6e31baa8"

# redis keys
KILLSWITCH_KEY = "killswitch"
KILLSWITCH_MESSAGE_KEY = "killswitch:message"
SCHEDULED_MAINTENANCE_KEY = "scheduled_maintenance"
SCHEDULED_MAINTENANCE_MESSAGE_KEY = "scheduled_maintenance:message"
KILLSWITCH_STATUS_KEY = "killswitch:status"

# banner messages
SYSTEM_UPGRADE_MSG = "System Upgrade in Progress: We're improving our services. Some functions may be temporarily unavailable."
GENERATION_NOT_AVAILABLE_MSG = "Making songs is currently disabled while we upgrade our infrastructure. You can still listen to your favorite tunes."
FREE_GENERATION_NOT_AVAILABLE_MSG = (
    "Generations are currently enabled, but only for Pro and Premier subscribers."
)
SCHEDULED_SYSTEM_MAINTENANCE_MSG = "Scheduled Maintenance: We're currently undergoing maintenance. Some functions may be temporarily unavailable."

redis_client = redis.from_url(os.environ["REDIS_URL"], health_check_interval=20)

logger = structlog.get_logger(__name__)


@app.get("/z")
async def root(request: Request):
    preferred_language = resolve_preferred_language_from_request(request)

    if await redis_client.get("killswitch"):
        message = await redis_client.get("killswitch:message")
        status = await redis_client.get("killswitch:status")
        if message:
            message = message.decode("utf-8")
        if status:
            status = status.decode("utf-8")
            message = _t(f"status.{status}", preferred_language, default=message)
        return {"status": "maintenance", "message": message}

    if await redis_client.get("scheduled_maintenance"):
        message = await redis_client.get("scheduled_maintenance:message")
        status = await redis_client.get("killswitch:status")
        if message:
            message = message.decode("utf-8")
        else:
            message = "Some features may be temporarily unusable while we complete an upcoming system maintenance."

        if status:
            status = status.decode("utf-8")
            message = _t(f"status.{status}", preferred_language, default=message)

        return {"status": "scheduled_maintenance", "message": message}

    return {}


@app.get("/s")
async def get_status_mode(request: Request):
    preferred_language = resolve_preferred_language_from_request(request)
    mode = await redis_client.get("killswitch")
    logger.info("mode", mode=mode)
    if mode:
        message = await redis_client.get("killswitch:message")
        status = await redis_client.get("killswitch:status")
        if message:
            message = message.decode("utf-8")
        if status:
            status = status.decode("utf-8")
            message = _t(f"status.{status}", preferred_language, default=message)

        return {
            "status": "maintenance",
            "mode": "all" if mode == b"all" or mode == "all" else "pro_only",
            "message": message,
        }

    mode = await redis_client.get("scheduled_maintenance")
    if mode:
        message = await redis_client.get("scheduled_maintenance:message")
        status = await redis_client.get("killswitch:status")
        if message:
            message = message.decode("utf-8")
        else:
            message = "Some features may be temporarily unusable while we complete an upcoming system maintenance."

        if status:
            status = status.decode("utf-8")
            message = _t(f"status.{status}", preferred_language, default=message)
        return {
            "status": "scheduled_maintenance",
            "mode": "all" if mode == b"all" or mode == "all" else "pro_only",
            "message": message,
        }

    return {}


async def display_system_upgrade_banner(msg: str | None = None):
    msg = msg or SYSTEM_UPGRADE_MSG
    await redis_client.delete(SCHEDULED_MAINTENANCE_KEY)
    await redis_client.set(KILLSWITCH_KEY, "all")
    await redis_client.set(KILLSWITCH_MESSAGE_KEY, msg)
    await redis_client.set(KILLSWITCH_STATUS_KEY, "upgrade")
    logger.info("System upgrade banner is set with message", msg=msg)
    return msg


async def display_scheduled_maintenance_banner(msg: str | None = None):
    msg = msg or SCHEDULED_SYSTEM_MAINTENANCE_MSG
    await redis_client.delete(KILLSWITCH_KEY)
    await redis_client.set(SCHEDULED_MAINTENANCE_KEY, "all")
    await redis_client.set(SCHEDULED_MAINTENANCE_MESSAGE_KEY, msg)
    await redis_client.set(KILLSWITCH_STATUS_KEY, "schedule")
    logger.info("Scheduled maintenance banner is set with message", msg=msg)
    return msg


async def remove_banner():
    await redis_client.delete(KILLSWITCH_KEY)
    await redis_client.delete(SCHEDULED_MAINTENANCE_KEY)
    await redis_client.delete(KILLSWITCH_STATUS_KEY)
    logger.info("Banner is removed")
    return "<Banner removed>"


async def display_banner_and_block_free_gens(msg: str | None = None):
    msg = msg or FREE_GENERATION_NOT_AVAILABLE_MSG
    await redis_client.set(KILLSWITCH_KEY, "true")
    await redis_client.set(KILLSWITCH_MESSAGE_KEY, msg)
    await redis_client.set(KILLSWITCH_STATUS_KEY, "block_free")
    logger.info("Free generation banner is set with message", msg=msg)
    return msg


async def display_banner_and_block_all_gens(msg: str | None = None):
    msg = msg or GENERATION_NOT_AVAILABLE_MSG
    await redis_client.set(KILLSWITCH_KEY, "all")
    await redis_client.set(KILLSWITCH_MESSAGE_KEY, msg)
    await redis_client.set(KILLSWITCH_STATUS_KEY, "block_all")
    logger.info("Generation banner is set with message", msg=msg)
    return msg


async def display_temporarily_down_banner(msg: str | None = None):
    msg = msg or GENERATION_NOT_AVAILABLE_MSG
    await redis_client.set(KILLSWITCH_KEY, "all")
    await redis_client.set(KILLSWITCH_MESSAGE_KEY, msg)
    await redis_client.set(KILLSWITCH_STATUS_KEY, "temporarily_down")
    logger.info("Generation banner is set with message", msg=msg)
    return msg


class UpdateStatusRequest(BaseModel):
    status: str
    message: str
    secret: str


@app.post("/update")
async def update_status(request: UpdateStatusRequest):
    if request.secret != BANNER_UPDATE_SECRET:
        raise HTTPException(status_code=401, detail="Unauthorized")

    status = request.status
    msg = request.message

    if status == "upgrade":
        set_message = await display_system_upgrade_banner(msg)
    elif status == "schedule":
        set_message = await display_scheduled_maintenance_banner(msg)
    elif status == "block_free":
        set_message = await display_banner_and_block_free_gens(msg)
    elif status == "block_all":
        set_message = await display_banner_and_block_all_gens(msg)
    elif status == "temporarily_down":
        set_message = await display_temporarily_down_banner(msg)
    elif status == "clear":
        set_message = await remove_banner()
    else:
        raise HTTPException(status_code=400, detail="Invalid status")

    return {
        "status": "OK",
        "message": set_message,
    }


# This is for interview purposes only, it alters the status every 30 seconds
@app.get("/t")
async def get_status_time_based():
    current_time = datetime.datetime.now()
    if current_time.second % 60 < 30:  # 0-29
        return {"status": "maintenance", "message": "Maintenance in progress."}
    else:  # 30-59
        return {"status": "OK"}
