# experiment_main.py
# Implementation of the SME Digital Intelligence Transformation Framework with Qwen-2

import json
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline
from sentence_transformers import SentenceTransformer
from bert_score import score as bert_score_eval
from bleurt import score as bleurt_score
from rouge_score import rouge_scorer
from sklearn.metrics.pairwise import cosine_similarity
import numpy as np
import re
from typing import List, Dict, Tuple
import warnings

warnings.filterwarnings("ignore")

# -----------------------------
# 1. Configuration
# -----------------------------

MODEL_NAME = "Qwen/Qwen2-7B-Instruct"  # Official Hugging Face ID
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
MAX_NEW_TOKENS = 256
TEST_DATA_PATH = "data/test_cases.jsonl"  # Format: {"input": "...", "target": "...", "task_type": "marketing|service"}

# Evaluation models
SBERT_MODEL = "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
BLEURT_MODEL_PATH = "google/bleurt-large-512"  # Or local path after download


# -----------------------------
# 2. Load Models
# -----------------------------

def load_models():
    print("Loading tokenizer and model...")
    tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, trust_remote_code=True)
    model = AutoModelForCausalLM.from_pretrained(
        MODEL_NAME,
        device_map="auto",
        torch_dtype=torch.bfloat16,
        trust_remote_code=True
    )
    
    generator = pipeline(
        "text-generation",
        model=model,
        tokenizer=tokenizer,
        device_map="auto",
        max_new_tokens=MAX_NEW_TOKENS,
        return_full_text=False
    )

    print("Loading SBERT evaluator...")
    sbert_model = SentenceTransformer(SBERT_MODEL, device=DEVICE)

    print("Loading BLEURT scorer...")
    bleurt_scorer = bleurt_score.BleurtScorer(BLEURT_MODEL_PATH)

    return generator, sbert_model, bleurt_scorer


# -----------------------------
# 3. Dynamic Prompt Builder (Enhancement Layer)
# -----------------------------

def build_dynamic_prompt(task_input: str, task_type: str, history_examples: List[Dict] = None) -> str:
    if task_type == "marketing":
        base_prompt = (
            "You are an e-commerce marketing expert. Based on the following customer review and product information, "
            "generate an engaging Chinese promotional copy. "
            "Requirements: vivid language, highlight key selling points, avoid repetitive expressions.\n\n"
            f"Input: {task_input}\n\n"
            "Output:"
        )
    elif task_type == "service":
        base_prompt = (
            "You are a customer service agent. Based on the following user request, generate a professional, polite, "
            "and complete response. If additional information is needed, clearly state so.\n\n"
            f"Request: {task_input}\n\n"
            "Response:"
        )
    else:
        base_prompt = f"Please complete the following task:\n{task_input}\n\nAnswer:"

    # Optional: Add few-shot examples from retrieval-augmented memory
    if history_examples:
        demo_str = "\n\n--- Reference Examples ---\n"
        for ex in history_examples[:2]:
            demo_str += f"Input: {ex['input']}\nOutput: {ex['output']}\n\n"
        base_prompt = demo_str + base_prompt

    return base_prompt


# -----------------------------
# 4. Generate Response
# -----------------------------

def generate_response(generator, prompt: str) -> str:
    try:
        outputs = generator(prompt, num_return_sequences=1)
        return outputs[0]["generated_text"].strip()
    except Exception as e:
        print(f"Generation error: {e}")
        return ""


# -----------------------------
# 5. Semantic Evaluation Metrics
# -----------------------------

def compute_rouge_l(generated: str, reference: str) -> float:
    scorer = rouge_scorer.RougeScorer(['rougeL'], use_stemmer=False)
    scores = scorer.score(reference, generated)
    return scores['rougeL'].fmeasure


def compute_bertscore(generated: List[str], reference: List[str]) -> float:
    P, R, F = bert_score_eval(cands=generated, refs=reference, lang="zh", verbose=False)
    return F.mean().item()


def compute_bleurt(generated: List[str], reference: List[str]) -> float:
    scores = bleurt_scorer.score(references=reference, candidates=generated)
    return float(np.mean(scores))


def compute_sbert_sim(generated: List[str], reference: List[str], sbert_model) -> float:
    gen_emb = sbert_model.encode(generated)
    ref_emb = sbert_model.encode(reference)
    sims = cosine_similarity(gen_emb, ref_emb).diagonal()
    return float(np.mean(sims))


def compute_distinct_n(text: str, n: int = 2) -> float:
    tokens = re.findall(r'\w+', text.lower())
    if len(tokens) < n:
        return 0.0
    ngrams = [' '.join(tokens[i:i+n]) for i in range(len(tokens)-n+1)]
    return len(set(ngrams)) / len(ngrams)


# -----------------------------
# 6. Business-Oriented Evaluation
# -----------------------------

def simulate_task_completion(generated: str, task_type: str) -> bool:
    required_keywords = {
        "marketing": ["recommended", "buy now", "limited-time offer", "discount"],
        "service": ["thank you", "solution provided", "will be processed", "handling"]
    }
    keywords = required_keywords.get(task_type, [])
    return any(kw in generated for kw in keywords)


def check_information_completeness(generated: str, required_fields: List[str]) -> bool:
    return all(field in generated for field in required_fields)


def simulate_first_turn_resolution(user_intent: str, generated: str) -> bool:
    resolution_indicators = ["resolved", "completed", "handled", "no further action needed", "issue closed"]
    return any(ind in generated for ind in resolution_indicators)


# -----------------------------
# 7. Main Experiment Loop
# -----------------------------

def main():
    print("🚀 Starting SME Digital Intelligence Transformation Experiment")
    generator, sbert_model, bleurt_scorer = load_models()

    # Mock retrieval-augmented memory (could be FAISS or Elasticsearch in real case)
    history_examples = [
        {
            "input": "User says phone battery drains too fast. What should I do?",
            "output": "Thank you for your feedback! We recommend closing background power-consuming apps and checking battery optimization settings. If the issue persists, free diagnostics are available."
        },
        {
            "input": "Is this headphone suitable for sports use?",
            "output": "Highly recommended! This model features IPX7 waterproof and sweatproof design, secure fit, and optimized audio for running and gym workouts."
        }
    ]

    results = {
        "rouge_l": [],
        "distinct_2": [],
        "bertscore": [],
        "bleurt": [],
        "sbert_sim": [],
        "task_completion_rate": [],
        "info_complete_rate": [],
        "first_turn_resolution": []
    }

    # Read test data line by line (simulate subsets)
    try:
        with open(TEST_DATA_PATH, 'r', encoding='utf-8') as f:
            test_cases = [json.loads(line.strip()) for line in f]
    except FileNotFoundError:
        print(f"Test file not found at {TEST_DATA_PATH}. Using mock data.")
        test_cases = [
            {
                "input": "The screen quality is excellent but delivery was slow.",
                "target": "We appreciate your positive feedback on display quality and apologize for the shipping delay. Your satisfaction matters to us.",
                "task_type": "service",
                "required_fields": ["apologize", "shipping delay", "appreciate"]
            }
        ] * 10

    for idx, case in enumerate(test_cases):
        print(f"\nProcessing Case {idx+1}: {case['task_type']}")

        # Build enhanced prompt
        prompt = build_dynamic_prompt(case["input"], case["task_type"], history_examples)

        # Generate response
        output = generate_response(generator, prompt)

        # Evaluate semantic quality
        rouge_l = compute_rouge_l(output, case["target"])
        distinct_2 = compute_distinct_n(output, n=2)
        bert_f1 = compute_bertscore([output], [case["target"]])
        bleurt_score_val = compute_bleurt([output], [case["target"]])
        sbert_sim = compute_sbert_sim([output], [case["target"]], sbert_model)

        # Simulate business KPIs
        tcr = simulate_task_completion(output, case["task_type"])
        icr = check_information_completeness(output, case.get("required_fields", []))
        ftrr = simulate_first_turn_resolution(case["input"], output)

        # Record results
        results["rouge_l"].append(rouge_l)
        results["distinct_2"].append(distinct_2)
        results["bertscore"].append(bert_f1)
        results["bleurt"].append(bleurt_score_val)
        results["sbert_sim"].append(sbert_sim)
        results["task_completion_rate"].append(int(tcr))
        results["info_complete_rate"].append(int(icr))
        results["first_turn_resolution"].append(int(ftrr))

        print(f"ROUGE-L: {rouge_l:.3f}, Distinct-2: {distinct_2:.3f}")
        print(f"BERTScore: {bert_f1:.3f}, BLEURT: {bleurt_score_val:.3f}, SBERT-Sim: {sbert_sim:.3f}")
        print(f"TCR: {tcr}, ICR: {icr}, FTRR: {ftrr}")

    # Final aggregate results (to match paper tables)
    print("\n" + "="*60)
    print("AGGREGATE RESULTS (AVERAGE ACROSS TEST SET)")
    print("="*60)
    for k, v in results.items():
        avg_val = np.mean(v)
        print(f"{k.upper():<30} {avg_val:.3f}")


if __name__ == "__main__":
    main()
