from fastapi import FastAPI

from validation.request_validator import IVTRequest
from features.feature_extractor import extract_features
from scoring.weighted_scorer import calculate_rule_score
from scoring.anomaly_model import (
    train_anomaly_model,
    get_anomaly_score,
    normalize_anomaly_score
)
from scoring.blender import calculate_final_risk_score
from explain.ai_explainer import generate_explanation
from features.rtb_adapter import adapt_rtb_request
from scoring.redis_state import (
    check_blocked,
    store_blocked,
    track_suspicious
)


app = FastAPI(title="IVT AI Plugin")


# Train ML model once when the application starts
anomaly_model = train_anomaly_model(
    "data/ml_features.csv"
)


def get_verdict(risk_score):

    if risk_score < 30:
        return "Clean"

    elif risk_score <= 70:
        return "Suspicious"

    else:
        return "Blocked"


@app.post("/validate")
def validate_request(rtb_request: dict):

    request = adapt_rtb_request(rtb_request)

    # Check if IP + UA was previously blocked
    blocked_check = check_blocked(
        request.ip,
        request.user_agent
    )

    if blocked_check["blocked"]:
        verdict = "Blocked"
        risk_score = 100
        rule_score = 100
        ml_risk_score = 100

        contributing_features = [
            "previously_blocked"
        ]

        explanation = generate_explanation(
            risk_score,
            verdict,
            contributing_features
        )

        return {
            "event_id": request.event_id,
            "rule_score": rule_score,
            "ml_score": ml_risk_score,
            "risk_score": risk_score,
            "verdict": verdict,
            "contributing_features": contributing_features,
            "explanation": explanation
        }

    # Extract IVT features
    features = extract_features(
        request,
        "data/blocklist_ips.csv",
        "data/blocklist_user_agents.csv"
    )

    # Rule-based score
    rule_score, contributing_features = calculate_rule_score(
        features
    )

    # ML anomaly score
    ml_score = get_anomaly_score(
        anomaly_model,
        [
            features["request_count"],
            int(features["high_velocity"]),
            int(features["duplicate"]),
            int(features["geo_mismatch"]),
            int(features["bad_ua"])
        ]
    )

    # Convert ML score to 0-100
    ml_risk_score = normalize_anomaly_score(
        ml_score
    )

    # Combine Rule + ML
    final_risk_score = calculate_final_risk_score(
        rule_score,
        ml_risk_score
    )

    # Determine verdict
    verdict = get_verdict(
        final_risk_score
    )

    # Track repeated suspicious traffic
    if verdict == "Suspicious":

        suspicious_result = track_suspicious(
            request.ip,
            request.user_agent
        )

        suspicious_count = suspicious_result["suspicious_count"]

        # Block after repeated suspicious activity
        if suspicious_count >= 3:
            verdict = "Blocked"

            contributing_features.append(
                "repeated_suspicious"
            )

            store_blocked(
                request.ip,
                request.user_agent,
                final_risk_score,
                contributing_features
            )   

    # Store newly blocked traffic in Redis
    if verdict == "Blocked":
        store_blocked(
            request.ip,
            request.user_agent,
            final_risk_score,
            contributing_features
        )

    # AI explanation
    explanation = generate_explanation(
        final_risk_score,
        verdict,
        contributing_features
    )

    return {
        "event_id": request.event_id,
        "rule_score": rule_score,
        "ml_score": ml_risk_score,
        "risk_score": final_risk_score,
        "verdict": verdict,
        "contributing_features": contributing_features,
        "explanation": explanation
    }