from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session

from agent.chat_agent import ChatAgent
from database import get_db
from memory.session_store import SessionStore
from models.schemas import (
    MessageOut,
    PresetAskRequest,
    PresetAskResponse,
    PresetCreate,
    PresetOut,
    PresetUpdate,
)
from presets.preset_store import PresetStore
from presets.template_utils import extract_variables

router = APIRouter(prefix="/presets", tags=["presets"])


def _preset_to_out(preset) -> PresetOut:
    return PresetOut(
        id=preset.id,
        title=preset.title,
        prompt_template=preset.prompt_template,
        category=preset.category,
        model=preset.model,
        sort_order=preset.sort_order,
        created_at=preset.created_at,
        variables=extract_variables(preset.prompt_template),
    )


@router.get("", response_model=list[PresetOut])
def list_presets(
    category: str | None = Query(None),
    db: Session = Depends(get_db),
):
    store = PresetStore(db)
    return [_preset_to_out(p) for p in store.list_presets(category)]


@router.get("/{preset_id}", response_model=PresetOut)
def get_preset(preset_id: str, db: Session = Depends(get_db)):
    store = PresetStore(db)
    preset = store.get_preset(preset_id)
    if not preset:
        raise HTTPException(status_code=404, detail="预设不存在")
    return _preset_to_out(preset)


@router.post("", response_model=PresetOut, status_code=201)
def create_preset(body: PresetCreate, db: Session = Depends(get_db)):
    store = PresetStore(db)
    if store.get_preset(body.id):
        raise HTTPException(status_code=409, detail="预设 ID 已存在")
    preset = store.create_preset(body.model_dump())
    return _preset_to_out(preset)


@router.put("/{preset_id}", response_model=PresetOut)
def update_preset(
    preset_id: str, body: PresetUpdate, db: Session = Depends(get_db)
):
    store = PresetStore(db)
    preset = store.update_preset(
        preset_id, body.model_dump(exclude_unset=True)
    )
    if not preset:
        raise HTTPException(status_code=404, detail="预设不存在")
    return _preset_to_out(preset)


@router.delete("/{preset_id}")
def delete_preset(preset_id: str, db: Session = Depends(get_db)):
    store = PresetStore(db)
    if not store.delete_preset(preset_id):
        raise HTTPException(status_code=404, detail="预设不存在")
    return {"ok": True}


@router.post("/{preset_id}/ask", response_model=PresetAskResponse)
async def ask_preset(
    preset_id: str, body: PresetAskRequest, db: Session = Depends(get_db)
):
    preset_store = PresetStore(db)
    session_store = SessionStore(db)
    if not session_store.get_session(body.session_id):
        raise HTTPException(status_code=404, detail="会话不存在")

    try:
        preset, prompt = preset_store.render_prompt(preset_id, body.variables)
    except LookupError as exc:
        raise HTTPException(status_code=404, detail=str(exc)) from exc
    except ValueError as exc:
        raise HTTPException(status_code=400, detail=str(exc)) from exc

    agent = ChatAgent(db)
    try:
        user_id, assistant_id, _ = await agent.reply(
            body.session_id, prompt, model=preset.model
        )
    except RuntimeError as exc:
        raise HTTPException(status_code=500, detail=str(exc)) from exc
    except Exception as exc:
        raise HTTPException(status_code=502, detail=f"DeepSeek 请求失败: {exc}") from exc

    messages = session_store.get_messages(body.session_id)
    user_message = next(m for m in messages if m.id == user_id)
    assistant_message = next(m for m in messages if m.id == assistant_id)
    return PresetAskResponse(
        preset_id=preset_id,
        rendered_prompt=prompt,
        user_message=MessageOut.model_validate(user_message),
        assistant_message=MessageOut.model_validate(assistant_message),
    )
