from fastapi import FastAPI, Request, Form, File, UploadFile, BackgroundTasks
from fastapi.responses import HTMLResponse, RedirectResponse, JSONResponse
from fastapi.templating import Jinja2Templates
from starlette.middleware.sessions import SessionMiddleware
from auth import (
    AdminAuthMiddleware,
    SESSION_SECRET,
    change_admin_password,
    check_admin_login,
    is_authenticated,
    safe_next_path,
    verify_admin,
)
import re
import os
import shutil
import threading
from typing import List, Dict, Optional
from pathlib import Path
from urllib.parse import urlencode
from uuid import uuid4
from sqlalchemy import Column, Integer, String, Text, Boolean, Float, DateTime, create_engine, text, inspect
from sqlalchemy.orm import declarative_base, sessionmaker
from sqlalchemy.sql import func
from spam_training import classify_new_imap_emails, learn_from_feedback, process_uploaded_spam_csv


DATABASE_URL = "postgresql+asyncpg://postgres:p4python@postgres.cqf3qroocfwm.us-east-1.rds.amazonaws.com:5432/email_filter"
SQL_ECHO = os.getenv("SQL_ECHO", "false").lower() == "true"
AUTO_CREATE_TABLES = os.getenv("AUTO_CREATE_TABLES", "false").lower() == "true"
ENSURE_FEEDBACK_COLUMN = os.getenv("ENSURE_FEEDBACK_COLUMN", "false").lower() == "true"

Base = declarative_base()

engine = create_engine(
    DATABASE_URL.replace("+asyncpg", ""),
    echo=SQL_ECHO,
    pool_pre_ping=True
)

SessionLocal = sessionmaker(bind=engine)

def get_all_emails_from_db():
    db = SessionLocal()
    try:
        emails = (
            db.query(EmailSpam)
            .filter(EmailSpam.imap_uid.isnot(None))
            .order_by(EmailSpam.received_at.desc())
            .all()
        )
        return emails or []   # always returns a list
    except Exception as e:
        print("DB error while fetching emails:", e)
        return []             # fail-safe
    finally:
        db.close()



app = FastAPI()
app.add_middleware(AdminAuthMiddleware)
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET, session_cookie="manjet_session")

# Get the directory where this file is located
BASE_DIR = Path(__file__).parent
TEMPLATES_DIR = BASE_DIR / "templates"
UPLOADS_DIR = BASE_DIR / "uploads"
MODEL_DIR = BASE_DIR / "models"
templates = Jinja2Templates(directory=str(TEMPLATES_DIR))

TRAINING_STATUS_LOCK = threading.Lock()
TRAINING_STATUS = {
    "state": "idle",
    "message": "",
    "filename": "",
}


def get_training_status():
    with TRAINING_STATUS_LOCK:
        status = TRAINING_STATUS.copy()
    if "Refreshed" in (status.get("message") or ""):
        status.update({"state": "idle", "message": "", "filename": ""})
    return status


def set_training_status(state: str, message: str = "", filename: str = ""):
    with TRAINING_STATUS_LOCK:
        TRAINING_STATUS.update({
            "state": state,
            "message": message,
            "filename": filename,
        })


def filter_results_exist() -> bool:
    db = SessionLocal()
    try:
        return db.execute(
            text(
                """
                SELECT 1
                FROM email_filter
                WHERE imap_uid IS NOT NULL
                  AND model_accuracy IS NOT NULL
                LIMIT 1
                """
            )
        ).first() is not None
    finally:
        db.close()


def run_feedback_learning_in_background(email_id: int, classification: str):
    try:
        learn_from_feedback(
            email_id=email_id,
            feedback_label=classification,
            session_factory=SessionLocal,
            model_dir=MODEL_DIR,
        )
        print("Updated model from feedback for email_id:", email_id)
    except Exception as e:
        print("Background model learning failed for email_id", email_id, ":", e)


def save_classification_feedback(
    email_id: int,
    classification: str,
    feedback: str,
) -> bool:
    db = SessionLocal()
    try:
        row = db.execute(
            text(
                """
                UPDATE email_filter
                SET predicted_label = :classification,
                    feedback_label = :classification,
                    feedback = :feedback,
                    is_reviewed = TRUE
                WHERE id = :email_id
                  AND imap_uid IS NOT NULL
                RETURNING id
                """
            ),
            {
                "classification": classification,
                "feedback": feedback,
                "email_id": email_id,
            },
        ).first()
        if not row:
            db.rollback()
            return False
        db.commit()
        return True
    except Exception as e:
        db.rollback()
        print("Full feedback save failed, trying minimal update:", e)
        try:
            row = db.execute(
                text(
                    """
                    UPDATE email_filter
                    SET predicted_label = :classification,
                        is_reviewed = TRUE
                    WHERE id = :email_id
                      AND imap_uid IS NOT NULL
                    RETURNING id
                    """
                ),
                {
                    "classification": classification,
                    "email_id": email_id,
                },
            ).first()
            if not row:
                db.rollback()
                return False
            db.commit()
            return True
        except Exception as fallback_error:
            db.rollback()
            print("Minimal feedback save failed:", fallback_error)
            return False
    finally:
        db.close()


def train_uploaded_dataset_in_background(csv_path: Path, filename: str):
    try:
        set_training_status("running", f"Training model from {filename}...", filename)
        result = process_uploaded_spam_csv(
            csv_path=csv_path,
            session_factory=SessionLocal,
            model_dir=MODEL_DIR,
        )
        message = (
            f"Training complete for {filename}. Trained on {result.training_rows} usable rows "
            f"from {result.total_rows} CSV rows and reclassified {result.reclassified_rows} existing emails. "
            f"Accuracy: {result.accuracy * 100:.1f}%. "
            f"Spam: {result.spam_count}, Ham: {result.ham_count}, Undecided: {result.undecided_count}."
        )
        set_training_status("completed", message, filename)
    except Exception as e:
        print("CSV upload/training failed:", e)
        set_training_status("error", str(e), filename)


class EmailSpam(Base):
    __tablename__ = "email_filter"
    
    id = Column(Integer, primary_key=True, index=True)
    sender_email = Column(String, nullable=False)
    subject = Column(String, nullable=True)
    body = Column(Text, nullable=True)

    spam_probability = Column(Float, nullable=True)
    predicted_label = Column(String, nullable=True)
    feedback_label = Column(String, nullable=True)
    feedback = Column(Text, nullable=True)
    is_reviewed = Column(Boolean, default=False)

    email_summary = Column(Text, nullable=True)      
    model_accuracy = Column(Float, nullable=True)  
    imap_uid = Column(String, nullable=True)
    received_at = Column(DateTime(timezone=True), server_default=func.now())

# ----------------------------
# Automated Analysis Functions
# ----------------------------
def analyze_email_subject(subject: str) -> List[str]:
    """Check email subject for spam-like patterns"""
    spam_patterns = []
    subject_lower = subject.lower()
    
    # Common spam indicators
    spam_keywords = [
        r'\b(win|winner|won|prize|free|urgent|limited time|act now|click here)\b',
        r'\$[\d,]+',
        r'!!!+',
        r'\b(viagra|cialis|pharmacy|pills)\b',
        r'\b(guaranteed|risk-free|no obligation)\b',
        r'\b(click|download|claim|verify)\b.*\b(now|immediately|today)\b'
    ]
    
    for pattern in spam_keywords:
        if re.search(pattern, subject_lower, re.IGNORECASE):
            spam_patterns.append(f"Subject contains suspicious pattern: '{pattern}'")
    
    return spam_patterns

# ----------------------------
# Admin login
# ----------------------------
@app.get("/admin/login", response_class=HTMLResponse)
async def admin_login_page(
    request: Request,
    next: Optional[str] = None,
    error: Optional[str] = None,
    username: Optional[str] = None,
):
    if is_authenticated(request):
        return RedirectResponse(safe_next_path(next), status_code=303)
    return templates.TemplateResponse(request,
        "admin_login.html",
        {
            "request": request,
            "error": error,
            "next": next,
            "username": username or "",
            "password": "",
            "error_field": None,
        },
    )


@app.post("/admin/login")
async def admin_login_submit(
    request: Request,
    username: str = Form(...),
    password: str = Form(...),
    next: Optional[str] = Form(None),
):
    success, error_code = check_admin_login(username, password)
    if success:
        request.session["is_admin"] = True
        return RedirectResponse(safe_next_path(next), status_code=303)

    if error_code == "wrong_username":
        error_message = "Wrong username. Please check your email address and try again."
    else:
        error_message = "Wrong password. Please try again."

    return templates.TemplateResponse(request,
        "admin_login.html",
        {
            "request": request,
            "error": error_message,
            "error_field": error_code,
            "next": next,
            "username": username,
            "password": password,
        },
        status_code=401,
    )


@app.post("/admin/logout")
async def admin_logout(request: Request):
    request.session.clear()
    return RedirectResponse("/", status_code=303)


@app.get("/admin/change-password", response_class=HTMLResponse)
async def change_password_page(
    request: Request,
    message: Optional[str] = None,
    message_type: Optional[str] = None,
):
    return templates.TemplateResponse(request,
        "change_password.html",
        {
            "request": request,
            "message": message,
            "message_type": message_type or "error",
        },
    )


@app.post("/admin/change-password")
async def change_password_submit(
    request: Request,
    current_password: str = Form(...),
    new_password: str = Form(...),
    confirm_password: str = Form(...),
):
    ok, result_message = change_admin_password(
        current_password,
        new_password,
        confirm_password,
    )
    message_type = "success" if ok else "error"
    return templates.TemplateResponse(request,
        "change_password.html",
        {
            "request": request,
            "message": result_message,
            "message_type": message_type,
        },
        status_code=200 if ok else 400,
    )


# ----------------------------
# Landing page
# ----------------------------
@app.get("/", response_class=HTMLResponse)
async def landing_page(request: Request):
    if is_authenticated(request):
        return RedirectResponse(url="/dashboard/spam", status_code=303)
    return templates.TemplateResponse(request,
        "landing.html",
        {"request": request},
    )

# ----------------------------
# LISTING PAGE
# ----------------------------

LABEL_FILTERS = {"spam", "ham", "undecided"}


@app.get("/dashboard/spam", response_class=HTMLResponse)
async def spam_dashboard(
    request: Request,
    page: int = 1,
    filter: Optional[str] = None,
    upload_status: Optional[str] = None,
    upload_message: Optional[str] = None,
):
    if upload_message and "Refreshed" in upload_message:
        upload_status = None
        upload_message = None

    PER_PAGE = 20
    offset = (page - 1) * PER_PAGE

    def normalize_label(label):
        return label.strip().lower() if label else None

    try:
        classify_new_imap_emails(SessionLocal, MODEL_DIR)
    except Exception as e:
        print("Could not classify newly fetched emails:", e)

    # 1. Fetch all IMAP emails once. These rows are display data and must not be deleted by uploads.
    all_emails = get_all_emails_from_db()
    total_emails = len(all_emails)

    db = SessionLocal()
    try:
        latest_accuracy_row = (
            db.query(EmailSpam.model_accuracy)
            .filter(EmailSpam.imap_uid.isnot(None))
            .filter(EmailSpam.model_accuracy.isnot(None))
            .order_by(EmailSpam.id.desc())
            .first()
        )
    finally:
        db.close()

    if latest_accuracy_row:
        model_accuracy_percent = round(latest_accuracy_row.model_accuracy * 100, 1)
        filter_results_available = True
    else:
        model_accuracy_percent = None
        filter_results_available = False

    # 2. Calculate dataset-based stats only after a user-uploaded dataset has run.
    if filter_results_available:
        all_labels = [normalize_label(e.predicted_label) for e in all_emails]
        spam_count = sum(1 for l in all_labels if l == "spam")
        ham_count = sum(1 for l in all_labels if l == "ham")
        undecided_count = sum(1 for l in all_labels if l == "undecided")
        not_classified_count = sum(1 for l in all_labels if l is None)
        classified_count = total_emails - not_classified_count
    else:
        spam_count = 0
        ham_count = 0
        undecided_count = 0
        classified_count = 0

    active_filter = (filter or "").strip().lower()
    if active_filter not in LABEL_FILTERS:
        active_filter = None
    if active_filter and not filter_results_available:
        active_filter = None

    display_emails = all_emails
    if active_filter:
        display_emails = [
            email
            for email in all_emails
            if normalize_label(email.predicted_label) == active_filter
        ]

    # 3. Paginate emails for display.
    db_emails = display_emails[offset:offset + PER_PAGE]

    enhanced_emails = []
    for email in db_emails:
        display_label = normalize_label(email.predicted_label) if filter_results_available else None
        spam_prob = email.spam_probability if filter_results_available else None
        ml_score = round(spam_prob * 100, 1) if spam_prob is not None else None

        enhanced_emails.append({
            "email_id": email.id,
            "sender": email.sender_email,
            "subject": email.subject,
            "summary": email.body,
            "label": display_label if display_label else "unfiltered",
            "ml_score": ml_score,
            "feedback": email.feedback or ""
        })

    # 4. Total pages for the current filtered view.
    filtered_total = len(display_emails)
    total_pages = (filtered_total + PER_PAGE - 1) // PER_PAGE if filtered_total else 1

    stats = {
        "total": total_emails,
        "spam": spam_count,
        "ham": ham_count,
        "undecided": undecided_count,
        "classified": classified_count
    }

    return templates.TemplateResponse(request,
        "spam_dashboard.html",
        {
            "request": request,
            "emails": enhanced_emails,
            "stats": stats,
            "page": page,
            "total_pages": total_pages,
            "model_accuracy": model_accuracy_percent,
            "filter_results_available": filter_results_available,
            "upload_status": upload_status,
            "upload_message": upload_message,
            "training_status": get_training_status(),
            "active_filter": active_filter,
            "filtered_total": filtered_total,
        }
    )


# ----------------------------
# ADD FORM PAGE
# ----------------------------
@app.get("/dashboard/spam/add", response_class=HTMLResponse)
async def spam_add_page(request: Request):
    return templates.TemplateResponse(request,
        "spam_add.html",
        {"request": request}
    )


@app.get("/dashboard/spam/upload-status")
async def spam_upload_status():
    return JSONResponse(get_training_status())


@app.post("/dashboard/spam/clear-analysis")
async def clear_spam_analysis():
    current_status = get_training_status()
    if current_status.get("state") == "running":
        query = urlencode({
            "upload_status": "error",
            "upload_message": "Training is still running. Please wait before clearing analysis.",
        })
        return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)

    db = SessionLocal()
    try:
        db.execute(
            text(
                """
                UPDATE email_filter
                SET spam_probability = NULL,
                    predicted_label = NULL,
                    model_accuracy = NULL
                WHERE imap_uid IS NOT NULL
                """
            )
        )
        db.commit()
    except Exception as e:
        db.rollback()
        print("Could not clear analysis:", e)
        query = urlencode({
            "upload_status": "error",
            "upload_message": "Could not clear analysis. Please try again.",
        })
        return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)
    finally:
        db.close()

    for model_file in (MODEL_DIR / "final_model.pkl", MODEL_DIR / "final_vectorizer.pkl"):
        try:
            model_file.unlink(missing_ok=True)
        except Exception as e:
            print("Could not remove model file:", model_file, e)

    set_training_status("idle", "", "")
    query = urlencode({
        "upload_status": "success",
        "upload_message": "Analysis cleared. Emails are still saved; upload a dataset to analyze them again.",
    })
    return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)


@app.post("/dashboard/spam/upload")
async def upload_spam_dataset(
    background_tasks: BackgroundTasks,
    dataset: UploadFile = File(...),
):
    current_status = get_training_status()
    if current_status.get("state") == "running":
        query = urlencode({
            "upload_status": "error",
            "upload_message": "A dataset is already training. Please wait for it to finish.",
        })
        return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)

    if filter_results_exist():
        query = urlencode({
            "upload_status": "error",
            "upload_message": "Analysis already exists. Clear analysis before training with another dataset.",
        })
        return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)

    filename = dataset.filename or ""
    if not filename.lower().endswith(".csv"):
        query = urlencode({
            "upload_status": "error",
            "upload_message": "Please upload a CSV file.",
        })
        return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)

    UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
    safe_name = re.sub(r"[^a-zA-Z0-9_.-]", "_", Path(filename).name)
    upload_path = UPLOADS_DIR / f"{uuid4().hex}_{safe_name}"

    try:
        with upload_path.open("wb") as buffer:
            shutil.copyfileobj(dataset.file, buffer)
        set_training_status("running", f"Training model from {filename}...", filename)
        background_tasks.add_task(train_uploaded_dataset_in_background, upload_path, filename)
        query = urlencode({
            "upload_status": "info",
            "upload_message": "Upload received. Training is running in the background and existing emails will be kept.",
        })
    except Exception as e:
        print("CSV upload failed:", e)
        set_training_status("error", str(e), filename)
        query = urlencode({
            "upload_status": "error",
            "upload_message": str(e),
        })
    finally:
        dataset.file.close()

    return RedirectResponse(f"/dashboard/spam?{query}", status_code=303)

# ----------------------------
# FORM SUBMIT
# ----------------------------

# @app.post("/dashboard/spam/add")
# async def spam_add_submit(
#     request: Request,
#     subject: str = Form(...),
#     sender: str = Form(...),
#     summary: str = Form(...)
# ):
#     analysis = generate_automated_analysis(subject, sender, summary)

#     db = SessionLocal()
#     try:
#         new_email = EmailSpam(
#             sender_email=sender,
#             subject=subject,
#             body=summary,
#             spam_probability=None,
#             predicted_label="Not Classified",
#             feedback_label=None,
#             is_reviewed=False
#         )
#         db.add(new_email)
#         db.commit()
#     finally:
#         db.close()

#     return RedirectResponse(url="/dashboard/spam", status_code=302)


# ----------------------------
# CLASSIFICATION UPDATE
# ----------------------------
# allowed values: "spam", "ham", "undecided"

@app.post("/dashboard/spam/classify")
async def classify_email(
    background_tasks: BackgroundTasks,
    email_id: int = Form(...),
    classification: Optional[str] = Form(None),
    feedback: Optional[str] = Form(None),
):
    try:
        classification = (classification or "").strip().lower()
        feedback = (feedback or "").strip()

        if classification not in ["spam", "ham", "undecided"]:
            return RedirectResponse(f"/dashboard/spam/email/{email_id}?form_error=1", status_code=303)
        if not feedback:
            return RedirectResponse(f"/dashboard/spam/email/{email_id}?form_error=1", status_code=303)

        print(
            "POST /dashboard/spam/classify",
            {"email_id": email_id, "classification": classification, "feedback_len": len(feedback)},
        )

        if not save_classification_feedback(email_id, classification, feedback):
            print("No row found or could not save classification for email_id:", email_id)
            return RedirectResponse(f"/dashboard/spam/email/{email_id}?db_error=1", status_code=303)

        print("Saved classification/feedback for email_id:", email_id)

        if classification in {"spam", "ham"}:
            background_tasks.add_task(
                run_feedback_learning_in_background,
                email_id,
                classification,
            )

        return RedirectResponse(
            f"/dashboard/spam/email/{email_id}?feedback_saved=1",
            status_code=303,
        )
    except Exception as e:
        print("Unexpected classify handler error:", e)
        return RedirectResponse(f"/dashboard/spam/email/{email_id}?db_error=1", status_code=303)



@app.get("/dashboard/spam/email/{email_id}", response_class=HTMLResponse)
async def email_details(request: Request, email_id: int):
    db = SessionLocal()
    try:
        email = (
            db.query(EmailSpam)
            .filter(EmailSpam.id == email_id)
            .filter(EmailSpam.imap_uid.isnot(None))
            .first()
        )
    finally:
        db.close()

    if not email:
        return RedirectResponse("/dashboard/spam", status_code=303)

    filter_results_available = email.model_accuracy is not None
    predicted_label = (email.predicted_label or "").strip().lower()

    # Detect if content is HTML
    body = email.body or ""
    is_html = bool(re.search(r'<[a-z][\s\S]*>', body, re.IGNORECASE))
    
    # Calculate ML score from spam_probability
    spam_prob = email.spam_probability if filter_results_available else None
    ml_score = round(spam_prob * 100, 1) if spam_prob is not None else None
    
    enhanced_email = {
        "id": email.id,
        "subject": email.subject or "No Subject",
        "sender": email.sender_email,
        "body": body,
        "summary": body,  # For template compatibility
        "is_html": is_html,
        "ml_score": ml_score,
        "label": predicted_label if predicted_label else "unfiltered",
        "feedback": email.feedback or "",
        "spam_probability": spam_prob,
        "received_at": email.received_at,
        "filter_results_available": filter_results_available,
        "show_classify_form": predicted_label == "undecided",
        "is_reviewed": bool(email.is_reviewed),
    }

    return templates.TemplateResponse(request,
        "email_details.html",
        {
            "request": request,
            "email": enhanced_email
        }
    )


if AUTO_CREATE_TABLES:
    Base.metadata.create_all(bind=engine)

FEEDBACK_COLUMN_EXISTS = False

def ensure_feedback_column():
    global FEEDBACK_COLUMN_EXISTS
    db = SessionLocal()
    try:
        db.execute(text("SET LOCAL lock_timeout = '5s'"))
        db.execute(text("ALTER TABLE email_filter ADD COLUMN IF NOT EXISTS feedback TEXT"))
        db.commit()
    except Exception as e:
        db.rollback()
        print("Could not ensure feedback column:", e)
    finally:
        db.close()

    # Verify existence (even if ALTER failed silently due to permissions, etc.)
    try:
        insp = inspect(engine)
        cols = [c.get("name") for c in insp.get_columns("email_filter")]
        FEEDBACK_COLUMN_EXISTS = "feedback" in cols
    except Exception as e:
        print("Could not verify feedback column existence:", e)
        FEEDBACK_COLUMN_EXISTS = False

if ENSURE_FEEDBACK_COLUMN:
    ensure_feedback_column()
else:
    FEEDBACK_COLUMN_EXISTS = True

if __name__ == "__main__":
    import uvicorn

    uvicorn.run(
        "main:app",          
        host="0.0.0.0",
        port=6202,
        reload=True,
        reload_dirs=[str(BASE_DIR)], 
        reload_excludes=["venv"]     
    )
