"""
evaluate_retrieval.py
----------------------
Μετράει την ποιότητα του retrieval pipeline (MMR -> Πορτιέρης -> Reranker)
ΧΩΡΙΣ να καλεί το Gemini -- γρήγορο, δωρεάν, και απαντάει στο ερώτημα:
"Βρίσκει το σύστημα τα σωστά κομμάτια πληροφορίας;"

Χρήση:
    python3 evaluate_retrieval.py

Χρειάζεται:
    - eval_questions_template.csv (ή δικό σου αρχείο -- άλλαξε το EVAL_FILE παρακάτω)
      στον ίδιο φάκελο με τα PDF/agreements.csv που ήδη χρησιμοποιεί το main.py
    - Το ίδιο venv/dependencies με το main.py (fitz, langchain, faiss-cpu,
      sentence-transformers, HuggingFace embeddings)

Παράγει:
    - evaluation_results.csv  (μία γραμμή ανά ερώτηση, με Precision/Recall/RR)
    - Εκτύπωση στο τερματικό με τους μέσους όρους (Precision@k, Recall@k, MRR)
"""

import os
import csv
import sys
import fitz  # PyMuPDF
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_community.vectorstores import FAISS
from langchain_community.embeddings import HuggingFaceEmbeddings
from sentence_transformers import CrossEncoder

# ============ ΡΥΘΜΙΣΕΙΣ (ίδιες με το main.py) ============
EVAL_FILE = "eval_questions_v1.csv"   
TOP_K_FINAL = 8                             # ίδιο με το top_docs cutoff στο main.py
MMR_K = 40
MMR_FETCH_K = 80
CHUNK_SIZE = 1500
CHUNK_OVERLAP = 300


def build_knowledge_base():
    """Ίδια λογική με το setup_knowledge_base() του main.py, αλλά επιστρέφει
    και τη λίστα όλων των chunks (χρειάζεται για τον σωστό υπολογισμό του Recall)."""
    texts = []
    metadatas = []

    files = [f for f in os.listdir('.') if f.endswith('.pdf')]
    for file in files:
        try:
            doc = fitz.open(file)
            file_text = ""
            for page in doc:
                file_text += page.get_text()
            file_text = file_text.strip()
            if file_text:
                texts.append(file_text)
                metadatas.append({"filename": file})
        except Exception as e:
            print(f"[WARN] Πρόβλημα στην ανάγνωση του {file}: {e}")

    if os.path.exists('agreements.csv'):
        try:
            with open('agreements.csv', mode='r', encoding='utf-8-sig') as csvfile:
                reader = csv.DictReader(csvfile, delimiter=';')
                for row in reader:
                    row_text = (
                        f"[ΠΙΝΑΚΑΣ - ΔΙΜΕΡΕΙΣ ΣΥΜΦΩΝΙΕΣ / ΣΥΝΕΡΓΑΖΟΜΕΝΑ ΠΑΝΕΠΙΣΤΗΜΙΑ ΕΚΠΑ] "
                        f"Αυτή η συμφωνία αποτελεί προαπαιτούμενο για κινητικότητα προσωπικού για διδασκαλία (STA) "
                        f"καθώς και για κινητικότητα φοιτητών για σπουδές (SMS / Long-term / Short-term). "
                        f"Το τμήμα {row.get('Τμήμα', '')} συνεργάζεται με το ίδρυμα {row.get('Συνεργαζόμενα Πανεπιστήμια', '')} "
                        f"(Κωδικός: {row.get('Κωδικός  Συνεργαζόμενων Πανεπιστημίων', '')}) στη χώρα {row.get('Χώρα', '')}. "
                        f"Υπεύθυνος Καθηγητής: ο/η {row.get('Υπεύθυνος Καθηγητής', '')}. "
                        f"Αριθμός διαθέσιμων θέσεων φοιτητών: {row.get('Αριθμός φοιτητών σύμφωνα με τη διμερή συμφ.', '')} "
                        f"για συνολικά {row.get('Σύνολο φοιτητομηνών', '')} φοιτητομήνες. "
                        f"Κύκλος σπουδών: {row.get('Κύκλος σπουδών', '')}. "
                        f"Γλώσσα διδασκαλίας: {row.get('Γλώσσα διδασκαλίας', '')} (Δεύτερη Γλώσσα: {row.get('Δεύτερη Γλώσσα', '')}). "
                        f"Τομέας Σπουδών: {row.get('Κωδικός Τομέα Σπουδών', '')}."
                    )
                    texts.append(row_text)
                    metadatas.append({"filename": "AGREEMENTS_CSV_STA_SMS_LONG"})
        except Exception as e:
            print(f"[WARN] Πρόβλημα στην ανάγνωση του agreements.csv: {e}")

    if not texts:
        print("[ΣΦΑΛΜΑ] Δεν βρέθηκαν PDF ή CSV αρχεία. Σταματάω.")
        sys.exit(1)

    splitter = RecursiveCharacterTextSplitter(chunk_size=CHUNK_SIZE, chunk_overlap=CHUNK_OVERLAP)
    chunks = splitter.create_documents(texts, metadatas=metadatas)

    print(f"[INFO] Φόρτωση embeddings model (BAAI/bge-m3)...")
    embeddings = HuggingFaceEmbeddings(model_name="BAAI/bge-m3")
    vector_store = FAISS.from_documents(chunks, embedding=embeddings)

    return vector_store, chunks


def apply_portiere_filter(docs, category):
    """Ίδια λογική "πορτιέρη" με το main.py -- φιλτράρει βάσει ενεργής κατηγορίας."""
    if not category:
        return docs  # καμία κατηγορία -> κανένα φιλτράρισμα

    filtered = []
    for doc in docs:
        filename = doc.metadata.get("filename", "").upper()
        is_stt = "STT" in filename
        is_sta = "STA" in filename
        is_long = "LONG" in filename or "SMS_OUT_LONG" in filename
        is_short = any(x in filename for x in ["SHORT", "BIP", "ΒΡΑΧΥ"])

        if category == "STT" and (is_sta or is_long or is_short):
            continue
        if category == "STA" and (is_stt or is_long or is_short):
            continue
        if category == "SMS_LONG" and (is_stt or is_sta or is_short):
            continue
        if category == "SMS_SHORT" and (is_stt or is_sta or is_long):
            continue

        filtered.append(doc)

    return filtered if filtered else docs  # ίδιο fallback με το main.py


def is_relevant(text, keywords):
    """True αν ΟΠΟΙΑΔΗΠΟΤΕ από τις keyword φράσεις (χωρισμένες με '|') βρεθεί στο κείμενο."""
    text_lower = text.lower()
    for kw in keywords.split("|"):
        if kw.strip().lower() in text_lower:
            return True
    return False


def evaluate():
    print("=" * 70)
    print("ΦΟΡΤΩΣΗ ΒΑΣΗΣ ΓΝΩΣΗΣ (μία φορά, μπορεί να πάρει λίγα λεπτά)...")
    print("=" * 70)
    vector_store, all_chunks = build_knowledge_base()

    print("[INFO] Φόρτωση reranker model (BAAI/bge-reranker-v2-m3)...")
    reranker = CrossEncoder("BAAI/bge-reranker-v2-m3")

    if not os.path.exists(EVAL_FILE):
        print(f"[ΣΦΑΛΜΑ] Δεν βρέθηκε το αρχείο {EVAL_FILE}. Δημιούργησέ το πρώτα.")
        sys.exit(1)

    with open(EVAL_FILE, mode='r', encoding='utf-8-sig') as f:
        eval_rows = list(csv.DictReader(f, delimiter=';'))

    print(f"[INFO] Φορτώθηκαν {len(eval_rows)} ερωτήσεις αξιολόγησης.")
    print("=" * 70)

    results = []

    for row in eval_rows:
        qid = row["question_id"]
        question = row["question"]
        category = row["category"].strip() or None
        keywords = row["required_keywords"]

        print(f"\n[{qid}] {question}  (κατηγορία: {category or 'καμία'})")

        # 1. MMR retrieval (ίδιο με main.py)
        query_lower = question.lower()
        docs = vector_store.max_marginal_relevance_search(query_lower, k=MMR_K, fetch_k=MMR_FETCH_K)

        # 2. Πορτιέρης
        filtered_docs = apply_portiere_filter(docs, category)

        # 3. Reranking
        if filtered_docs:
            pairs = [[query_lower, doc.page_content] for doc in filtered_docs]
            scores = reranker.predict(pairs)
            reranked = sorted(zip(filtered_docs, scores), key=lambda x: x[1], reverse=True)
        else:
            reranked = []

        top_docs = [doc for doc, score in reranked[:TOP_K_FINAL]]

        # ---- Μετρικές ----

        # Precision@k: πόσα από τα top-k είναι πράγματι σχετικά
        relevant_in_top_k = sum(1 for doc in top_docs if is_relevant(doc.page_content, keywords))
        precision = relevant_in_top_k / len(top_docs) if top_docs else 0.0

        # Recall@k: πόσα από ΟΛΑ τα σχετικά chunks στη βάση βρέθηκαν στα top-k
        total_relevant_in_corpus = sum(1 for c in all_chunks if is_relevant(c.page_content, keywords))
        recall = (relevant_in_top_k / total_relevant_in_corpus) if total_relevant_in_corpus > 0 else None

        # Reciprocal Rank: θέση του πρώτου σχετικού chunk μέσα στο ΠΛΗΡΕΣ reranked (όχι μόνο top-k)
        rr = 0.0
        for idx, (doc, score) in enumerate(reranked, start=1):
            if is_relevant(doc.page_content, keywords):
                rr = 1.0 / idx
                break

        results.append({
            "question_id": qid,
            "question": question,
            "category": category or "",
            "total_relevant_in_corpus": total_relevant_in_corpus,
            "relevant_in_top_k": relevant_in_top_k,
            "precision_at_k": round(precision, 3),
            "recall_at_k": round(recall, 3) if recall is not None else "N/A",
            "reciprocal_rank": round(rr, 3),
        })

        # ΔΙΑΓΝΩΣΤΙΚΟ: αν το RR=0 (πλήρης αποτυχία) αλλά ΥΠΑΡΧΟΥΝ σχετικά chunks
        # στη βάση, δείξε μας το περιεχόμενό τους -- βοηθάει να καταλάβουμε
        # ΓΙΑΤΙ δεν βρέθηκαν (κακό chunking, διαφορετική διατύπωση, portiere κλπ.)
        if rr == 0.0 and total_relevant_in_corpus > 0:
            print(f"    [DIAGNOSTIC] Το σχετικό chunk ΔΕΝ βρέθηκε πουθενά στα αποτελέσματα.")
            print(f"    [DIAGNOSTIC] Ο πορτιέρης άφησε {len(filtered_docs)}/{len(docs)} chunks μετά το φιλτράρισμα.")
            shown = 0
            for c in all_chunks:
                if is_relevant(c.page_content, keywords):
                    filename = c.metadata.get("filename", "?")
                    preview = c.page_content[:200].replace("\n", " ")
                    print(f"    [DIAGNOSTIC] Πραγματικό σχετικό chunk (από: {filename}):")
                    print(f"    [DIAGNOSTIC]   \"{preview}...\"")
                    shown += 1
                    if shown >= 3:
                        break

        print(f"    -> Precision@{TOP_K_FINAL}: {precision:.3f} | "
              f"Recall@{TOP_K_FINAL}: {recall if recall is None else round(recall,3)} | "
              f"RR: {rr:.3f}  (σχετικά στη βάση: {total_relevant_in_corpus})")

    # ---- Αποθήκευση αναλυτικών αποτελεσμάτων ----
    out_path = "evaluation_results.csv"
    with open(out_path, mode='w', newline='', encoding='utf-8-sig') as f:
        writer = csv.DictWriter(f, fieldnames=list(results[0].keys()))
        writer.writeheader()
        writer.writerows(results)

    # ---- Συνολικοί μέσοι όροι ----
    valid_recalls = [r["recall_at_k"] for r in results if r["recall_at_k"] != "N/A"]
    mean_precision = sum(r["precision_at_k"] for r in results) / len(results)
    mean_recall = sum(valid_recalls) / len(valid_recalls) if valid_recalls else float('nan')
    mean_rr = sum(r["reciprocal_rank"] for r in results) / len(results)

    print("\n" + "=" * 70)
    print("ΣΥΝΟΛΙΚΑ ΑΠΟΤΕΛΕΣΜΑΤΑ")
    print("=" * 70)
    print(f"Μέσο Precision@{TOP_K_FINAL}:  {mean_precision:.3f}")
    print(f"Μέσο Recall@{TOP_K_FINAL}:     {mean_recall:.3f}")
    print(f"MRR (Mean Reciprocal Rank): {mean_rr:.3f}")
    print(f"\nΑναλυτικά αποτελέσματα αποθηκεύτηκαν στο: {out_path}")
    print("=" * 70)


if __name__ == "__main__":
    evaluate()
