"""LLM 호출 DB 로거 — 모든 LLM 호출의 입력/출력을 자동 기록."""

import json
import logging
import uuid
from datetime import datetime, timezone
from typing import Optional

logger = logging.getLogger(__name__)


def log_llm_call(
    model_name: str,
    user_prompt: str,
    status: str,
    system_prompt: Optional[str] = None,
    output_text: Optional[str] = None,
    duration_ms: int = 0,
    project_id: Optional[str] = None,
    episode_id: Optional[str] = None,
    operation_type: Optional[str] = None,
    step_name: Optional[str] = None,
    reference_image_ids: Optional[list[str]] = None,
    input_tokens: Optional[int] = None,
    output_tokens: Optional[int] = None,
    error_message: Optional[str] = None,
    metadata: Optional[dict] = None,
) -> str:
    """LLM 호출 기록을 DB에 저장. Returns log ID. 실패해도 메인 흐름 중단 안 함.

    Args:
        metadata (Phase 4 iter 7 I1): free-form trace context — scene_index /
            shot_index / still_id / entity_id 등. column 추가 비용을 줄이기 위해
            JSON serialize 후 단일 TEXT column (`metadata_json`) 에 저장.
    """
    log_id = str(uuid.uuid4())
    now = datetime.now(timezone.utc).isoformat()

    MAX_PROMPT_CHARS = 10000
    if system_prompt and len(system_prompt) > MAX_PROMPT_CHARS:
        system_prompt = system_prompt[:MAX_PROMPT_CHARS] + "...[truncated]"
    if len(user_prompt) > MAX_PROMPT_CHARS:
        user_prompt = user_prompt[:MAX_PROMPT_CHARS] + "...[truncated]"
    if output_text and len(output_text) > MAX_PROMPT_CHARS:
        output_text = output_text[:MAX_PROMPT_CHARS] + "...[truncated]"

    try:
        from app.core.database import SessionLocal
        from app.models.project import LLMCallLog

        db = SessionLocal()
        try:
            record = LLMCallLog(
                id=log_id,
                project_id=project_id,
                episode_id=episode_id,
                operation_type=operation_type,
                step_name=step_name,
                model_name=model_name,
                system_prompt=system_prompt,
                user_prompt=user_prompt,
                output_text=output_text,
                reference_image_ids=json.dumps(reference_image_ids or []),
                duration_ms=duration_ms,
                input_tokens=input_tokens,
                output_tokens=output_tokens,
                status=status,
                error_message=error_message,
                metadata_json=(
                    json.dumps(metadata, ensure_ascii=False, default=str)
                    if metadata else None
                ),
                created_at=now,
            )
            db.add(record)
            db.commit()
        finally:
            db.close()
    except Exception as exc:
        logger.warning("LLM call log failed (non-fatal): %s", exc)

    return log_id
