# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from __future__ import annotations

from dataclasses import dataclass, field
from enum import StrEnum
from typing import Protocol, TYPE_CHECKING

from renderers import Message

from torchtitan.rl.observability import metrics as m
from torchtitan.rl.types import Completion, RolloutTurnID

if TYPE_CHECKING:
    # Type-only: importing the generator module here would pull in vLLM at import time.
    from torchtitan.rl.generator import SamplingConfig


_TRUNCATED = frozenset(
    {"truncated_length", "truncated_prompt_too_long", "truncated_max_turns"}
)
_ERROR = frozenset({"error_parse", "error_timeout", "error_abort", "error"})


class GenerateFn(Protocol):
    """Generate one model completion for a prompt.

    The Rollouter calls this once per turn and gets back a `Completion`. It does not need to know
    how or where generation runs. It can be a monarch actor, a router, an http endpoint, etc.
    """

    async def __call__(
        self,
        prompt_token_ids: list[int],
        *,
        request_id: str,
        group_id: int,
        routing_session_id: str | None = None,
        sampling_config: SamplingConfig | None = None,
    ) -> Completion | None:
        """Run one generation.

        Args:
            prompt_token_ids: The tokenized prompt to generate from.
            request_id: Unique per call; identifies the exact turn (e.g. ".../turn=2") in logs.
            group_id: Rollout group this call belongs to. Calls of one group share a prefix
                cache namespace, so siblings and later turns can reuse the group's KV.
            routing_session_id: Optional stable key for the routing session this call
                belongs to. A router may use it for session affinity, routing same-key
                calls to the same generator when possible. `None` means no affinity.
            sampling_config: Optional per-call sampling overrides.

        Returns:
            The generated `Completion`, or `None` if none is produced.
        """


class RolloutStatus(StrEnum):
    """Per-rollout status."""

    ONGOING = "ongoing"
    COMPLETED = "completed"
    TRUNCATED_LENGTH = "truncated_length"
    TRUNCATED_PROMPT_TOO_LONG = "truncated_prompt_too_long"
    TRUNCATED_MAX_TURNS = "truncated_max_turns"
    ERROR_PARSE = "error_parse"
    ERROR_TIMEOUT = "error_timeout"
    ERROR_ABORT = "error_abort"
    ERROR = "error"

    def is_truncated(self) -> bool:
        return self.value in _TRUNCATED

    def is_error(self) -> bool:
        return self.value in _ERROR

    def is_terminal(self) -> bool:
        return self is not RolloutStatus.ONGOING


@dataclass(kw_only=True, slots=True)
class RolloutTurn:
    """Full per-turn snapshot: the prompt fed to the generator + the sampled completion +
    the env's reply, in both token and message space. Rubrics score it and
    `rollout_to_training_samples` packs the rollout into training tokens."""

    # TODO: add a `logs` field (raw prompt/response text, finish_reason, timings)
    # so a turn can be dumped and inspected without re-deriving from tokens.

    rollout_id: RolloutTurnID
    """Identifies this turn (group, sibling index, turn index)."""

    # Fields needed for training
    prompt_token_ids: list[int]  # [num_prompt_tokens]
    """Tokenized conversation up to this turn, used to generate this turn's completion."""

    completion_token_ids: list[int]  # [num_completion_tokens]
    """This turn's completion token ids."""

    completion_logprobs: list[float]  # [num_completion_tokens]
    """This turn's completion token logprobs generated by the generator policy."""

    # Filtering
    min_policy_version: int | None = None
    """Oldest policy version this turn was sampled under; `None` if no generation happened."""

    max_policy_version: int | None = None
    """Newest policy version this turn was sampled under; `None` if no generation happened."""

    # Logging
    prompt_messages: list[Message] = field(
        default_factory=list
    )  # [num_prompt_messages]
    """Full conversation up to this turn; equivalent to `prompt_token_ids`."""

    completion_message: Message | None = None
    """This turn's completion decoded into a message by the renderer (the TokenEnv's parse,
    not the raw generator output)."""

    env_messages: list[Message] = field(default_factory=list)  # [num_env_messages]
    """This turn env's reply messages (tool / user)."""

    # For rubrics
    env_rewards: dict[str, float] = field(default_factory=dict)
    """This turn optional reward signals the env attached; the rubric decides how to use them."""

    metrics: list[m.Metric] = field(default_factory=list)
    """Per-turn metrics produced during rollouts"""


@dataclass(kw_only=True, slots=True)
class Rollout:
    """A complete rollout: ordered turns + terminal state + reward + identifier."""

    # TODO: add a `logs` field (per-turn debug records / event trace) to make a
    # full rollout reconstructable for debugging.

    group_id: int
    """Prompt-group ID; siblings share it for advantage centering."""

    rollout_id: int
    """Sibling index within the group (0..group_size-1)."""

    turns: list[RolloutTurn] = field(default_factory=list)  # [num_turns]
    """Ordered rollout turns. Each turn stores its full prompt (redundant across turns); kept so a
    rollout can be replayed/branched and divergences found, then collapsed at training_sample assembly."""
    # TODO: represent shared history as graph nodes so branching does not repeat
    # complete prompt prefixes and consume O(num_turns**2) token storage.

    status: RolloutStatus
    """Rollout-level terminal status."""

    reward: float | None = None
    """Final weighted reward, filled by the rubric."""

    reward_breakdown: dict[str, float] = field(default_factory=dict)
    """Raw per-reward-function values, filled by the rubric."""

    # TODO: make it per token
    advantage: float | None = None
    """Advantage for this sample."""


@dataclass(kw_only=True, slots=True)
class RolloutGroup:
    group_id: int
    """Prompt-group ID; siblings share it for advantage centering."""

    rollouts: list[Rollout]  # [group_size]
    """Sibling rollouts sampled from the group's shared prompt. Empty on a failed generation."""

    metrics: list[m.Metric] = field(default_factory=list)
    """Rollout-origin metrics that ride with this group to the trainer (a failed group carries its failure metric)."""
