378 lines
14 KiB
Python
378 lines
14 KiB
Python
"""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()
|
||
_flask_app = None # web_server 注册时注入(后台线程 db 操作需 app context)
|
||
|
||
|
||
def set_app(app):
|
||
global _flask_app
|
||
_flask_app = app
|
||
|
||
|
||
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:
|
||
if _flask_app is None:
|
||
return ""
|
||
with _flask_app.app_context():
|
||
_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):
|
||
"""保存经验(后台线程调用,包 app context)。"""
|
||
try:
|
||
if _flask_app is None:
|
||
return
|
||
with _flask_app.app_context():
|
||
_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)
|