# server_agent_script.py
# -*- coding: utf-8 -*-
"""
✨ 簡化版 Agent Orchestrator
- 配置驅動：所有 MCP Server 定義在 tools_config.json
- 易擴展：添加新 Server 只需編輯 JSON
- 少 bug：邏輯清晰，參數自動填充
"""
from __future__ import annotations

from typing import Any, Dict, List
import json
import os
from pathlib import Path
import subprocess
import time

from fastmcp import FastMCP, Client
from ollama_client import ask_ollama_json, unload_ollama_model

mcp = FastMCP("action_script_mcp")

# =========================================
# 📋 配置加載
# =========================================
BASE_DIR = Path(__file__).resolve().parent          # /home/imtc/mcp_stack/mcp_main
STACK_ROOT = BASE_DIR.parent                        # /home/imtc/mcp_stack
DATA_JOBS = STACK_ROOT / "data_jobs"

CONFIG_PATH = Path(__file__).parent / "tools_config.json"
DEFAULT_PLANNER_MODEL = "gpt-oss:120b-cloud"
LOCAL_GPU_MODEL = "gpt-oss:120b"

def _load_config() -> Dict[str, Any]:
    """載入工具配置"""
    try:
        with open(CONFIG_PATH, "r", encoding="utf-8") as f:
            config = json.load(f)
        return config.get("servers", {})
    except Exception as e:
        print(f"⚠️ 無法載入配置: {e}，使用空配置")
        return {}
    
# ============ Job log store (for UI polling) ============
JOB_LOGS: dict[str, list[str]] = {}

def job_log(job_id: str | None, msg: str):
    """append log line for a job_id"""
    if not job_id:
        return
    JOB_LOGS.setdefault(job_id, []).append(str(msg))
    JOB_LOGS[job_id] = JOB_LOGS[job_id][-2000:]  # avoid memory blow
    

SERVERS = _load_config()


# =========================================
# 🛠️ 工具函數
# =========================================
def _safe_json(obj: Any, max_len: int = 2000) -> str:
    """安全轉 JSON"""
    try:
        return json.dumps(obj, ensure_ascii=False)[:max_len]
    except:
        return str(obj)[:max_len]


def _env_bool(name: str, default: str = "1") -> bool:
    return str(os.environ.get(name, default)).strip().lower() not in ("0", "false", "no", "off")


def _gpu_free_mb() -> int | None:
    try:
        out = subprocess.check_output(
            ["nvidia-smi", "--query-gpu=memory.free", "--format=csv,noheader,nounits"],
            text=True,
            stderr=subprocess.DEVNULL,
            timeout=5,
        )
        values = [int(x.strip()) for x in out.splitlines() if x.strip()]
        return max(values) if values else None
    except Exception:
        return None


def _wait_gpu_free_for_local_model(model: str, logline) -> None:
    if model != LOCAL_GPU_MODEL or not _env_bool("GPU_WAIT_BEFORE_STT", "1"):
        return

    min_free_mb = int(os.environ.get("GPU_WAIT_FREE_MB", "10000"))
    timeout_sec = int(os.environ.get("GPU_WAIT_TIMEOUT_SEC", "180"))
    interval_sec = float(os.environ.get("GPU_WAIT_INTERVAL_SEC", "2"))
    deadline = time.time() + timeout_sec

    while True:
        free_mb = _gpu_free_mb()
        if free_mb is None:
            logline("[gpu] nvidia-smi unavailable; skip GPU wait")
            return
        if free_mb >= min_free_mb:
            logline(f"[gpu] free={free_mb} MiB >= {min_free_mb} MiB; continue")
            return
        if time.time() >= deadline:
            raise RuntimeError(
                f"GPU memory still busy after unloading {model}. "
                f"Required free >= {min_free_mb} MiB, current free = {free_mb} MiB."
            )
        logline(f"[gpu] waiting free VRAM: {free_mb}/{min_free_mb} MiB")
        time.sleep(interval_sec)


def _unwrap_result(r: Any) -> Dict[str, Any]:
    """解包 FastMCP 結果"""
    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:
            txt = c.get("text") if isinstance(c, dict) else getattr(c, "text", None)
            if txt:
                try:
                    return json.loads(txt)
                except:
                    pass
    
    return {"_raw": str(r)}

def _resolve_data_path(project_id: str, collection: str) -> Path:
    if project_id:
        return DATA_JOBS / project_id / collection
    return DATA_JOBS / collection

def append_run_meta(project_id: str, collection: str, meta: dict):
    base = _resolve_data_path(project_id, collection)
    base.mkdir(parents=True, exist_ok=True)
    path = base / "runs.jsonl"
    meta = dict(meta)
    meta["saved_at"] = int(time.time())
    path.open("a", encoding="utf-8").write(json.dumps(meta, ensure_ascii=False) + "\n")

def write_meta_json(project_id: str, collection: str, meta: dict):
    base = _resolve_data_path(project_id, collection)
    base.mkdir(parents=True, exist_ok=True)
    meta = dict(meta)
    meta["saved_at"] = int(time.time())
    (base / "meta.json").write_text(
        json.dumps(meta, ensure_ascii=False, indent=2),
        encoding="utf-8"
    )

# =========================================
# 🎯 核心功能 1: Dispatch（調用工具）
# =========================================
async def _dispatch(target: str, tool: str, args: Dict[str, Any]) -> Dict[str, Any]:
    """
    通用調度器：調用任何 MCP Server 的任何工具
    """
    # 檢查 server 是否存在
    if target not in SERVERS:
        return {
            "ok": False,
            "error": f"未知的 server: {target}",
            "available": list(SERVERS.keys())
        }
    
    # 檢查工具是否存在
    server_info = SERVERS[target]
    if tool not in server_info.get("tools", {}):
        return {
            "ok": False,
            "error": f"{target} 沒有工具: {tool}",
            "available": list(server_info.get("tools", {}).keys())
        }
    
    url = server_info["url"]
    print(f"[dispatch] {target}.{tool} -> {url}")
    #job_log(args.get("job_id"), f"[dispatch] {target}.{tool} -> {url}")
    jid = args.get("_job_id_for_log") or args.get("job_id")
    job_log(jid, f"[dispatch] {target}.{tool} -> {url}")
    
    try:
        send_args = dict(args)
        send_args.pop("_job_id_for_log", None)
        #send_args.pop("job_id", None)
        #send_args.pop("upload_time", None)
        if not (target == "script" and tool == "extract_key_info"):
            send_args.pop("job_id", None)
            send_args.pop("upload_time", None)
        async with Client(url) as client:
            result = await client.call_tool(tool, send_args)
        
        return {
            "ok": True,
            "target": target,
            "tool": tool,
            "result": _unwrap_result(result)
        }
    except Exception as e:
        print(f"❌ [dispatch] 錯誤: {e}")
        return {
            "ok": False,
            "target": target,
            "tool": tool,
            "error": str(e)
        }


@mcp.tool()
async def dispatch(target: str, tool: str, args: Dict[str, Any] | None = None) -> Dict[str, Any]:
    """MCP 工具：供外部調用"""
    return await _dispatch(target, tool, args or {})


# =========================================
# 🎯 核心功能 2: Planner（規劃步驟）
# =========================================
def _build_tool_catalog() -> str:
    """從配置生成工具目錄"""
    lines = []
    for server_id, server_info in SERVERS.items():
        lines.append(f"\n[{server_id}] - {server_info.get('name', server_id)}")
        lines.append(f"  用途: {server_info.get('description', '')}")
        lines.append("  工具:")
        
        for tool_name, tool_info in server_info.get("tools", {}).items():
            desc = tool_info.get("description", "")
            lines.append(f"    - {tool_name}: {desc}")
    
    return "\n".join(lines)


def _plan(user_goal: str, has_audio: bool, model: str, selected_tools: list[str] | None = None) -> List[Dict[str, Any]]:
    """
    讓 LLM 規劃步驟
    """
    catalog = _build_tool_catalog()

    selected_tools = selected_tools or []     # ✅ 1) 確保是 list
    #allowed = set(selected_tools)
    allowed = set(selected_tools) if selected_tools else None 
    
    # 🔑 關鍵改進：不提供「典型流程」
    prompt = f"""你是智能任務規劃器。請根據目標選擇**最少必要**的工具。

{catalog}

[用戶目標]
{user_goal}

[當前狀態]
- 有音檔: {has_audio}

你“只能”使用下列工具（target.tool）：
{json.dumps(selected_tools, ensure_ascii=False)}

[注意]
- 若 selected_tools 非空，steps 內每一步都必須來自 selected_tools
- args 若不確定可省略，系統會自動補齊
- 沒有音檔時不要使用 upload_and_stt

[輸出格式]（純 JSON，不要其他文字）
{{
  "steps": [
    {{"target": "server_name", "tool": "tool_short_name", "args": {{}}}}
  ]
}}

[欄位規則]
- target 必須是 selected_tools 中任一項目的前綴（例如 "stt" 或 "script"）
- tool 必須是工具短名（例如 "upload_and_stt"、"script_generation"），不要包含 "stt." / "script." 前綴
"""


    
    print(f"[planner] 使用模型: {model}")
    if allowed:
        print(f"[planner] 限制工具: {selected_tools}")
    
    try:
        obj = ask_ollama_json(prompt, model=model)
        print(f"[planner] 規劃: {_safe_json(obj)}")
        
        steps = obj.get("steps", [])
        if not isinstance(steps, list):
            return []
        
        # 驗證步驟格式 + normalize
        valid_steps = []
        for s in steps:
            if not (isinstance(s, dict) and "target" in s and "tool" in s):
                continue

            target = str(s["target"])
            tool = str(s["tool"])

            # ✅ NEW: target 亂填（<server>/server），但 allowed 只有 1 個 server 時，直接推回正確 server
            if allowed and (target in ("<server>", "server", "Server", "SERVER")):
                allowed_targets = sorted({t.split(".", 1)[0] for t in allowed})
                if len(allowed_targets) == 1:
                    target = allowed_targets[0]

            # ✅ 情況 1：tool 寫成 "script.xxx" / "stt.xxx"（全名），就以 prefix 修正 target，tool 轉短名
            if "." in tool:
                prefix, rest = tool.split(".", 1)
                if prefix in SERVERS:
                    target = prefix
                    tool = rest

            # ✅ 情況 2：tool 寫成 "<target>.xxx" 再剝一次
            if tool.startswith(f"{target}."):
                tool = tool.split(".", 1)[1]

            tool_id = f"{target}.{tool}"


            if allowed and tool_id not in allowed:
                print(f"⚠️ [planner] 跳過不允許的工具: {tool_id}")
                continue

            # ✅ 一定要用 normalize 後的 target/tool
            valid_steps.append({
                "target": target,
                "tool": tool,
                "args": s.get("args", {})
            })

        
        return valid_steps
        
    except Exception as e:
        print(f"❌ [planner] 錯誤: {e}")
        return []


# =========================================
# 🎯 核心功能 3: 參數填充
# =========================================
def _fill_args(
    target: str,
    tool: str,
    step_args: Dict[str, Any],
    context: Dict[str, Any]
) -> Dict[str, Any]:
    """
    自動填充工具參數

    context 包含:
    - project_id, collection, flag
    - filename, b64, language, llm_refine
    - comment, query_text, job_id, upload_time
    - audio_text (從前面步驟拿到的 transcript)
    """

    args = dict(step_args)

    # =============================
    # 讀 tools_config.json 預設參數
    # =============================
    if target in SERVERS:
        tool_info = SERVERS[target].get("tools", {}).get(tool, {})
        default_params = tool_info.get("params", {})

        for key, value in default_params.items():
            args.setdefault(key, value)

    # =============================
    # 基本 context
    # =============================
    pid = context.get("project_id", "")
    col = context.get("collection", "audio_v1")

    # =============================
    # STT server
    # =============================
    if target == "stt":

        # 新架構：永遠傳 project_id + collection
        args.setdefault("project_id", pid)
        args.setdefault("collection", col)

        # upload_and_stt 需要音檔
        if tool == "upload_and_stt":
            args.setdefault("filename", context.get("filename"))
            args.setdefault("b64", context.get("b64"))
            args.setdefault("language", context.get("language", "zh"))

        # refine
        if tool == "refine":
            args.setdefault("llm_refine", context.get("llm_refine", True))

        # query
        if tool == "query":
            if "q" not in args:
                args["q"] = context.get("query_text", "")

    # =============================
    # Script server
    # =============================
    if target == "script":

        # text2script 需要的 comp_name/industry 由 script_gen/server_script.py
        # 從 meta.json 讀取後再往下游送。這裡只呼叫本地 script MCP，
        # 避免 planner 多塞參數時觸發 FastMCP unexpected keyword。
        if tool in ("save_markdown_text", "get_feature_result", "get_script_result_from_ui"):
            args = {
                "project_id": pid,
                "collection": col,
            }
        else:
            args.setdefault("project_id", pid)
            args.setdefault("collection", col)

        if tool == "extract_key_info":
            args.setdefault("job_id", context.get("job_id"))
            args.setdefault("upload_time", context.get("upload_time"))

    # =============================
    # VLM server
    # =============================
    if target == "vlm":
        args.setdefault("project_id", pid)
        args.setdefault("collection", col)

        if tool in ("upload_and_vlm", "get_vlm_transcript_pipeline", "parse_video"):
            args.setdefault("filename", context.get("video_filename"))
            args.setdefault("b64", context.get("video_b64"))
            args.setdefault("language", context.get("language", "zh"))

        if tool == "refine":                                    # ← 只有 refine 才需要
            args.setdefault("llm_refine", context.get("llm_refine", True))

        if tool == "query":
            if "q" not in args:
                args["q"] = context.get("query_text", "")

    # Claude VLM server
    if target == "claude_vlm":
        args.setdefault("project_id", pid)
        args.setdefault("collection", col)
        vf = context.get("video_filename")
        vb = context.get("video_b64")
        vp = context.get("video_path")
        if vf: args.setdefault("filename", vf)
        if vb: args.setdefault("b64", vb)
        if vp: args.setdefault("video_path", vp)

    # Claude Doc server
    if target == "claude_doc":
        if tool in ("sync_raw", "generate_docs", "update_claudemd"):
            args.setdefault("project_name", pid)
        if tool == "update_claudemd":
            args.setdefault("filename", context.get("claude_filename"))
            args.setdefault("b64", context.get("claude_b64"))

    # docparse server
    if target == "docparse":
        # ✅ 優先用 doc 專屬欄位，沒有才 fallback 到通用 filename/b64
        args.setdefault("filename", context.get("doc_filename") or context.get("filename"))
        args.setdefault("b64",      context.get("doc_b64")      or context.get("b64"))
        args.setdefault("project_id", pid)
        args.setdefault("collection", context.get("collection", "audio_v1"))   

    # picparser server
    if target == "picparser":
        args.setdefault("filename", context.get("pic_filename") or context.get("filename"))
        args.setdefault("b64",      context.get("pic_b64")      or context.get("b64"))
        args.setdefault("project_id", pid)
        args.setdefault("collection", context.get("collection", "audio_v1"))        
    return args


# =========================================
# 🎯 主入口：run_agent
# =========================================
@mcp.tool()
async def run_agent(
    filename: str | None = None,
    b64: str | None = None,
    user_goal: str = "將音檔轉文字並優化品質",
    language: str = "zh",
    llm_refine: bool = True,
    planner_model: str = DEFAULT_PLANNER_MODEL,
    project_id: str = "",
    collection: str = "audio_v1",
    flag: int = 1,
    comment: str = "",
    industry: str = "",
    vendor_name: str = "",
    vendor_id: str = "",
    query_text: str | None = None,
    selected_tools: List[str] | None = None,
    job_id: str | None = None,
    upload_time: str | None = None,
    video_filename: str | None = None,
    video_b64: str | None = None,
    video_path: str | None = None,
    filetype: str | None = None,
    # ✅ 新增 doc 相關欄位
    doc_filename: str | None = None,
    doc_b64: str | None = None,
    pic_filename: str | None = None,
    pic_b64: str | None = None,
    claude_filename: str | None = None,
    claude_b64: str | None = None,
) -> Dict[str, Any]:
    """
    智能 Agent：根據目標動態規劃和執行
    """
    logs = []
    results = []
    def logline(s: str):
        logs.append(s)          # 原本回傳用的 logs
        job_log(job_id, s)      # ✅ 給 UI 輪詢用
        print(s)                # 你 terminal 也想看就保留
    
    has_audio = bool(filename and b64)
    has_video = bool(video_filename and video_b64)
    vendor_id = vendor_id or vendor_name
    
    logs.append(f"🎯 目標: {user_goal}")
    logs.append(f"📂 Project: {project_id} / Collection: {collection} / Flag: {flag}")
    logs.append(f"🎤 有音檔: {has_audio}")
    
    logline(f"[run_agent] 目標={user_goal}")
    logline(f"[run_agent] project_id={project_id}, collection={collection}, flag={flag}")
    logline(f"[run_agent] industry={industry}, vendor_name={vendor_name}, vendor_id={vendor_id}")
    logline(f"[run_agent] 有音檔={has_audio}")
    logline(f"[run_agent] 🆔 job_id: {job_id}, ⏱ upload_time: {upload_time}")

    meta_payload = {
        "project_id": project_id,
        "collection": collection,
        "flag": flag,
        "job_id": job_id,
        "upload_time": upload_time,
        "comment": comment,
        "industry": industry,
        "vendor_name": vendor_name,
        "vendor_id": vendor_id,
        "selected_tools": selected_tools or [],
        "has_audio": has_audio,
        "user_goal": user_goal,
        "filetype": filetype,
    }

    append_run_meta(project_id, collection, meta_payload)
    write_meta_json(project_id, collection, meta_payload)
    
    # 1️⃣ 規劃步驟
    steps = _plan(user_goal, has_audio, planner_model, selected_tools)
    
    # 2️⃣ Fallback（最小可行）
    if not steps:
        logs.append("⚠️ Planner 未產生步驟，使用 fallback")

        if selected_tools:
            # ✅ 完全照 UI 勾選跑（不讓 LLM 決定）
            steps = []
            for tool_id in selected_tools:
                target, tool = tool_id.split(".", 1)
                steps.append({"target": target, "tool": tool, "args": {}})
        elif has_audio:
            steps = [{"target": "stt", "tool": "upload_and_stt", "args": {}}]
        else:
            logs.append("（no-audio）未提供音檔，且 planner 未產生步驟 → 直接返回 ok")
            return {"ok": True, "logs": logs, "results": []}

    
    logs.append(f"📋 規劃 {len(steps)} 步驟")

    needs_gpu_stt = any(
        step.get("target") == "stt" and step.get("tool") == "upload_and_stt"
        for step in steps
    )
    if needs_gpu_stt and planner_model == LOCAL_GPU_MODEL:
        if _env_bool("OLLAMA_UNLOAD_AFTER_PLANNER", "1"):
            try:
                did_unload = unload_ollama_model(planner_model)
                if did_unload:
                    logline(f"[ollama] unloaded planner model: {planner_model}")
            except Exception as e:
                logline(f"[ollama] unload planner model failed: {e}")
        _wait_gpu_free_for_local_model(planner_model, logline)
    
    # 3️⃣ 執行步驟
    context = {
        "project_id": project_id,
        "collection": collection,
        "flag": flag,
        "filename": filename,
        "b64": b64,
        "language": language,
        "llm_refine": llm_refine,
        "comment": comment,
        "industry": industry,
        "vendor_name": vendor_name,
        "vendor_id": vendor_id,
        "query_text": query_text,
        "job_id": job_id,
        "upload_time": upload_time,
        "audio_text": "",  # 會在 stt 步驟後填入
        "video_filename": video_filename,
        "video_b64": video_b64,
        "video_path": video_path,
        "has_video": has_video,
        "doc_filename": doc_filename,
        "doc_b64": doc_b64,
        "pic_filename": pic_filename,
        "pic_b64": pic_b64,
        "claude_filename": claude_filename,
        "claude_b64": claude_b64,
    }
    
    for i, step in enumerate(steps, start=1):
        target = step["target"]
        tool = step["tool"]
        
        logs.append(f"[{i}] {target}.{tool}")
        logline(f"[step {i}] 執行 {target}.{tool}")
        
        # 防呆：沒音檔不能 upload_and_stt
        if tool in ("upload_and_vlm", "get_vlm_transcript_pipeline") and not context.get("has_video"):
            result = {"ok": False, "error": "需要操作影片但未提供"}
            results.append({"step": i, "target": target, "tool": tool, "result": result})
            logs.append(f"  ❌ 失敗: 需要影片")
            break

        if tool == "upload_and_stt" and not has_audio:
            result = {"ok": False, "error": "需要音檔但未提供"}
            results.append({"step": i, "target": target, "tool": tool, "result": result})
            logs.append(f"  ❌ 失敗: 需要音檔")
            break

        # 防呆：parse_video 需要影片
        if target == "claude_vlm" and tool == "parse_video" and not context.get("video_filename"):
            result = {"ok": False, "error": "需要影片但未提供（video_filename/video_b64 為空）"}
            results.append({"step": i, "target": target, "tool": tool, "result": result})
            logs.append(f"  ❌ 失敗: 需要影片")
            break

        # 防呆：docparse 需要文件檔
        if tool == "parse_file" and not (context.get("doc_filename") or context.get("filename")):
            result = {"ok": False, "error": "需要文件但未提供（doc_filename/b64 為空）"}
            results.append({"step": i, "target": target, "tool": tool, "result": result})
            logs.append(f"  ❌ 失敗: 需要文件")
            break
        
        # 防呆：picparser 需要圖片
        if tool == "parse_image" and not (context.get("pic_filename") or context.get("filename")):
            result = {"ok": False, "error": "需要圖片但未提供（pic_filename/b64 為空）"}
            results.append({"step": i, "target": target, "tool": tool, "result": result})
            logs.append(f"  ❌ 失敗: 需要圖片")
            break

        if target == "claude_doc" and tool == "update_claudemd" and not (context.get("claude_filename") and context.get("claude_b64")):
            result = {"ok": False, "error": "需要 CLAUDE.md 但未提供（claude_filename/claude_b64 為空）"}
            results.append({"step": i, "target": target, "tool": tool, "result": result})
            logs.append("  ❌ 失敗: 需要 CLAUDE.md")
            break
        
        # 填充參數
        args = _fill_args(target, tool, step.get("args", {}), context)
        args["_job_id_for_log"] = job_id
        
        # 執行
        result = await _dispatch(target, tool, args)
        results.append({"step": i, "target": target, "tool": tool, "result": result})
        
        # 檢查是否成功
        if not result.get("ok"):
            logline(f"  ❌ 失敗: {result.get('error', '未知錯誤')}")
            break
        else:
            logline(f"  ✅ 成功")
    
    return {
        "ok": True,
        "logs": logs,
        "results": results
    }

@mcp.tool()
async def get_job_logs(job_id: str, from_idx: int = 0) -> Dict[str, Any]:
    items = JOB_LOGS.get(job_id, [])
    return {
        "ok": True,
        "job_id": job_id,
        "from_idx": from_idx,
        "next_idx": len(items),
        "lines": items[from_idx:],
    }

# =========================================
# 🚀 啟動服務
# =========================================
if __name__ == "__main__":
    print("="*50)
    print("🚀 Agent Orchestrator 啟動中...")
    print(f"📋 已載入 {len(SERVERS)} 個 MCP Server:")
    for server_id, server_info in SERVERS.items():
        tool_count = len(server_info.get("tools", {}))
        print(f"  - {server_id}: {tool_count} 工具")
    print("="*50)
    
    mcp.run(transport="http", host="0.0.0.0", port=8081, path="/stt")
