Files
auto_control/web/agent_api.py
T

158 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""AI 控制台 API:Agent 配置(模型/Key 前端可配)+ 运行 + 状态轮询。
Agent 在平台进程内跑(后台线程),连本机 MCP Server(8033)执行工具。
单实例:同一时间只允许一个 Agent 运行,避免多路操作设备冲突。
"""
import asyncio
import base64
import io
import os
import sys
import threading
import uuid
from flask import Blueprint, jsonify, request
from core.logger import get_logger
from core.models import db
from web.auth import admin_required
_log = get_logger("web.agent")
bp = Blueprint("agent", __name__)
# ---------- 配置键(app_meta) ----------
_CFG_KEYS = {"api_base": "agent_api_base",
"model": "agent_model",
"api_key": "agent_api_key",
"default_serial": "agent_default_serial"}
# ---------- 运行状态(单实例) ----------
_run = {"id": None, "state": "idle", "prompt": "", "serial": "",
"steps": [], "answer": "", "error": ""}
_lock = threading.Lock()
def _meta_get(key):
return db.session.execute(
db.text("SELECT value FROM app_meta WHERE key=:k"), {"k": key}).scalar() or ""
def _meta_put(key, value):
db.session.execute(
db.text("INSERT OR REPLACE INTO app_meta(key,value) VALUES(:k,:v)"),
{"k": key, "v": str(value)})
def _read_cfg():
return {k: _meta_get(v) for k, v in _CFG_KEYS.items()}
@bp.route("/api/agent/config", methods=["GET"])
@admin_required
def agent_config_get():
"""读 Agent 配置(key 打码返回)。"""
cfg = _read_cfg()
if cfg["api_key"]:
k = cfg["api_key"]
cfg["api_key_masked"] = k[:6] + "***" + k[-4:]
return jsonify({"ok": True, **cfg})
@bp.route("/api/agent/config", methods=["POST"])
@admin_required
def agent_config_save():
"""保存 Agent 配置:{api_base?, model?, api_key?, default_serial?} 部分更新。"""
data = request.json or {}
for key, meta_key in _CFG_KEYS.items():
if key in data and data[key] is not None:
_meta_put(meta_key, str(data[key]).strip())
db.session.commit()
return jsonify({"ok": True, "msg": "已保存"})
@bp.route("/api/agent/run", methods=["POST"])
@admin_required
def agent_run():
"""启动 Agent 执行指令:{prompt, serial?}。运行中返回 409。"""
data = request.json or {}
prompt = (data.get("prompt") or "").strip()
if not prompt:
return jsonify({"ok": False, "error": "请输入指令"}), 400
cfg = _read_cfg()
if not cfg.get("api_key"):
return jsonify({"ok": False, "error": "请先在配置区填写 API Key"}), 400
with _lock:
if _run["state"] == "running":
return jsonify({"ok": False, "error": "已有 Agent 运行中,请等待完成"}), 409
run_id = uuid.uuid4().hex[:8]
_run.update(id=run_id, state="running", prompt=prompt,
serial=(data.get("serial") or "").strip(),
steps=[], answer="", error="")
_log.info(f"Agent 启动: {prompt[:60]} serial={_run['serial']}")
threading.Thread(target=_agent_thread, args=(run_id, prompt, cfg),
daemon=True).start()
return jsonify({"ok": True, "run_id": run_id})
@bp.route("/api/agent/status")
@admin_required
def agent_status():
"""Agent 运行状态(前端轮询):{state, prompt, serial, steps, answer, error}。"""
with _lock:
return jsonify({"ok": True, **_run})
def _shrink_image(b64, width=220, quality=50):
"""截图降采样(web 展示用,避免大图撑爆轮询响应)。失败原样返回。"""
try:
from PIL import Image
img = Image.open(io.BytesIO(base64.b64decode(b64)))
if img.width > width:
img = img.resize((width, int(img.height * width / img.width)))
buf = io.BytesIO()
img.convert("RGB").save(buf, "JPEG", quality=quality)
return base64.b64encode(buf.getvalue()).decode()
except Exception:
return b64
def _agent_thread(run_id, prompt, cfg):
"""后台线程:跑 Agent,回调记录步骤。"""
try:
# 平台进程 cwd=/app(含 mcp_agent 包),进程内 import
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from mcp_agent.agent import Agent
def on_step(step):
with _lock:
if _run["id"] != run_id:
return
rec = {"tool": step.get("tool"), "args": str(step.get("args"))[:120]}
if step.get("image_b64"):
rec["image"] = _shrink_image(step["image_b64"])
# 保留最近 6 步截图,防止响应过大
_run["steps"] = _run["steps"][-5:]
_run["steps"].append(rec)
agent = Agent()
# 用 web 配置覆盖 Agent 默认(api key 等由前端配置)
agent.s.api_base = cfg.get("api_base") or agent.s.api_base
agent.s.model = cfg.get("model") or agent.s.model
agent.s.api_key = cfg.get("api_key") or agent.s.api_key
agent.s.default_serial = cfg.get("default_serial") or agent.s.default_serial
agent.on_step = on_step
async def _run():
await agent._load_tools()
return await agent.run(prompt, cfg.get("default_serial") or "")
answer = asyncio.run(_run())
with _lock:
_run["state"] = "done"
_run["answer"] = answer
except Exception as e:
_log.warning(f"Agent 运行异常: {e}")
with _lock:
_run["state"] = "error"
_run["error"] = f"{type(e).__name__}: {str(e)[:200]}"