import os
import sqlite3
import pyotp
from fastapi import FastAPI, Request, Form, Response, Depends, HTTPException, status
from fastapi.responses import HTMLResponse, RedirectResponse
from fastapi.templating import Jinja2Templates
from fastapi.staticfiles import StaticFiles
import psutil

app = FastAPI()
templates = Jinja2Templates(directory="templates")

TOTP_SECRET_FILE = os.path.join(os.path.dirname(__file__), "totp_secret.txt")

if "ADMIN_TOTP_SECRET" in os.environ:
    TOTP_SECRET = os.environ["ADMIN_TOTP_SECRET"]
else:
    if os.path.exists(TOTP_SECRET_FILE):
        with open(TOTP_SECRET_FILE, "r") as f:
            TOTP_SECRET = f.read().strip()
    else:
        TOTP_SECRET = pyotp.random_base32()
        with open(TOTP_SECRET_FILE, "w") as f:
            f.write(TOTP_SECRET)

print(f"===========================================================")
print(f"ADMIN TOTP SECRET: {TOTP_SECRET}")
print(f"===========================================================")

DB_PATH = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "data", "game_design_matrix.db"))

def get_db():
    conn = sqlite3.connect(DB_PATH)
    conn.row_factory = sqlite3.Row
    return conn

active_sessions = set()

class AuthException(Exception):
    pass

@app.exception_handler(AuthException)
async def auth_exception_handler(request: Request, exc: AuthException):
    return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND)

def require_auth(request: Request):
    session_id = request.cookies.get("session_id")
    if not session_id or session_id not in active_sessions:
        raise AuthException()

@app.get("/login", response_class=HTMLResponse)
async def login_get(request: Request):
    return templates.TemplateResponse(request, "login.html", {"secret": TOTP_SECRET})

@app.post("/login")
async def login_post(request: Request, response: Response, code: str = Form(...)):
    totp = pyotp.TOTP(TOTP_SECRET)
    if totp.verify(code):
        session_id = os.urandom(16).hex()
        active_sessions.add(session_id)
        res = RedirectResponse(url="/", status_code=status.HTTP_302_FOUND)
        res.set_cookie(key="session_id", value=session_id, httponly=True, samesite="lax")
        return res
    return templates.TemplateResponse(request, "login.html", {"error": "Invalid 2FA code"})

@app.get("/logout")
async def logout(request: Request):
    session_id = request.cookies.get("session_id")
    if session_id in active_sessions:
        active_sessions.remove(session_id)
    res = RedirectResponse(url="/login")
    res.delete_cookie("session_id")
    return res

@app.get("/", response_class=HTMLResponse)
async def index(request: Request, _=Depends(require_auth)):
    conn = get_db()
    tables = [row[0] for row in conn.execute("SELECT name FROM sqlite_master WHERE type='table'").fetchall()]
    
    stats = {
        "cpu": psutil.cpu_percent(),
        "memory": psutil.virtual_memory().percent,
    }
    
    return templates.TemplateResponse(request, "dashboard.html", {"tables": tables, "stats": stats})

@app.get("/table/{table_name}", response_class=HTMLResponse)
async def view_table(request: Request, table_name: str, _=Depends(require_auth)):
    conn = get_db()
    try:
        tables = [row[0] for row in conn.execute("SELECT name FROM sqlite_master WHERE type='table'").fetchall()]
        if table_name not in tables:
            return HTMLResponse(content="Invalid table name", status_code=400)

        cursor = conn.execute(f"SELECT * FROM {table_name} LIMIT 100")
        rows = cursor.fetchall()
        columns = [description[0] for description in cursor.description]
        
        pk_info = conn.execute(f"PRAGMA table_info({table_name})").fetchall()
        pk_cols = [col['name'] for col in pk_info if col['pk'] > 0]
        if not pk_cols and columns:
            pk_cols = [columns[0]]

    except Exception as e:
        return HTMLResponse(content=f"Error: {e}", status_code=500)
        
    return templates.TemplateResponse(request, "table.html", {
        "table_name": table_name, 
        "columns": columns, 
        "rows": rows,
        "pk_cols": pk_cols
    })

@app.post("/table/{table_name}/edit")
async def edit_row(request: Request, table_name: str, _=Depends(require_auth)):
    form = await request.form()
    conn = get_db()
    
    tables = [row[0] for row in conn.execute("SELECT name FROM sqlite_master WHERE type='table'").fetchall()]
    if table_name not in tables:
        return HTMLResponse(content="Invalid table name", status_code=400)
        
    pk_info = conn.execute(f"PRAGMA table_info({table_name})").fetchall()
    valid_cols = [col['name'] for col in pk_info]
    pk_cols = [col['name'] for col in pk_info if col['pk'] > 0]
    if not pk_cols and valid_cols:
        pk_cols = [valid_cols[0]]
        
    for pk_col in pk_cols:
        if pk_col not in form:
            return HTMLResponse(content=f"Missing primary key: {pk_col}", status_code=400)
            
    updates = []
    params = []
    pk_params = []
    
    for pk_col in pk_cols:
        pk_params.append(form.get(pk_col))
        
    for key, value in form.items():
        if key not in valid_cols:
            continue
        if key not in pk_cols and key != 'action':
            updates.append(f"{key} = ?")
            params.append(value)
            
    action = form.get('action')
    try:
        if action == 'delete':
            where_clause = " AND ".join([f"{col} = ?" for col in pk_cols])
            conn.execute(f"DELETE FROM {table_name} WHERE {where_clause}", pk_params)
        elif action == 'update':
            if updates:
                where_clause = " AND ".join([f"{col} = ?" for col in pk_cols])
                query = f"UPDATE {table_name} SET {', '.join(updates)} WHERE {where_clause}"
                conn.execute(query, params + pk_params)
        conn.commit()
    except sqlite3.Error as e:
        return HTMLResponse(content=f"Database Error: {e}", status_code=400)
        
    return RedirectResponse(url=f"/table/{table_name}", status_code=status.HTTP_302_FOUND)

@app.post("/table/{table_name}/add")
async def add_row(request: Request, table_name: str, _=Depends(require_auth)):
    form = await request.form()
    conn = get_db()
    
    tables = [row[0] for row in conn.execute("SELECT name FROM sqlite_master WHERE type='table'").fetchall()]
    if table_name not in tables:
        return HTMLResponse(content="Invalid table name", status_code=400)
        
    valid_cols = [col['name'] for col in conn.execute(f"PRAGMA table_info({table_name})").fetchall()]
    
    columns = []
    params = []
    placeholders = []
    
    for key, value in form.items():
        if key in valid_cols:
            columns.append(key)
            params.append(value)
            placeholders.append("?")
            
    if not columns:
        return HTMLResponse(content="No valid columns provided", status_code=400)
        
    try:
        query = f"INSERT INTO {table_name} ({', '.join(columns)}) VALUES ({', '.join(placeholders)})"
        conn.execute(query, params)
        conn.commit()
    except sqlite3.Error as e:
        return HTMLResponse(content=f"Database Error: {e}", status_code=400)
        
    return RedirectResponse(url=f"/table/{table_name}", status_code=status.HTTP_302_FOUND)

if __name__ == "__main__":
    import uvicorn
    uvicorn.run("main:app", host="127.0.0.1", port=9091, reload=True)
