Files
auto_control/web/agent_api.py
T

365 lines
14 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:模型/Key 前端配置 + Agent 流式执行(SSE 事件流)。
流程:
POST /api/agent/run {prompt} → 启动 Agent 线程,返回 run_id
GET /api/agent/stream?run_id= → SSE 事件流(EventSource 订阅):
event: delta {text, kind: content|reasoning} 流式文本增量
event: step {tool, args, image?} 工具调用完成(MCP 步骤)
event: done {answer} 完成
event: error {message} 失败
GET/POST /api/agent/config → 配置读写(key 打码回显)
单实例:同时只允许一个 Agent 运行。
"""
import asyncio
import base64
import io
import json
import os
import queue
import sys
import threading
import uuid
from flask import Blueprint, Response, 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": "",
"answer": "", "error": "",
"history": []} # 多轮对话历史 [{role: user|assistant, content}]
_queues = {} # run_id -> queue.Queue(SSE 消费者读取)
_stop_events = {} # run_id -> threading.Event(用户中断)
_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()}
# ================== 经验记忆(自进化) ==================
# agent_experience:任务成功后的操作配方,下次相似任务检索注入 system prompt。
# 原始 SQLite(CREATE IF NOT EXISTS 幂等),不进模型层迁移。
_EXP_TABLE = """
CREATE TABLE IF NOT EXISTS agent_experience (
id INTEGER PRIMARY KEY AUTOINCREMENT,
task_prompt TEXT DEFAULT '',
recipe TEXT DEFAULT '',
tool_seq TEXT DEFAULT '',
hits INTEGER DEFAULT 0,
created_at VARCHAR(20) DEFAULT '')"""
def _ensure_exp_table():
try:
db.session.execute(db.text(_EXP_TABLE))
db.session.commit()
except Exception:
pass
def _bigrams(text):
"""中文/英文文本 bigram 集合(无空格分词,粗粒度相似度)。"""
t = "".join(c for c in (text or "").lower() if c.isalnum() or "\u4e00" <= c <= "\u9fff")
return {t[i:i + 2] for i in range(len(t) - 1)}
def _find_experiences(prompt, limit=2, threshold=0.10):
"""按 bigram 重叠检索相似历史经验(prompt 与任务描述的字符相似度)。"""
try:
_ensure_exp_table()
rows = db.session.execute(db.text(
"SELECT task_prompt, recipe, hits FROM agent_experience "
"WHERE recipe != '' ORDER BY hits DESC, id DESC LIMIT 50")).fetchall()
except Exception:
return ""
if not rows:
return ""
cur = _bigrams(prompt)
if not cur:
return ""
scored = []
for task_prompt, recipe, hits in rows:
sim = len(cur & _bigrams(task_prompt)) / len(cur)
if sim >= threshold:
scored.append((sim, hits or 0, recipe))
scored.sort(key=lambda x: (-x[0], -x[1]))
parts = []
for sim, _hits, recipe in scored[:limit]:
parts.append(f"- {recipe[:600]}")
return "\n".join(parts)
def _distill_experience(cfg, prompt, tool_seq):
"""任务完成后用模型把操作序列提炼为可复用配方(失败静默,不阻塞)。"""
try:
import httpx
body = {
"model": cfg.get("model") or "deepseek-v4-flash-vision-exp",
"messages": [{"role": "user",
"content": "以下是一次成功的手机自动化操作记录。请提炼成简洁的"
"「操作配方」(2-6 步,每步:目标 → 用哪个工具),"
"供下次同类任务参考。不要解释,直接输出配方。\n"
f"任务:{prompt[:300]}\n操作序列:{tool_seq[:800]}"}],
"max_tokens": 600,
}
headers = {"Authorization": f"Bearer {cfg.get('api_key', '')}",
"Content-Type": "application/json"}
r = httpx.post(f"{(cfg.get('api_base') or 'https://api.deepseek.com').rstrip('/')}/chat/completions",
json=body, headers=headers, timeout=60)
if r.status_code != 200:
return ""
j = r.json()
recipe = ((j.get("choices") or [{}])[0].get("message") or {}).get("content") or ""
return recipe.strip()[:1500]
except Exception as e:
_log.warning(f"经验提炼失败: {e}")
return ""
def _save_experience(prompt, recipe, tool_seq):
try:
_ensure_exp_table()
from datetime import datetime
db.session.execute(db.text(
"INSERT INTO agent_experience(task_prompt, recipe, tool_seq, hits, created_at) "
"VALUES(:p, :r, :t, 0, :c)"),
{"p": prompt[:500], "r": recipe, "t": tool_seq[:1000],
"c": datetime.now().strftime("%Y-%m-%d %H:%M")})
db.session.commit()
_log.info("经验已保存(配方 %d 字符)", len(recipe))
except Exception as e:
_log.warning(f"经验保存失败: {e}")
# ================== 配置 ==================
@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
serial = (data.get("serial") or "").strip()
cfg = _read_cfg()
if not cfg.get("api_key"):
return jsonify({"ok": False, "error": "请先在配置区填写 API Key"}), 400
if not cfg.get("model"):
return jsonify({"ok": False, "error": "请先填写模型名"}), 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(),
answer="", error="")
# history 保留(同会话多轮对话),由前端「清空对话」调用 clear 重置
_queues[run_id] = queue.Queue()
_stop_events[run_id] = threading.Event()
_log.info(f"Agent 启动: {prompt[:60]}")
threading.Thread(target=_agent_thread,
args=(run_id, prompt, serial, cfg),
daemon=True).start()
return jsonify({"ok": True, "run_id": run_id})
@bp.route("/api/agent/stream")
@admin_required
def agent_stream():
"""SSE 事件流(EventSource):delta/step/done/error。"""
run_id = request.args.get("run_id", "")
with _lock:
if run_id != _run["id"]:
return jsonify({"ok": False, "error": "run_id 不存在"}), 404
q = _queues.get(run_id)
def gen():
while True:
try:
evt = q.get(timeout=15)
except queue.Empty:
yield ": keepalive\n\n" # 心跳防超时
continue
if evt is None:
break
kind, payload = evt
yield f"event: {kind}\ndata: {json.dumps(payload, ensure_ascii=False)}\n\n"
if kind in ("done", "error"):
break
return Response(gen(), mimetype="text/event-stream",
headers={"Cache-Control": "no-cache",
"X-Accel-Buffering": "no"})
@bp.route("/api/agent/stop", methods=["POST"])
@admin_required
def agent_stop():
"""中断当前运行的 Agent(下一个检查点生效,通常在数秒内)。"""
with _lock:
if _run["state"] != "running":
return jsonify({"ok": False, "error": "当前没有运行中的任务"}), 400
evt = _stop_events.get(_run["id"])
if evt:
evt.set()
_log.info("用户请求中断 Agent")
return jsonify({"ok": True, "msg": "已请求停止"})
@bp.route("/api/agent/clear", methods=["POST"])
@admin_required
def agent_clear():
"""清空对话历史。"""
with _lock:
_run["history"] = []
return jsonify({"ok": True, "msg": "已清空"})
def _shrink_image(b64, width=220, quality=50):
"""截图降采样(SSE step 事件用,控制传输体积)。失败原样返回。"""
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, serial, cfg):
"""后台线程:Agent 流式执行,事件推入队列供 SSE 消费。"""
q = _queues.get(run_id)
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_delta(text, kind):
q.put(("delta", {"text": text, "kind": kind}))
tool_seq = [] # 本轮工具序列(经验提炼用)
def on_tool(step):
rec = {"tool": step.get("tool"),
"args": str(step.get("args"))[:200]}
if step.get("image"):
rec["image"] = _shrink_image(step["image"])
q.put(("step", rec))
# 记录精简工具序列
try:
args = step.get("args") or {}
brief = {k: v for k, v in args.items() if k != "serial"}
tool_seq.append(f"{step.get('tool')}({str(brief)[:60]})")
except Exception:
pass
agent = Agent()
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
with _lock:
history = list(_run.get("history") or [])
target = serial or cfg.get("default_serial") or ""
stop_evt = _stop_events.get(run_id)
# 经验检索:相似历史任务的操作配方注入 system(自进化记忆)
exp_ctx = _find_experiences(prompt)
if exp_ctx:
_log.info("命中历史经验,注入参考配方")
q.put(("step", {"tool": "🧠 经验记忆",
"args": f"命中 {exp_ctx.count(chr(10) + '- ')} 条同类历史经验,已注入参考",
"image": None}))
async def _execute():
await agent._load_tools()
return await agent.run_stream(prompt, target,
history=history,
on_delta=on_delta, on_tool=on_tool,
should_stop=lambda: bool(
stop_evt and stop_evt.is_set()),
extra_context=exp_ctx)
# 整体超时保护:卡死时结束,释放单实例
answer = asyncio.run(asyncio.wait_for(_execute(), timeout=900))
with _lock:
_run["state"] = "done"
_run["answer"] = answer
# 追加本轮进历史(多轮连续性;上限 12 轮防 token 膨胀)
hist = _run.setdefault("history", [])
hist.append({"role": "user", "content": prompt[:2000]})
hist.append({"role": "assistant", "content": (answer or "")[:4000]})
_run["history"] = hist[-24:]
q.put(("done", {"answer": answer}))
# 自进化:成功后提炼操作配方存为经验(尽力而为,不阻塞/不影响结果)
if tool_seq:
try:
recipe = _distill_experience(cfg, prompt, " -> ".join(tool_seq))
if recipe and "配方" not in recipe[:50]:
_save_experience(prompt, recipe, " -> ".join(tool_seq))
except Exception as e:
_log.warning(f"经验保存异常: {e}")
except Exception as e:
_log.warning(f"Agent 运行异常: {e}")
with _lock:
_run["state"] = "error"
_run["error"] = f"{type(e).__name__}: {str(e)[:200]}"
q.put(("error", {"message": str(e)[:200]}))
finally:
q.put(None) # 关闭 SSE
_queues.pop(run_id, None)
_stop_events.pop(run_id, None)