"""
Utility classes and functions for preference learning.

Supports pairwise and multi-response preference data.
"""

import torch
from typing import List, Optional, Dict, Any, Union
from dataclasses import dataclass
import json


@dataclass
class ResponsePair:
    """
    Pairwise preference data.

    Attributes:
        prompt: Input prompt tokens
        chosen: Preferred response tokens
        rejected: Rejected response tokens
        margin: Optional preference margin (for calibrated preferences)
    """
    prompt: torch.Tensor
    chosen: torch.Tensor
    rejected: torch.Tensor
    margin: Optional[float] = None

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return {
            'prompt': self.prompt.tolist() if isinstance(self.prompt, torch.Tensor) else self.prompt,
            'chosen': self.chosen.tolist() if isinstance(self.chosen, torch.Tensor) else self.chosen,
            'rejected': self.rejected.tolist() if isinstance(self.rejected, torch.Tensor) else self.rejected,
            'margin': self.margin
        }


@dataclass
class StepPreference:
    """
    Step-wise preference for reasoning tasks.

    Each step has its own preference signal for fine-grained
    feedback on reasoning processes.

    Attributes:
        prompt: Input prompt tokens
        chosen_steps: List of preferred reasoning steps
        rejected_steps: List of rejected reasoning steps
        step_weights: Optional weights for each step
    """
    prompt: torch.Tensor
    chosen_steps: List[torch.Tensor]
    rejected_steps: List[torch.Tensor]
    step_weights: Optional[List[float]] = None

    def __post_init__(self):
        """Validate step preference data."""
        if len(self.chosen_steps) != len(self.rejected_steps):
            raise ValueError(
                f"Chosen and rejected steps must have same length, "
                f"got {len(self.chosen_steps)} vs {len(self.rejected_steps)}"
            )

        if self.step_weights is not None:
            if len(self.step_weights) != len(self.chosen_steps):
                raise ValueError(
                    f"Step weights length {len(self.step_weights)} must match "
                    f"number of steps {len(self.chosen_steps)}"
                )


@dataclass
class PreferenceData:
    """
    General preference data container.

    Supports multiple preference formats:
    - Pairwise comparisons
    - Rankings (multiple responses)
    - Binary feedback (good/bad)
    """
    prompt: torch.Tensor
    responses: List[torch.Tensor]
    rankings: Optional[List[int]] = None  # Lower is better
    binary_labels: Optional[List[bool]] = None  # True = good, False = bad
    scores: Optional[List[float]] = None  # Explicit preference scores

    def to_pairwise(self) -> List[ResponsePair]:
        """
        Convert to pairwise preferences.

        Returns:
            List of ResponsePair objects for all pairwise comparisons.
        """
        pairs = []

        if self.rankings is not None:
            # Convert rankings to pairs
            sorted_indices = sorted(range(len(self.rankings)), key=lambda i: self.rankings[i])
            for i in range(len(sorted_indices) - 1):
                better_idx = sorted_indices[i]
                worse_idx = sorted_indices[i + 1]
                pairs.append(ResponsePair(
                    prompt=self.prompt,
                    chosen=self.responses[better_idx],
                    rejected=self.responses[worse_idx]
                ))

        elif self.binary_labels is not None:
            # Pair good responses with bad responses
            good_indices = [i for i, label in enumerate(self.binary_labels) if label]
            bad_indices = [i for i, label in enumerate(self.binary_labels) if not label]

            for good_idx in good_indices:
                for bad_idx in bad_indices:
                    pairs.append(ResponsePair(
                        prompt=self.prompt,
                        chosen=self.responses[good_idx],
                        rejected=self.responses[bad_idx]
                    ))

        elif self.scores is not None:
            # Pair based on scores
            sorted_indices = sorted(range(len(self.scores)), key=lambda i: -self.scores[i])
            for i in range(len(sorted_indices) - 1):
                better_idx = sorted_indices[i]
                worse_idx = sorted_indices[i + 1]
                margin = self.scores[better_idx] - self.scores[worse_idx]
                pairs.append(ResponsePair(
                    prompt=self.prompt,
                    chosen=self.responses[better_idx],
                    rejected=self.responses[worse_idx],
                    margin=margin
                ))

        return pairs


def compute_log_prob(
    model,
    prompt_vectors: torch.Tensor,
    response_vectors: torch.Tensor,
    reduction: str = 'mean'
) -> torch.Tensor:
    """
    Compute log probability of response given prompt for continuous model.

    Since we work with continuous vectors, we compute:
    - Cosine similarity between predicted and actual vectors
    - Convert to log-probability via softmax over candidates

    Args:
        model: ContinuousAutoregressiveModel
        prompt_vectors: [batch, num_prompt_chunks, vector_dim]
        response_vectors: [batch, num_response_chunks, vector_dim]
        reduction: 'mean', 'sum', or 'none'

    Returns:
        log_prob: Log probability of response
    """
    batch_size = prompt_vectors.shape[0]

    # Combine prompt and response
    full_sequence = torch.cat([prompt_vectors, response_vectors], dim=1)

    # Forward pass
    output = model.forward(full_sequence, use_cache=False)
    predicted_vectors = output['predicted_vectors']

    # Get predictions for response positions
    prompt_len = prompt_vectors.shape[1]
    response_len = response_vectors.shape[1]
    predicted_response = predicted_vectors[:, prompt_len - 1:prompt_len + response_len - 1, :]

    # Compute cosine similarity (proxy for log probability in continuous space)
    cos_sim = torch.nn.functional.cosine_similarity(
        predicted_response,
        response_vectors,
        dim=-1
    )  # [batch, response_len]

    # Convert to log probability via log-softmax over sequence
    # Higher cosine similarity = higher probability
    log_prob = torch.nn.functional.log_softmax(cos_sim, dim=-1)

    if reduction == 'mean':
        return log_prob.mean(dim=-1)
    elif reduction == 'sum':
        return log_prob.sum(dim=-1)
    else:
        return log_prob


def compute_sequence_log_prob(
    model,
    continuous_vectors: torch.Tensor,
    reduction: str = 'mean'
) -> torch.Tensor:
    """
    Compute log probability of full sequence.

    Args:
        model: ContinuousAutoregressiveModel
        continuous_vectors: [batch, num_chunks, vector_dim]
        reduction: 'mean', 'sum', or 'none'

    Returns:
        log_prob: Log probability
    """
    # Forward pass
    output = model.forward(continuous_vectors, use_cache=False)
    predicted_vectors = output['predicted_vectors']

    # Shift for autoregressive prediction
    predicted = predicted_vectors[:, :-1, :]
    target = continuous_vectors[:, 1:, :]

    # Cosine similarity
    cos_sim = torch.nn.functional.cosine_similarity(predicted, target, dim=-1)

    # Convert to log probability
    log_prob = torch.nn.functional.log_softmax(cos_sim, dim=-1)

    if reduction == 'mean':
        return log_prob.mean(dim=-1)
    elif reduction == 'sum':
        return log_prob.sum(dim=-1)
    else:
        return log_prob


def estimate_kl_divergence(
    policy_log_prob: torch.Tensor,
    reference_log_prob: torch.Tensor
) -> torch.Tensor:
    """
    Estimate KL divergence between policy and reference.

    KL(π || π_ref) ≈ log π - log π_ref

    Args:
        policy_log_prob: Log probability under policy
        reference_log_prob: Log probability under reference

    Returns:
        kl: KL divergence estimate
    """
    return policy_log_prob - reference_log_prob


def margin_scaling(margin: Optional[float], beta: float = 1.0) -> float:
    """
    Scale preference margin for calibrated learning.

    Args:
        margin: Raw preference margin
        beta: Scaling factor

    Returns:
        scaled_margin: Scaled margin
    """
    if margin is None:
        return beta

    # Apply exponential scaling for larger margins
    return beta * (1.0 + torch.sigmoid(torch.tensor(margin)).item())


def batch_preference_pairs(
    pairs: List[ResponsePair],
    batch_size: int
) -> List[List[ResponsePair]]:
    """
    Batch preference pairs for efficient training.

    Args:
        pairs: List of preference pairs
        batch_size: Batch size

    Returns:
        batches: List of batches
    """
    batches = []
    for i in range(0, len(pairs), batch_size):
        batches.append(pairs[i:i + batch_size])
    return batches


def normalize_preference_scores(scores: List[float]) -> List[float]:
    """
    Normalize preference scores to [0, 1] range.

    Args:
        scores: Raw preference scores

    Returns:
        normalized_scores: Normalized scores
    """
    if not scores:
        return scores

    min_score = min(scores)
    max_score = max(scores)

    if max_score == min_score:
        return [0.5] * len(scores)

    return [(s - min_score) / (max_score - min_score) for s in scores]
