"""
Chain-of-Thought (CoT) Prompting for Enhanced Reasoning

Implements Chain-of-Thought prompting techniques to improve model reasoning:
- Zero-shot CoT ("Let's think step by step")
- Few-shot CoT (with reasoning examples)
- Structured reasoning templates
- Answer extraction from CoT generations

References:
- Wei et al. 2022: "Chain-of-Thought Prompting Elicits Reasoning in Large Language Models"
- Kojima et al. 2022: "Large Language Models are Zero-Shot Reasoners"
"""

import torch
import torch.nn.functional as F
from typing import Optional, List, Dict, Tuple, Callable
from dataclasses import dataclass, field
import re


@dataclass
class CoTConfig:
    """Configuration for Chain-of-Thought generation."""

    # CoT prompting style
    use_zero_shot: bool = True  # Use "Let's think step by step"
    use_few_shot: bool = False  # Use example demonstrations

    # Zero-shot prompts
    zero_shot_prompt: str = "Let's think step by step."
    alternative_prompts: List[str] = field(default_factory=lambda: [
        "Let's solve this step by step.",
        "Let's break this down:",
        "Let's approach this systematically:",
        "Let's reason through this carefully:",
    ])

    # Answer extraction
    extract_final_answer: bool = True
    answer_pattern: str = r"(?:answer|result|solution|conclusion)(?:\s+is)?:\s*(.+?)(?:\n|$)"
    answer_markers: List[str] = field(default_factory=lambda: [
        "Therefore,",
        "Thus,",
        "So,",
        "In conclusion,",
        "The answer is",
        "The result is",
        "The solution is",
    ])

    # Generation parameters for reasoning
    reasoning_max_length: int = 512
    reasoning_temperature: float = 0.7
    reasoning_top_p: float = 0.9

    # Structured reasoning
    use_structured_steps: bool = False
    step_markers: List[str] = field(default_factory=lambda: [
        "Step 1:",
        "Step 2:",
        "Step 3:",
        "Step 4:",
        "Step 5:",
    ])


@dataclass
class CoTPrompt:
    """
    Chain-of-Thought prompt with optional examples.
    """
    question: str
    examples: List[Dict[str, str]] = field(default_factory=list)
    system_prompt: Optional[str] = None

    def format_few_shot(self, cot_config: CoTConfig) -> str:
        """
        Format prompt with few-shot examples.

        Args:
            cot_config: CoT configuration

        Returns:
            Formatted prompt string
        """
        parts = []

        # Add system prompt if provided
        if self.system_prompt:
            parts.append(self.system_prompt)
            parts.append("")

        # Add examples
        for i, example in enumerate(self.examples):
            parts.append(f"Question: {example['question']}")
            parts.append(f"Answer: {example['reasoning']}")
            if 'final_answer' in example:
                parts.append(f"Therefore, the answer is: {example['final_answer']}")
            parts.append("")

        # Add current question
        parts.append(f"Question: {self.question}")
        parts.append("Answer:")

        return "\n".join(parts)

    def format_zero_shot(self, cot_config: CoTConfig) -> str:
        """
        Format prompt for zero-shot CoT.

        Args:
            cot_config: CoT configuration

        Returns:
            Formatted prompt string
        """
        parts = []

        # Add system prompt if provided
        if self.system_prompt:
            parts.append(self.system_prompt)
            parts.append("")

        # Add question
        parts.append(f"Question: {self.question}")
        parts.append("")

        # Add zero-shot prompt
        parts.append(cot_config.zero_shot_prompt)
        parts.append("")

        return "\n".join(parts)

    def format_structured(self, cot_config: CoTConfig) -> str:
        """
        Format prompt for structured step-by-step reasoning.

        Args:
            cot_config: CoT configuration

        Returns:
            Formatted prompt string
        """
        parts = []

        # Add system prompt if provided
        if self.system_prompt:
            parts.append(self.system_prompt)
            parts.append("")

        # Add question
        parts.append(f"Question: {self.question}")
        parts.append("")
        parts.append("Let's solve this step by step:")
        parts.append("")

        # Add step markers
        for step in cot_config.step_markers[:3]:  # Start with first 3 steps
            parts.append(step)

        return "\n".join(parts)


class ChainOfThoughtGenerator:
    """
    Generate text using Chain-of-Thought prompting for enhanced reasoning.
    """

    def __init__(
        self,
        model,
        tokenizer,
        config: Optional[CoTConfig] = None
    ):
        """
        Initialize CoT generator.

        Args:
            model: The language model (should have a generate method)
            tokenizer: Tokenizer for encoding/decoding
            config: CoT configuration
        """
        self.model = model
        self.tokenizer = tokenizer
        self.config = config or CoTConfig()

    def generate_with_cot(
        self,
        prompt: CoTPrompt,
        max_length: Optional[int] = None,
        temperature: Optional[float] = None,
        top_p: Optional[float] = None,
        num_return_sequences: int = 1,
        return_reasoning: bool = True,
        device: Optional[torch.device] = None
    ) -> Dict[str, any]:
        """
        Generate answer with Chain-of-Thought reasoning.

        Args:
            prompt: CoT prompt with question and optional examples
            max_length: Maximum generation length
            temperature: Sampling temperature
            top_p: Nucleus sampling parameter
            num_return_sequences: Number of reasoning paths to generate
            return_reasoning: Whether to return intermediate reasoning
            device: Device for computation

        Returns:
            Dictionary containing:
            - 'reasoning': Generated reasoning steps (if return_reasoning=True)
            - 'final_answer': Extracted final answer
            - 'full_text': Complete generated text
            - 'confidence': Confidence score (if multiple sequences)
        """
        # Use config defaults if not provided
        max_length = max_length or self.config.reasoning_max_length
        temperature = temperature or self.config.reasoning_temperature
        top_p = top_p or self.config.reasoning_top_p
        device = device or next(self.model.parameters()).device

        # Format prompt based on config
        if self.config.use_few_shot and prompt.examples:
            formatted_prompt = prompt.format_few_shot(self.config)
        elif self.config.use_structured_steps:
            formatted_prompt = prompt.format_structured(self.config)
        else:
            formatted_prompt = prompt.format_zero_shot(self.config)

        # Tokenize prompt
        prompt_ids = self.tokenizer.encode(formatted_prompt, add_special_tokens=True)
        prompt_tensor = torch.tensor([prompt_ids], dtype=torch.long, device=device)

        # Repeat for multiple sequences
        if num_return_sequences > 1:
            prompt_tensor = prompt_tensor.repeat(num_return_sequences, 1)

        # Generate reasoning
        self.model.eval()
        with torch.no_grad():
            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 generated text
        results = []
        for i in range(num_return_sequences):
            full_text = self.tokenizer.decode(generated_ids[i], skip_special_tokens=True)

            # Extract reasoning (remove prompt)
            reasoning = full_text[len(formatted_prompt):].strip()

            # Extract final answer
            final_answer = None
            if self.config.extract_final_answer:
                final_answer = self._extract_answer(reasoning)

            results.append({
                'full_text': full_text,
                'reasoning': reasoning if return_reasoning else None,
                'final_answer': final_answer,
                'prompt': formatted_prompt
            })

        # If single sequence, return single result
        if num_return_sequences == 1:
            return results[0]

        # If multiple sequences, compute confidence and return all
        return {
            'results': results,
            'num_sequences': num_return_sequences,
            'consensus_answer': self._find_consensus([r['final_answer'] for r in results])
        }

    def _extract_answer(self, reasoning_text: str) -> Optional[str]:
        """
        Extract final answer from reasoning text.

        Args:
            reasoning_text: Generated reasoning text

        Returns:
            Extracted answer or None
        """
        # Try regex pattern first
        match = re.search(self.config.answer_pattern, reasoning_text, re.IGNORECASE)
        if match:
            return match.group(1).strip()

        # Try answer markers
        for marker in self.config.answer_markers:
            if marker.lower() in reasoning_text.lower():
                # Get text after marker
                parts = reasoning_text.lower().split(marker.lower(), 1)
                if len(parts) > 1:
                    answer = parts[1].strip()
                    # Take first sentence or up to newline
                    answer = answer.split('.')[0].split('\n')[0].strip()
                    if answer:
                        return answer

        # Fallback: return last sentence
        sentences = reasoning_text.split('.')
        if sentences:
            last_sentence = sentences[-1].strip()
            if last_sentence:
                return last_sentence

        return None

    def _find_consensus(self, answers: List[Optional[str]]) -> Optional[str]:
        """
        Find consensus answer from multiple reasoning paths.

        Args:
            answers: List of extracted answers

        Returns:
            Most common answer or None
        """
        if not answers:
            return None

        # Filter out None values
        valid_answers = [a for a in answers if a is not None]
        if not valid_answers:
            return None

        # Count frequencies (case-insensitive)
        from collections import Counter
        normalized_answers = [a.lower().strip() for a in valid_answers]
        counter = Counter(normalized_answers)

        # Return most common
        most_common = counter.most_common(1)[0][0]

        # Return original casing
        for answer in valid_answers:
            if answer.lower().strip() == most_common:
                return answer

        return valid_answers[0]

    def generate_with_verification(
        self,
        prompt: CoTPrompt,
        verification_prompt: str = "Is this reasoning correct? Let's verify step by step.",
        **generation_kwargs
    ) -> Dict[str, any]:
        """
        Generate with Chain-of-Thought and then verify the reasoning.

        This implements a two-stage process:
        1. Generate initial reasoning
        2. Verify the reasoning with a second pass

        Args:
            prompt: CoT prompt
            verification_prompt: Prompt for verification step
            **generation_kwargs: Additional generation arguments

        Returns:
            Dictionary with initial reasoning, verification, and final answer
        """
        # First pass: generate reasoning
        initial_result = self.generate_with_cot(prompt, **generation_kwargs)

        # Second pass: verify reasoning
        verification_question = f"{initial_result['reasoning']}\n\n{verification_prompt}"
        verification_cot = CoTPrompt(
            question=verification_question,
            system_prompt=prompt.system_prompt
        )

        verification_result = self.generate_with_cot(
            verification_cot,
            **generation_kwargs
        )

        return {
            'initial_reasoning': initial_result['reasoning'],
            'initial_answer': initial_result['final_answer'],
            'verification': verification_result['reasoning'],
            'verified_answer': verification_result['final_answer'],
            'full_initial_text': initial_result['full_text'],
            'full_verification_text': verification_result['full_text']
        }


def create_math_cot_examples() -> List[Dict[str, str]]:
    """
    Create few-shot examples for mathematical reasoning.

    Returns:
        List of CoT examples for math problems
    """
    return [
        {
            'question': 'If a train travels 60 miles in 1 hour, how far will it travel in 3.5 hours?',
            'reasoning': '''
                Let's think step by step:
                1. The train travels 60 miles in 1 hour
                2. We need to find the distance in 3.5 hours
                3. Since speed is constant, distance = speed × time
                4. Distance = 60 miles/hour × 3.5 hours
                5. Distance = 210 miles
            ''',
            'final_answer': '210 miles'
        },
        {
            'question': 'A rectangle has a length of 8 cm and a width of 5 cm. What is its area?',
            'reasoning': '''
                Let's solve this step by step:
                1. We have a rectangle with length = 8 cm and width = 5 cm
                2. Area of a rectangle = length × width
                3. Area = 8 cm × 5 cm
                4. Area = 40 cm²
            ''',
            'final_answer': '40 cm²'
        }
    ]


def create_reasoning_cot_examples() -> List[Dict[str, str]]:
    """
    Create few-shot examples for logical reasoning.

    Returns:
        List of CoT examples for reasoning problems
    """
    return [
        {
            'question': 'All roses are flowers. Some flowers fade quickly. Can we conclude that some roses fade quickly?',
            'reasoning': '''
                Let's analyze this step by step:
                1. Premise 1: All roses are flowers (roses ⊂ flowers)
                2. Premise 2: Some flowers fade quickly
                3. Question: Do some roses fade quickly?
                4. Analysis: While all roses are flowers, we only know that SOME flowers fade quickly
                5. These "some flowers" that fade quickly might or might not include roses
                6. We cannot determine from these premises alone whether roses are among the flowers that fade quickly
            ''',
            'final_answer': 'No, we cannot conclude that some roses fade quickly from the given premises.'
        }
    ]
