import asyncio
import json
import os
import time
import uuid
from datetime import datetime, timezone
from typing import Any, Callable, Dict, List, Optional

from fastapi import APIRouter, WebSocket, WebSocketDisconnect
from openai import AsyncOpenAI
from pydantic import BaseModel, ConfigDict
import tiktoken
from langsmith import traceable
from langsmith.wrappers import wrap_openai
router = APIRouter()

# ---------------------------------------------------------------------------
# Models
# ---------------------------------------------------------------------------

class HistoryItem(BaseModel):
    title: str
    text: str

class PendingCall(BaseModel):
    id: str
    call_id: str
    name: str
    arguments: str

class Session(BaseModel):
    model_config = ConfigDict(arbitrary_types_allowed=True)

    messages: List[Dict[str, Any]] = []
    history: Dict[str, HistoryItem] = {}
    cancel_requested: bool = False
    websocket: Optional[WebSocket] = None
    last_active: float = 0.0
    output_tokens_used: int = 0

# ---------------------------------------------------------------------------
# Session state — in-memory, ephemeral (no persistence by design)
# ---------------------------------------------------------------------------
sessions: Dict[str, Session] = {}
websocket_to_session: Dict[WebSocket, str] = {}

OPENAI_API_KEY = os.getenv("OPENAI_API_KEY", "")
print(f"OPENAI_API_KEY: {OPENAI_API_KEY}")
client = wrap_openai(AsyncOpenAI(api_key=OPENAI_API_KEY)) if OPENAI_API_KEY else None
tiktoken_encoding = tiktoken.get_encoding("cl100k_base")

USAGE_UPDATE_FREQUENCY = 5

_PROMPT_PATH = os.path.join(os.path.dirname(__file__), "..", "system_prompt.md")
try:
    with open(_PROMPT_PATH, encoding="utf-8") as _f:
        SYSTEM_PROMPT = _f.read().strip()
except FileNotFoundError:
    SYSTEM_PROMPT = (
        "You are an AI assistant on Ross Klein's portfolio website. "
        "Answer questions about Ross's background, skills, experience, and projects conversationally. "
        "Be concise and friendly."
    )
SYSTEM_PROMPT = f"Today's date is {datetime.now().strftime('%Y-%m-%d')}." + SYSTEM_PROMPT

NOTES_DIR = os.getenv(
    "NOTES_DIR",
    os.path.join(os.path.dirname(__file__), "..", "notes"),
)

try:
    _notes_lines: list[str] = []
    for _section in sorted(os.listdir(NOTES_DIR)):
        _section_dir = os.path.join(NOTES_DIR, _section)
        if not os.path.isdir(_section_dir):
            continue
        _section_files = sorted(
            f for f in os.listdir(_section_dir) if f.endswith(".md")
        )
        _main = f"{_section}.md"
        _subs = [f for f in _section_files if f != _main]
        if _main in _section_files:
            _notes_lines.append(f"{_main} -")
            for _sub in _subs:
                _notes_lines.append(f"    {_sub}")
            if not _subs:
                _notes_lines.append("")
    if _notes_lines:
        SYSTEM_PROMPT += "\n\n# Available Notes\n\n" + "\n".join(_notes_lines)
except FileNotFoundError:
    pass

print(f"SYSTEM_PROMPT: {SYSTEM_PROMPT}")
TOOLS = [
    {
        "type": "function",
        "name": "read_notes",
        "description": "Read the contents of one or more note files about Ross Klein.",
        "parameters": {
            "type": "object",
            "properties": {
                "filenames": {
                    "type": "array",
                    "items": {"type": "string"},
                    "description": (
                        "Note filenames to read (without .md extension), "
                        "e.g. ['experience', 'skills']"
                    ),
                }
            },
            "required": ["filenames"],
        },
    },
]

MAX_TOOL_ITERATIONS = 50
SESSION_OUTPUT_TOKEN_LIMIT = 50_000


def _note_path(name: str) -> str:
    """Resolve a note filename (without .md) to its full path.

    Files live under notes/{section}/{name}.md where the section is the
    first underscore-delimited token of the filename.
    e.g. 'experience'              → notes/experience/experience.md
         'experience_amazon'       → notes/experience/experience_amazon.md
         'projects_sheets_agent'   → notes/projects/projects_sheets_agent.md
    """
    stem = name[:-3] if name.endswith(".md") else name
    section = stem.split("_")[0]
    return os.path.join(NOTES_DIR, section, stem + ".md")


async def _read_notes(args: Dict[str, Any]) -> Dict[str, Any]:
    filenames = args.get("filenames", [])
    results: Dict[str, str] = {}
    for name in filenames:
        stem = name[:-3] if name.endswith(".md") else name
        path = _note_path(stem)
        try:
            with open(path, encoding="utf-8") as fh:
                results[stem] = fh.read()
        except FileNotFoundError:
            results[stem] = f"[Note not found: {path}]"
        except Exception as exc:
            results[stem] = f"[Error reading {stem}: {exc}]"
    return results


HANDLER_REGISTRY: Dict[str, Callable] = {
    "read_notes": _read_notes,
}

# ---------------------------------------------------------------------------
# Session helpers
# ---------------------------------------------------------------------------

SESSION_TTL_SECONDS = 3600  # 1 hour


def _new_token() -> str:
    return str(uuid.uuid4())

def _get_token(ws: WebSocket) -> Optional[str]:
    return websocket_to_session.get(ws)

def _get_session(token: str) -> Optional[Session]:
    return sessions.get(token)

def _sweep_sessions() -> None:
    cutoff = time.time() - SESSION_TTL_SECONDS
    expired = [t for t, s in sessions.items() if s.last_active < cutoff]
    for t in expired:
        sessions.pop(t, None)
    if expired:
        print(f"[chat] swept {len(expired)} expired session(s)")

# ---------------------------------------------------------------------------
# Token counting
# ---------------------------------------------------------------------------

def _count_tokens(messages: List[Dict[str, Any]]) -> int:
    total = 0
    for msg in messages:
        for field in ("role", "content"):
            val = msg.get(field, "")
            if val:
                total += len(tiktoken_encoding.encode(str(val)))
        total += 4
    return total

def _count_text_tokens(text: str) -> int:
    return len(tiktoken_encoding.encode(str(text))) if text else 0

# ---------------------------------------------------------------------------
# Usage stats
# ---------------------------------------------------------------------------

def _build_usage(input_tok: int, output_tok: int, accurate_total: int,
                 has_accurate: bool, elapsed: float) -> Dict[str, Any]:
    total = accurate_total if has_accurate and accurate_total > 0 else input_tok + output_tok
    tps = output_tok / elapsed if elapsed > 0 else 0.0
    return {
        "input_tokens": input_tok,
        "output_tokens": output_tok,
        "total_tokens": total,
        "tokens_per_second": round(tps, 2),
        "reasoning_tokens": None,
    }

async def _send_usage(ws: WebSocket, usage: Dict[str, Any], phase: str = "update") -> None:
    token = _get_token(ws)
    if not token:
        return
    session = _get_session(token)
    if not session or not session.websocket:
        return
    try:
        await session.websocket.send_json({
            "type": "usage",
            "phase": phase,
            "usage": usage,
            "timestamp": datetime.now(timezone.utc).isoformat(),
        })
    except Exception:
        pass

# ---------------------------------------------------------------------------
# Chat history sender
# ---------------------------------------------------------------------------

async def _send_chat(ws: WebSocket, item_id: str, item: HistoryItem) -> None:
    token = _get_token(ws)
    if not token:
        return
    session = _get_session(token)
    if not session or not session.websocket:
        return
    try:
        await session.websocket.send_json({
            "type": "chat",
            "item_id": item_id,
            "title": item.title,
            "text": item.text,
            "role": item.title,
            "timestamp": datetime.now(timezone.utc).isoformat(),
        })
    except Exception:
        pass

# ---------------------------------------------------------------------------
# Context builder — reads full history, returns trimmed view for the model.
# The raw history list is never modified here.
#
# Rules:
#   - System prompt is always included.
#   - Regular messages: keep from the (MAX_TURNS)th-from-last user message onward.
#     "Turn" = one user message, so MAX_TURNS=100 keeps the last 100 exchanges.
#   - Tool items (function_call / function_call_output): included only when
#     fewer than TOOL_TURNS user messages have occurred after them.
# ---------------------------------------------------------------------------

MAX_TURNS = 100
TOOL_TURNS = 10


def _build_context(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
    if not messages:
        return messages

    system: List[Dict[str, Any]] = []
    rest = messages[:]
    if rest and rest[0].get("role") == "system":
        system = [rest[0]]
        rest = rest[1:]

    user_indices = [i for i, m in enumerate(rest) if m.get("role") == "user"]
    regular_cutoff = user_indices[-MAX_TURNS] if len(user_indices) > MAX_TURNS else 0

    result = system[:]
    for i, msg in enumerate(rest):
        is_tool = msg.get("type") in ("function_call", "function_call_output")
        if is_tool:
            users_after = sum(1 for j in range(i + 1, len(rest)) if rest[j].get("role") == "user")
            if users_after < TOOL_TURNS:
                result.append(msg)
        else:
            if i >= regular_cutoff:
                result.append(msg)

    return result

# ---------------------------------------------------------------------------
# LLM streaming loop
# ---------------------------------------------------------------------------
async def _run_llm(ws: WebSocket, user_content: str) -> None:
    token = _get_token(ws)
    if not token:
        return
    session = _get_session(token)
    if not session:
        return

    conversation = session.messages
    history = session.history
    session.cancel_requested = False

    # Inject system prompt on first message
    if not conversation:
        conversation.append({"role": "system", "content": SYSTEM_PROMPT})

    conversation.append({"role": "user", "content": user_content})
    history[f"user-{int(time.time() * 1000)}"] = HistoryItem(title="user", text=user_content)

    if not client:
        if session.websocket:
            try:
                await session.websocket.send_json({
                    "type": "error",
                    "message": "OpenAI client not initialised. Set OPENAI_API_KEY.",
                })
            except Exception:
                pass
        return

    if session.output_tokens_used >= SESSION_OUTPUT_TOKEN_LIMIT:
        if session.websocket:
            try:
                await session.websocket.send_json({
                    "type": "error",
                    "message": "This session has reached its token limit. Please refresh to start a new session.",
                })
            except Exception:
                pass
        return

    try:
        input_tok = _count_tokens(conversation)
        output_tok = 0
        accurate_total = 0
        has_accurate = False
        start = time.time()
        chunk_counter = 0

        iteration_active = True
        tool_iterations = 0
        while iteration_active:
            segments: List[str] = []
            pending_calls: List[PendingCall] = []
            message_id = f"assistant-{int(time.time() * 1000)}"
            assistant_item = HistoryItem(title="assistant", text="")

            reasoning_entries: Dict[str, Dict[str, str]] = {}
            completed_reasoning_ids: List[str] = []

            def _ensure_reasoning(item_id: Any) -> str:
                nid = str(item_id) if item_id is not None else ""
                reasoning_entries.setdefault(nid, {"summary": "", "content": ""})
                if nid not in history:
                    history[nid] = HistoryItem(title="thinking", text="")
                history[nid].title = "thinking"
                return nid

            def _reasoning_text(nid: str) -> str:
                entry = reasoning_entries.get(nid, {})
                parts = []
                if entry.get("summary"):
                    parts.append(entry["summary"].strip())
                if entry.get("content"):
                    parts.append(entry["content"].strip())
                return "\n\n".join(parts)

            async def _push_reasoning(item_id: Any) -> None:
                nid = _ensure_reasoning(item_id)
                history[nid].text = _reasoning_text(nid)
                await _send_chat(ws, nid, history[nid])

            async with client.responses.stream(
                model="gpt-5.4",
                reasoning={"effort": "low", "summary": "auto"},
                input=_build_context(conversation),
                tools=TOOLS,
                store=False,
                truncation="auto",
            ) as stream:
                async for chunk in stream:
                    if session.cancel_requested:
                        if session.websocket:
                            try:
                                await session.websocket.send_json({
                                    "type": "chat",
                                    "item_id": message_id,
                                    "title": "assistant",
                                    "text": "".join(segments) + "\n\n[Cancelled]",
                                    "role": "assistant",
                                    "timestamp": datetime.now(timezone.utc).isoformat(),
                                })
                            except Exception:
                                pass
                        iteration_active = False
                        break

                    chunk_counter += 1

                    if chunk.type == "response.output_text.delta":
                        delta = getattr(chunk, "delta", "")
                        if delta:
                            segments.append(delta)
                            assistant_item.text = "".join(segments)
                            await _send_chat(ws, message_id, assistant_item)
                            output_tok += _count_text_tokens(delta)

                    if chunk.type == "response.output_text.done":
                        final = getattr(chunk, "text", "")
                        if final:
                            segments = [final]
                            assistant_item.text = final
                            await _send_chat(ws, message_id, assistant_item)

                    if chunk.type in ("response.reasoning_summary_part.added", "response.reasoning_summary_text.delta"):
                        if chunk.type == "response.reasoning_summary_part.added":
                            part = getattr(chunk, "part", None)
                            delta = "\n\n" + (getattr(part, "text", "") if part else "")
                        else:
                            delta = getattr(chunk, "delta", "")
                        if delta:
                            nid = _ensure_reasoning(getattr(chunk, "item_id", None))
                            reasoning_entries[nid]["summary"] += delta
                            output_tok += _count_text_tokens(delta)
                            await _push_reasoning(nid)

                    if chunk.type == "response.reasoning_summary_text.done":
                        nid = _ensure_reasoning(getattr(chunk, "item_id", None))
                        await _push_reasoning(nid)
                        entry = reasoning_entries.get(nid, {})
                        if nid and nid not in completed_reasoning_ids:
                            if entry.get("summary", "").strip() or entry.get("content", "").strip():
                                completed_reasoning_ids.append(nid)

                    if chunk.type == "response.output_item.done":
                        item = getattr(chunk, "item", None)
                        if item and getattr(item, "type", None) == "function_call":
                            pending_calls.append(PendingCall(
                                id=getattr(item, "id", ""),
                                call_id=getattr(item, "call_id", ""),
                                name=getattr(item, "name", ""),
                                arguments=getattr(item, "arguments", "{}"),
                            ))

                    if chunk.type == "response.completed":
                        resp_obj = getattr(chunk, "response", None)
                        usage = getattr(resp_obj, "usage", None) if resp_obj else None
                        if usage:
                            input_tok = getattr(usage, "input_tokens", input_tok)
                            output_tok = getattr(usage, "output_tokens", output_tok)
                            accurate_total = getattr(usage, "total_tokens", accurate_total)
                            has_accurate = bool(accurate_total)
                            session.output_tokens_used += output_tok

                    if chunk_counter % USAGE_UPDATE_FREQUENCY == 0:
                        elapsed = max(time.time() - start, 0.001)
                        await _send_usage(ws, _build_usage(input_tok, output_tok, accurate_total, has_accurate, elapsed))

            # Write assistant item into history now — after thinking entries — so
            # dict insertion order is: user → thinking → assistant
            history[message_id] = assistant_item

            # Add completed reasoning to conversation context
            for nid in completed_reasoning_ids:
                text = _reasoning_text(nid)
                if text.strip():
                    conversation.append({"role": "assistant", "content": f"<thinking>\n{text}\n</thinking>"})

            final_text = "".join(segments).strip()

            if pending_calls and tool_iterations < MAX_TOOL_ITERATIONS:
                # Add the model's function call items so they appear in the next request's context
                for call in pending_calls:
                    conversation.append({
                        "type": "function_call",
                        "id": call.id,
                        "call_id": call.call_id,
                        "name": call.name,
                        "arguments": call.arguments,
                    })
                # Execute each tool and append its result
                for call in pending_calls:
                    handler = HANDLER_REGISTRY.get(call.name)
                    try:
                        args = json.loads(call.arguments or "{}")
                    except json.JSONDecodeError:
                        args = {}
                    result = await handler(args) if handler else {"error": f"Unknown tool: {call.name}"}
                    conversation.append({
                        "type": "function_call_output",
                        "call_id": call.call_id,
                        "output": json.dumps(result),
                    })
                tool_iterations += 1
                # Loop again so the model can respond using the tool results
            else:
                if final_text:
                    conversation.append({"role": "assistant", "content": final_text})
                await _send_chat(ws, "complete", HistoryItem(title="assistant", text=""))
                iteration_active = False

        elapsed = max(time.time() - start, 0.001)
        await _send_usage(ws, _build_usage(input_tok, output_tok, accurate_total, has_accurate, elapsed), phase="final")

    except Exception as exc:
        if session and session.websocket:
            try:
                await session.websocket.send_json({"type": "error", "message": f"LLM error: {exc}"})
            except Exception:
                pass

# ---------------------------------------------------------------------------
# WebSocket endpoint
# ---------------------------------------------------------------------------

@router.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket) -> None:
    await websocket.accept()
    session_token: Optional[str] = None

    try:
        # First message: optional session_token for reconnection
        raw = await websocket.receive_text()
        initial: Optional[Dict[str, Any]] = None
        try:
            initial = json.loads(raw)
            session_token = initial.get("session_token") if initial else None
        except json.JSONDecodeError:
            pass

        _sweep_sessions()

        if not session_token or session_token not in sessions:
            session_token = _new_token()
            sessions[session_token] = Session(last_active=time.time())
        else:
            sessions[session_token].last_active = time.time()

        sessions[session_token].websocket = websocket
        websocket_to_session[websocket] = session_token

        await websocket.send_json({"type": "session", "session_token": session_token})

        # Handle initial message if it contained a chat payload
        if initial and isinstance(initial, dict) and initial.get("type") == "chat":
            await _handle_message(websocket, initial)

        while True:
            data = await websocket.receive_text()
            try:
                msg = json.loads(data)
                await _handle_message(websocket, msg)
            except json.JSONDecodeError as exc:
                await websocket.send_json({"type": "error", "message": f"Invalid JSON: {exc}"})

    except WebSocketDisconnect:
        pass
    finally:
        if session_token and session_token in sessions:
            sessions[session_token].websocket = None
        websocket_to_session.pop(websocket, None)


async def _handle_message(ws: WebSocket, msg: Dict[str, Any]) -> None:
    token = _get_token(ws)
    if not token:
        await ws.send_json({"type": "error", "message": "No session. Please reconnect."})
        return
    session = _get_session(token)
    if not session:
        await ws.send_json({"type": "error", "message": "Session not found. Please reconnect."})
        return

    session.last_active = time.time()
    msg_type = msg.get("type")

    if msg_type == "chat":
        content = msg.get("content", "")
        asyncio.create_task(_run_llm(ws, content))

    elif msg_type == "stop":
        session.cancel_requested = True
        if session.websocket:
            try:
                await session.websocket.send_json({"type": "stopped", "message": "Cancelled"})
            except Exception:
                pass

    else:
        await ws.send_json({"type": "error", "message": f"Unknown message type: {msg_type}"})

# ---------------------------------------------------------------------------
# REST helpers
# ---------------------------------------------------------------------------

@router.get("/health")
async def health():
    return {"status": "ok", "sessions": len(sessions)}


@router.get("/history")
async def get_history(session_token: str):
    session = _get_session(session_token)
    if not session:
        return {"items": []}
    items = [
        {
            "id": item_id,
            "role": item.title,
            "text": item.text,
            "isThinking": item.title == "thinking",
        }
        for item_id, item in session.history.items()
    ]
    return {"items": items}


@router.post("/clear-history")
async def clear_history(body: dict):
    token = body.get("session_token")
    if not token or token not in sessions:
        return {"status": "error", "message": "Session not found"}
    sessions[token].messages.clear()
    sessions[token].history.clear()
    return {"status": "ok", "message": "History cleared"}
