"""프롬프트 트레이서 — 모든 LLM 호출의 입출력을 JSON으로 기록.

각 호출마다:
- timestamp
- stage (turn0, turn1, scene_detail, prompt_translation, image_gen 등)
- model
- input (system prompt + user prompt)
- output (응답)
- metadata (scene_index, entity_name 등)
"""

import json
import logging
import os
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, Optional

logger = logging.getLogger(__name__)


class PromptTracer:
    """LLM 호출 추적기. JSON Lines 파일에 기록."""

    def __init__(self, trace_dir: str, episode_id: str = ""):
        self._trace_dir = Path(trace_dir)
        self._trace_dir.mkdir(parents=True, exist_ok=True)
        self._episode_id = episode_id
        ts = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S")
        self._trace_file = self._trace_dir / f"trace_{episode_id[:8]}_{ts}.jsonl"
        self._count = 0

    def log(
        self,
        stage: str,
        model: str,
        input_data: Dict[str, Any],
        output_data: Any,
        metadata: Optional[Dict[str, Any]] = None,
        elapsed_ms: int = 0,
    ) -> None:
        """LLM 호출 1건 기록."""
        record = {
            "timestamp": datetime.now(timezone.utc).isoformat(),
            "seq": self._count,
            "stage": stage,
            "model": model,
            "episode_id": self._episode_id,
            "input": input_data,
            "output": output_data if isinstance(output_data, (dict, list, str)) else str(output_data),
            "metadata": metadata or {},
            "elapsed_ms": elapsed_ms,
        }
        try:
            with open(self._trace_file, "a", encoding="utf-8") as f:
                f.write(json.dumps(record, ensure_ascii=False) + "\n")
            self._count += 1
        except Exception as exc:
            logger.warning("Trace write failed: %s", exc)

    def log_ref_matching(
        self,
        scene_index: int,
        t2i_prompt: str,
        matched_refs: list,
        final_prompt: str,
    ) -> None:
        """씬 이미지 생성 시 참조 매칭 + 최종 프롬프트 기록."""
        self.log(
            stage="scene_ref_matching",
            model="",
            input_data={
                "scene_index": scene_index,
                "original_t2i": t2i_prompt,
                "matched_references": [
                    {"index": i + 1, "label": label, "bytes": len(img_bytes)}
                    for i, (label, img_bytes) in enumerate(matched_refs)
                ],
            },
            output_data={"final_prompt": final_prompt},
        )

    @property
    def trace_file(self) -> str:
        return str(self._trace_file)
