# agent_gateway.py
# -*- coding: utf-8 -*-
from __future__ import annotations

import asyncio
import json
from typing import Any, Dict, List, Tuple

from flask import Flask, request, jsonify
from flask_cors import CORS
from fastmcp import Client
import sys
import threading
from datetime import datetime
import uuid

app = Flask(__name__)
CORS(app)

# 連到你的 orchestrator MCP (server_action_script.py)
ACTION_MCP_URL = "http://192.168.41.20:8081/stt"


# -------------------------
# Helpers
# -------------------------
def _safe_json(obj: Any, max_len: int = 1500) -> str:
    try:
        s = json.dumps(obj, ensure_ascii=False)
    except Exception:
        s = str(obj)
    return s[:max_len]


def _unwrap_tool_result(r: Any) -> Dict[str, Any]:
    """
    把 fastmcp 的回傳（可能是 dict / ToolResult / TextContent list）
    轉成「可 jsonify」的 dict。
    """
    if isinstance(r, dict):
        return r

    data = getattr(r, "data", None)
    if isinstance(data, dict):
        return data

    content = getattr(r, "content", None)
    if content:
        for c in content:
            if isinstance(c, dict):
                c_type = c.get("type")
                txt = c.get("text") if c_type == "text" else None
            else:
                c_type = getattr(c, "type", None)
                txt = getattr(c, "text", None) if c_type == "text" else None

            if txt:
                # 很多 MCP server 會把 JSON 用文字包回來
                try:
                    obj = json.loads(txt)
                    if isinstance(obj, dict):
                        return obj
                except Exception:
                    pass

    # 最後保底：轉字串，避免 "not JSON serializable"
    return {"_raw": str(r)}

# -------------------------
# HTML LOG
# 全域 job log（你前面已經有類似的可以沿用）
JOB_LOGS = {}

def push_log(job_id: str, msg: str, level="INFO"):
    JOB_LOGS.setdefault(job_id, []).append({
        "ts": datetime.now().strftime("%H:%M:%S"),
        "level": level,
        "msg": msg,
    })
    # 保護一下，不然跑久會爆
    JOB_LOGS[job_id] = JOB_LOGS[job_id][-1000:]


class StdoutTee:
    """
    同時：
    1) 印到原本 terminal
    2) 收集到 UI system log
    """
    def __init__(self, job_id: str):
        self.job_id = job_id
        self._stdout = sys.__stdout__   # 原本的 terminal
        self._buf = ""
        self._lock = threading.Lock()

    def write(self, s: str):
        with self._lock:
            # 1️⃣ 原樣印到 terminal
            self._stdout.write(s)
            self._stdout.flush()

            # 2️⃣ 收集給 UI
            self._buf += s
            while "\n" in self._buf:
                line, self._buf = self._buf.split("\n", 1)
                line = line.rstrip()
                if line:
                    push_log(self.job_id, line)

    def flush(self):
        with self._lock:
            if self._buf.strip():
                push_log(self.job_id, self._buf.strip())
                self._buf = ""
            self._stdout.flush()

# -------------------------

def _extract_debug(action_result: Dict[str, Any]) -> Tuple[List[str], Dict[str, Any]]:
    """
    從 server_action_script.run_agent() 的回傳，把 logs/steps/results 抽出來，
    讓前端能顯示「跑到哪」「planner 怎麼排的」。
    """
    debug_logs: List[str] = []

    # action_result 正常會長這樣：
    # {"ok": True, "logs": [...], "results": [{"tool":..., "result":...}]}
    if isinstance(action_result, dict):
        if "logs" in action_result and isinstance(action_result["logs"], list):
            debug_logs.extend([str(x) for x in action_result["logs"]])

        # 額外把 steps / results 的摘要加進 log（前端看得懂）
        results = action_result.get("results", [])
        if isinstance(results, list):
            debug_logs.append(f"results_count={len(results)}")
            for i, rr in enumerate(results, start=1):
                # ✅ 新格式：{step,target,tool,result:{ok,...}}
                target = rr.get("target")
                tool = rr.get("tool") or rr.get("name")  # 保底
                step_no = rr.get("step") or i

                ok = None
                r = rr.get("result")
                if isinstance(r, dict):
                    ok = r.get("ok")
                    # 如果 dispatch 回 {"ok": True, "result": {...}}，也抓一下內層 ok
                    inner = r.get("result")
                    if ok is True and isinstance(inner, dict) and "ok" in inner:
                        ok = inner.get("ok")

                if target:
                    debug_logs.append(f"[{step_no}] {target}.{tool} ok={ok}")
                else:
                    # ✅ 舊格式 fallback
                    debug_logs.append(f"[{step_no}] tool={tool} ok={ok}")

    return debug_logs, action_result


async def _call_action_run_agent(payload: Dict[str, Any]) -> Dict[str, Any]:
    """
    呼叫 orchestrator MCP：run_agent
    """
    filename = payload.get("filename")
    b64 = payload.get("b64")
    language = payload.get("language", "zh")
    llm_refine = bool(payload.get("llm_refine", True))
    user_goal = payload.get("user_goal", "將音檔做語音轉文字並進行字典/LLM refine")
    planner_model = payload.get("planner_model", "gpt-oss:120b-cloud")
    stt_mcp_url = payload.get("stt_mcp_url")  # 可選：不給就用 server_action_script 的預設
    project_id = payload.get("project_id") or ""
    collection = payload.get("collection") or "audio_v1"
    comment = payload.get("comment") or ""
    industry = payload.get("industry") or ""
    vendor_name = payload.get("vendor_name") or ""
    vendor_id = payload.get("vendor_id") or vendor_name
    flag = payload.get("flag", 1)
    selected_tools = payload.get("selected_tools")
    upload_time = payload.get("upload_time")
    job_id = payload.get("job_id")
    video_filename = payload.get("video_filename")
    video_b64 = payload.get("video_b64")
    filetype = payload.get("filetype")

    video_path = None
    doc_filename = payload.get("doc_filename")
    doc_b64 = payload.get("doc_b64")
    pic_filename = payload.get("pic_filename")
    pic_b64 = payload.get("pic_b64")
    claude_filename = payload.get("claude_filename")
    claude_b64 = payload.get("claude_b64")

    print("[gateway] -> call MCP run_agent")
    print(f"[gateway] ACTION_MCP_URL={ACTION_MCP_URL}")
    print(f"[gateway] filename={filename} b64_len={(len(b64) if b64 else 0)} language={language} llm_refine={llm_refine}")
    print(f"[gateway] user_goal={user_goal} planner_model={planner_model} stt_mcp_url={stt_mcp_url}")
    print(f"[gateway] project_id={project_id} collection={collection} flag={flag} comment={comment}")
    print(f"[gateway] industry={industry} vendor_name={vendor_name} vendor_id={vendor_id}")
    print(f"[gateway] upload_time={upload_time} job_id={job_id}")
    print(f"[gateway] video_filename={video_filename} has_video_b64={bool(video_b64)}")
    print(f"[gateway] filetype={filetype}")
    print(f"[gateway] doc_filename={doc_filename} has_doc_b64={bool(doc_b64)}")
    print(f"[gateway] claude_filename={claude_filename} has_claude_b64={bool(claude_b64)}")

    async with Client(ACTION_MCP_URL) as client:
        args = {
            "filename": filename,
            "b64": b64,
            "language": language,
            "llm_refine": llm_refine,
            "user_goal": user_goal,
            "planner_model": planner_model,
            "project_id": project_id,
            "collection": collection,
            "flag": flag,
            "comment": comment,
            "industry": industry,
            "vendor_name": vendor_name,
            "vendor_id": vendor_id,
            "selected_tools": selected_tools,
            "job_id": job_id,
            "upload_time": upload_time,
            "filetype": filetype,
            "doc_filename": doc_filename,
            "doc_b64": doc_b64,
            "pic_filename": pic_filename,
            "pic_b64": pic_b64,
            "claude_filename": claude_filename,
            "claude_b64": claude_b64,
        }
        if stt_mcp_url:
            args["stt_mcp_url"] = stt_mcp_url
        if video_filename:
            args["video_filename"] = video_filename
        if video_b64:
            args["video_b64"] = video_b64
        if video_path:
            args["video_path"] = video_path

        result = await client.call_tool("run_agent", args)

    print("[gateway] <- MCP raw_type=", type(result))
    return _unwrap_tool_result(result)


# -------------------------
# Routes
# -------------------------
@app.get("/health")
def health():
    return jsonify({"ok": True, "service": "agent_gateway", "action_mcp_url": ACTION_MCP_URL})

@app.route("/agent/logs", methods=["GET", "POST", "OPTIONS"])
def agent_logs():
    """
    UI 輪詢用：
      GET  /agent/logs?job_id=xxx&from_idx=0
      POST /agent/logs?job_id=xxx&from_idx=0   (也支援，避免你前端不小心用 POST)
    轉呼 orchestrator MCP tool: get_job_logs
    """
    job_id = request.args.get("job_id") or (request.get_json(silent=True) or {}).get("job_id")
    from_idx = request.args.get("from_idx") or (request.get_json(silent=True) or {}).get("from_idx") or 0

    if not job_id:
        return jsonify({"ok": False, "error": "missing job_id"}), 400

    try:
        from_idx = int(from_idx)
    except Exception:
        from_idx = 0

    async def _call_get_logs():
        async with Client(ACTION_MCP_URL) as client:
            r = await client.call_tool("get_job_logs", {"job_id": job_id, "from_idx": from_idx})
            return _unwrap_tool_result(r)

    try:
        data = asyncio.run(_call_get_logs())
        # data 會是 {ok, job_id, from_idx, next_idx, lines}
        return jsonify(data)
    except Exception as e:
        return jsonify({"ok": False, "error": str(e)}), 500


@app.post("/agent/start")
def agent_start():
    """
    前端 POST JSON：
    {
      "filename": "...",
      "b64": "...",
      "language": "zh",
      "llm_refine": true,
      "user_goal": "...",          (optional)
      "planner_model": "...",      (optional)
      "stt_mcp_url": "http://..."  (optional)
    }
    """
    payload = request.get_json(force=True, silent=True) or {}

    #if not payload.get("filename") and not payload.get("b64") and not payload.get("query_text"):
        #return jsonify({"ok": False, "error": "need filename & b64, or query_text"}), 400

    selected_tools = payload.get("selected_tools") or []
    job_id = payload.get("job_id") or f"job_{uuid.uuid4().hex[:12]}"
    payload["job_id"] = job_id

    # ✅ 只有勾了 stt.upload_and_stt 才需要音檔
    need_audio = ("stt.upload_and_stt" in selected_tools)

    filename = payload.get("filename")
    b64 = payload.get("b64")
    has_audio = bool(filename and b64)

    # ✅ 入口不要被訂死：沒勾 upload_and_stt 就不要求音檔、不要求 query_text
    if need_audio and (not has_audio):
        return jsonify({"ok": False, "error": "stt.upload_and_stt selected -> need filename & b64"}), 400
    
    need_doc = ("docparse.parse_file" in selected_tools)
    doc_filename = payload.get("doc_filename")
    doc_b64 = payload.get("doc_b64")
    has_doc = bool(doc_filename and doc_b64)

    if need_doc and not has_doc:
        return jsonify({
            "ok": False,
            "error": "docparse.parse_file selected -> need doc_filename & doc_b64"
        }), 400
    
    need_pic = ("picparser.parse_image" in selected_tools)
    pic_filename = payload.get("pic_filename")
    pic_b64 = payload.get("pic_b64")
    has_pic = bool(pic_filename and pic_b64)

    if need_pic and not has_pic:
        return jsonify({
            "ok": False,
            "error": "picparser.parse_image selected -> need pic_filename & pic_b64"
        }), 400

    need_claude_md = ("claude_doc.update_claudemd" in selected_tools)
    claude_filename = payload.get("claude_filename")
    claude_b64 = payload.get("claude_b64")
    has_claude_md = bool(claude_filename and claude_b64)

    if need_claude_md and not has_claude_md:
        return jsonify({
            "ok": False,
            "error": "claude_doc.update_claudemd selected -> need claude_filename & claude_b64"
        }), 400

    # 這個 logs 會讓前端「即使 action 失敗」也知道卡在哪
    gateway_logs: List[str] = []
    gateway_logs.append("啟動 agent_gateway")
    gateway_logs.append("送出 run_agent 到 action_script MCP")

    try:
        # ✅ 注意：Flask sync handler 用 asyncio.run OK（你現在就是這樣用）
        old_stdout = sys.stdout
        sys.stdout = StdoutTee(job_id)
        try:
            action_result = asyncio.run(_call_action_run_agent(payload))
        finally:
            sys.stdout = old_stdout    

        debug_logs, safe_result = _extract_debug(action_result)

        # 後端 terminal 可見
        print("[gateway] action_result=", _safe_json(safe_result, max_len=4000))

        gateway_logs.append("action_script MCP 完成回傳")
        captured = [x["msg"] for x in JOB_LOGS.get(job_id, [])]

        # 回給前端：logs 給 UI 顯示、result 放完整結果
        return jsonify({
            "ok": True,
            "logs": captured + debug_logs,
            "result": safe_result,
        })

    except Exception as e:
        print("[gateway] ERROR:", repr(e))
        gateway_logs.append(f"失敗: {e}")
        return jsonify({"ok": False, "logs": gateway_logs, "error": str(e)}), 500


if __name__ == "__main__":
    # 給瀏覽器打的 port
    app.run(host="0.0.0.0", port=8082, debug=False)
