"""
Self-Consistency: Sampling Multiple Reasoning Paths

Implements self-consistency to improve reasoning accuracy by:
1. Sampling multiple diverse reasoning paths
2. Aggregating answers from different paths
3. Selecting the most consistent answer

This significantly improves accuracy on complex reasoning tasks.

Reference:
- Wang et al. 2022: "Self-Consistency Improves Chain of Thought Reasoning in Language Models"
"""

import torch
import torch.nn.functional as F
from typing import Optional, List, Dict, Tuple, Callable, Union
from collections import Counter, defaultdict
from dataclasses import dataclass
import re
import numpy as np


@dataclass
class SelfConsistencyConfig:
    """Configuration for self-consistency generation."""

    # Number of reasoning paths
    num_paths: int = 5

    # Diversity settings
    temperature: float = 0.7  # Higher temperature for diversity
    top_p: float = 0.9
    use_diverse_prompts: bool = True

    # Aggregation method
    aggregation_method: str = "majority_vote"  # "majority_vote", "weighted_vote", "clustering"

    # Answer extraction
    normalize_answers: bool = True
    case_sensitive: bool = False

    # Confidence threshold
    min_agreement: float = 0.4  # Minimum fraction of paths that must agree


class SelfConsistencyGenerator:
    """
    Generate multiple reasoning paths and aggregate results.
    """

    def __init__(
        self,
        model,
        tokenizer,
        config: Optional[SelfConsistencyConfig] = None
    ):
        """
        Initialize self-consistency generator.

        Args:
            model: The language model
            tokenizer: Tokenizer for encoding/decoding
            config: Self-consistency configuration
        """
        self.model = model
        self.tokenizer = tokenizer
        self.config = config or SelfConsistencyConfig()

    def generate_multiple_paths(
        self,
        prompt: str,
        num_paths: Optional[int] = None,
        max_length: int = 512,
        temperature: Optional[float] = None,
        top_p: Optional[float] = None,
        device: Optional[torch.device] = None,
        answer_extractor: Optional[Callable[[str], str]] = None
    ) -> List[Dict[str, any]]:
        """
        Generate multiple diverse reasoning paths.

        Args:
            prompt: Input prompt/question
            num_paths: Number of reasoning paths to generate
            max_length: Maximum generation length
            temperature: Sampling temperature (higher = more diverse)
            top_p: Nucleus sampling parameter
            device: Device for computation
            answer_extractor: Function to extract answer from generated text

        Returns:
            List of dictionaries, each containing:
            - 'text': Generated reasoning text
            - 'answer': Extracted answer
            - 'tokens': Generated token IDs
        """
        num_paths = num_paths or self.config.num_paths
        temperature = temperature or self.config.temperature
        top_p = top_p or self.config.top_p
        device = device or next(self.model.parameters()).device

        paths = []

        self.model.eval()
        with torch.no_grad():
            for i in range(num_paths):
                # Optionally vary prompt for diversity
                if self.config.use_diverse_prompts and i > 0:
                    current_prompt = self._create_diverse_prompt(prompt, i)
                else:
                    current_prompt = prompt

                # Tokenize
                prompt_ids = self.tokenizer.encode(current_prompt, add_special_tokens=True)
                prompt_tensor = torch.tensor([prompt_ids], dtype=torch.long, device=device)

                # Generate with higher temperature for diversity
                generated_ids = self.model.generate(
                    initial_tokens=prompt_tensor,
                    max_new_vectors=max_length // self.model.chunk_size,
                    temperature=temperature,
                    top_p=top_p
                )

                # Decode
                full_text = self.tokenizer.decode(generated_ids[0], skip_special_tokens=True)
                reasoning_text = full_text[len(current_prompt):].strip()

                # Extract answer
                if answer_extractor:
                    answer = answer_extractor(reasoning_text)
                else:
                    answer = self._default_answer_extraction(reasoning_text)

                paths.append({
                    'text': reasoning_text,
                    'answer': answer,
                    'tokens': generated_ids[0],
                    'prompt_variant': i
                })

        return paths

    def aggregate_answers(
        self,
        paths: List[Dict[str, any]],
        method: Optional[str] = None
    ) -> Dict[str, any]:
        """
        Aggregate answers from multiple reasoning paths.

        Args:
            paths: List of path dictionaries from generate_multiple_paths
            method: Aggregation method ("majority_vote", "weighted_vote", "clustering")

        Returns:
            Dictionary containing:
            - 'final_answer': Selected answer
            - 'confidence': Confidence score (0-1)
            - 'vote_distribution': Distribution of answers
            - 'agreement': Fraction of paths agreeing with final answer
            - 'all_answers': All extracted answers
        """
        method = method or self.config.aggregation_method

        # Extract and normalize answers
        answers = [path['answer'] for path in paths if path['answer'] is not None]

        if not answers:
            return {
                'final_answer': None,
                'confidence': 0.0,
                'vote_distribution': {},
                'agreement': 0.0,
                'all_answers': []
            }

        # Normalize answers if configured
        if self.config.normalize_answers:
            answers = [self._normalize_answer(ans) for ans in answers]

        # Aggregate based on method
        if method == "majority_vote":
            result = self._majority_vote(answers)
        elif method == "weighted_vote":
            result = self._weighted_vote(answers, paths)
        elif method == "clustering":
            result = self._clustering_aggregate(answers)
        else:
            raise ValueError(f"Unknown aggregation method: {method}")

        # Add metadata
        result['all_answers'] = answers
        result['num_paths'] = len(paths)

        return result

    def generate_with_self_consistency(
        self,
        prompt: str,
        num_paths: Optional[int] = None,
        max_length: int = 512,
        temperature: Optional[float] = None,
        answer_extractor: Optional[Callable[[str], str]] = None,
        return_all_paths: bool = False
    ) -> Dict[str, any]:
        """
        Complete self-consistency pipeline: generate paths and aggregate.

        Args:
            prompt: Input prompt/question
            num_paths: Number of reasoning paths
            max_length: Maximum generation length
            temperature: Sampling temperature
            answer_extractor: Custom answer extraction function
            return_all_paths: Whether to return all reasoning paths

        Returns:
            Dictionary with final answer, confidence, and optionally all paths
        """
        # Generate multiple paths
        paths = self.generate_multiple_paths(
            prompt=prompt,
            num_paths=num_paths,
            max_length=max_length,
            temperature=temperature,
            answer_extractor=answer_extractor
        )

        # Aggregate answers
        result = self.aggregate_answers(paths)

        # Add paths if requested
        if return_all_paths:
            result['paths'] = paths

        # Check if agreement meets threshold
        if result['agreement'] < self.config.min_agreement:
            result['warning'] = f"Low agreement ({result['agreement']:.2%}), answer may be unreliable"

        return result

    def _create_diverse_prompt(self, base_prompt: str, variant_idx: int) -> str:
        """
        Create a diverse prompt variant to encourage different reasoning paths.

        Args:
            base_prompt: Original prompt
            variant_idx: Variant index

        Returns:
            Modified prompt
        """
        # Different reasoning triggers
        triggers = [
            "Let's think step by step.",
            "Let's solve this carefully.",
            "Let's approach this systematically:",
            "Let's break this down:",
            "Let's work through this:",
        ]

        # Append a different trigger
        trigger = triggers[variant_idx % len(triggers)]
        return f"{base_prompt}\n\n{trigger}"

    def _default_answer_extraction(self, text: str) -> Optional[str]:
        """
        Default answer extraction from reasoning text.

        Args:
            text: Reasoning text

        Returns:
            Extracted answer or None
        """
        # Try common patterns
        patterns = [
            r"(?:the\s+)?(?:final\s+)?answer\s+is\s*:?\s*(.+?)(?:\.|$)",
            r"therefore\s*,?\s*(.+?)(?:\.|$)",
            r"thus\s*,?\s*(.+?)(?:\.|$)",
            r"so\s*,?\s*(.+?)(?:\.|$)",
        ]

        for pattern in patterns:
            match = re.search(pattern, text, re.IGNORECASE)
            if match:
                answer = match.group(1).strip()
                # Clean up
                answer = answer.split('\n')[0]  # Take first line
                return answer

        # Fallback: last sentence
        sentences = text.split('.')
        if sentences:
            return sentences[-1].strip()

        return text.strip() if text else None

    def _normalize_answer(self, answer: str) -> str:
        """
        Normalize answer for comparison.

        Args:
            answer: Raw answer string

        Returns:
            Normalized answer
        """
        if answer is None:
            return ""

        # Remove punctuation and extra whitespace
        normalized = re.sub(r'[^\w\s]', '', answer)
        normalized = ' '.join(normalized.split())

        # Case normalization
        if not self.config.case_sensitive:
            normalized = normalized.lower()

        return normalized

    def _majority_vote(self, answers: List[str]) -> Dict[str, any]:
        """
        Simple majority voting.

        Args:
            answers: List of answers

        Returns:
            Result dictionary
        """
        # Count votes
        counter = Counter(answers)
        total = len(answers)

        # Get most common
        most_common_answer, most_common_count = counter.most_common(1)[0]

        # Calculate confidence and agreement
        agreement = most_common_count / total
        confidence = agreement  # Simple confidence = agreement fraction

        return {
            'final_answer': most_common_answer,
            'confidence': confidence,
            'agreement': agreement,
            'vote_distribution': dict(counter)
        }

    def _weighted_vote(
        self,
        answers: List[str],
        paths: List[Dict[str, any]]
    ) -> Dict[str, any]:
        """
        Weighted voting based on answer likelihood or other metrics.

        Args:
            answers: List of answers
            paths: List of path dictionaries

        Returns:
            Result dictionary
        """
        # For now, use simple majority vote
        # Could be extended to weight by:
        # - Answer probability
        # - Reasoning length/quality
        # - Model confidence
        return self._majority_vote(answers)

    def _clustering_aggregate(self, answers: List[str]) -> Dict[str, any]:
        """
        Aggregate answers using semantic clustering.

        Groups similar answers and selects from largest cluster.

        Args:
            answers: List of answers

        Returns:
            Result dictionary
        """
        # For text answers, use string similarity
        clusters = defaultdict(list)

        for answer in answers:
            # Find most similar cluster
            best_cluster = None
            best_similarity = 0.0

            for cluster_rep in clusters.keys():
                similarity = self._answer_similarity(answer, cluster_rep)
                if similarity > best_similarity and similarity > 0.8:  # Threshold
                    best_similarity = similarity
                    best_cluster = cluster_rep

            if best_cluster:
                clusters[best_cluster].append(answer)
            else:
                # Create new cluster
                clusters[answer].append(answer)

        # Find largest cluster
        largest_cluster = max(clusters.items(), key=lambda x: len(x[1]))
        cluster_rep, cluster_members = largest_cluster

        # Calculate statistics
        total = len(answers)
        agreement = len(cluster_members) / total
        confidence = agreement

        # Use most common answer in cluster as final answer
        counter = Counter(cluster_members)
        final_answer = counter.most_common(1)[0][0]

        return {
            'final_answer': final_answer,
            'confidence': confidence,
            'agreement': agreement,
            'vote_distribution': dict(Counter(answers)),
            'num_clusters': len(clusters)
        }

    def _answer_similarity(self, answer1: str, answer2: str) -> float:
        """
        Compute similarity between two answers.

        Args:
            answer1: First answer
            answer2: Second answer

        Returns:
            Similarity score (0-1)
        """
        # Normalize
        a1 = self._normalize_answer(answer1)
        a2 = self._normalize_answer(answer2)

        # Exact match
        if a1 == a2:
            return 1.0

        # Jaccard similarity on words
        words1 = set(a1.split())
        words2 = set(a2.split())

        if not words1 or not words2:
            return 0.0

        intersection = len(words1 & words2)
        union = len(words1 | words2)

        return intersection / union if union > 0 else 0.0


def aggregate_answers(
    answers: List[str],
    method: str = "majority_vote",
    normalize: bool = True,
    case_sensitive: bool = False
) -> Dict[str, any]:
    """
    Convenience function to aggregate a list of answers.

    Args:
        answers: List of answer strings
        method: Aggregation method
        normalize: Whether to normalize answers
        case_sensitive: Whether comparison is case-sensitive

    Returns:
        Dictionary with final answer and statistics
    """
    # Create a minimal generator instance for aggregation
    config = SelfConsistencyConfig(
        aggregation_method=method,
        normalize_answers=normalize,
        case_sensitive=case_sensitive
    )

    generator = SelfConsistencyGenerator(
        model=None,  # Not needed for aggregation only
        tokenizer=None,
        config=config
    )

    # Create paths structure
    paths = [{'answer': ans, 'text': ''} for ans in answers]

    return generator.aggregate_answers(paths, method=method)


def compute_consistency_score(answers: List[str], normalize: bool = True) -> float:
    """
    Compute consistency score for a set of answers.

    Higher score = more agreement between answers.

    Args:
        answers: List of answer strings
        normalize: Whether to normalize answers before comparison

    Returns:
        Consistency score (0-1)
    """
    if not answers:
        return 0.0

    if len(answers) == 1:
        return 1.0

    # Normalize if requested
    if normalize:
        config = SelfConsistencyConfig(normalize_answers=True, case_sensitive=False)
        generator = SelfConsistencyGenerator(None, None, config)
        answers = [generator._normalize_answer(ans) for ans in answers]

    # Compute agreement as fraction agreeing with most common answer
    counter = Counter(answers)
    most_common_count = counter.most_common(1)[0][1]

    return most_common_count / len(answers)


def filter_by_consistency(
    paths: List[Dict[str, any]],
    min_consistency: float = 0.5,
    answer_key: str = 'answer'
) -> List[Dict[str, any]]:
    """
    Filter reasoning paths to keep only those consistent with majority.

    Args:
        paths: List of path dictionaries
        min_consistency: Minimum consistency threshold
        answer_key: Key in path dict containing the answer

    Returns:
        Filtered list of paths
    """
    if not paths:
        return []

    # Get all answers
    answers = [path[answer_key] for path in paths if answer_key in path]

    if not answers:
        return paths

    # Find most common answer
    counter = Counter(answers)
    most_common_answer = counter.most_common(1)[0][0]

    # Filter paths
    filtered = [
        path for path in paths
        if path.get(answer_key) == most_common_answer
    ]

    # Check if we meet threshold
    agreement = len(filtered) / len(paths)

    if agreement >= min_consistency:
        return filtered
    else:
        # Return all if threshold not met
        return paths
