Files
ai-worker-platform/app.py
T

397 lines
16 KiB
Python

# -*- coding: utf-8 -*-
"""
AI Worker 项目管理平台 - MVP
人派活 → AI 干活 → 人验收 最小闭环
"""
import os
import functools
from flask import Flask, request, jsonify, session, send_from_directory
import config
import db
import engine
import llm_gateway
app = Flask(__name__, static_folder='static', static_url_path='')
app.secret_key = config.SECRET_KEY
db.init_db()
# ---------------------------------------------------------------------------
# 鉴权(简单口令登录)
# ---------------------------------------------------------------------------
def auth_enabled():
return bool(config.AUTH_PASSWORD)
@app.route('/api/login', methods=['POST'])
def login():
data = request.get_json(force=True)
if not auth_enabled():
return jsonify({'ok': True})
if data.get('password') == config.AUTH_PASSWORD:
session['authed'] = True
return jsonify({'ok': True})
return jsonify({'ok': False, 'error': '口令错误'}), 401
@app.route('/api/logout', methods=['POST'])
def logout():
session.clear()
return jsonify({'ok': True})
@app.route('/api/me')
def me():
return jsonify({'ok': True, 'authed': not auth_enabled() or session.get('authed')})
def require_auth(fn):
@functools.wraps(fn)
def wrapper(*args, **kwargs):
if auth_enabled() and not session.get('authed'):
return jsonify({'ok': False, 'error': '未登录'}), 401
return fn(*args, **kwargs)
return wrapper
@app.route('/api/health')
def health():
return jsonify({'ok': True, 'service': 'ai-worker-platform', 'version': 'v1.0.0'})
# ---------------------------------------------------------------------------
# 项目
# ---------------------------------------------------------------------------
@app.route('/api/projects', methods=['GET', 'POST'])
@require_auth
def projects():
if request.method == 'POST':
d = request.get_json(force=True)
pid = db.w(
'INSERT INTO projects (name, description, objective, acceptance_criteria, '
'status, budget_limit, created_at, updated_at) VALUES (?,?,?,?,?,?,?,?)',
(d.get('name', '').strip(), d.get('description', ''), d.get('objective', ''),
d.get('acceptance_criteria', ''), d.get('status', 'active'),
float(d.get('budget_limit') or 0), db.now(), db.now()))
return jsonify({'ok': True, 'id': pid})
rows = db.q('SELECT * FROM projects ORDER BY id DESC')
for r in rows:
r['task_count'] = db.q('SELECT COUNT(*) c FROM tasks WHERE project_id=?', (r['id'],))[0]['c']
r['done_count'] = db.q('SELECT COUNT(*) c FROM tasks WHERE project_id=? AND status="done"',
(r['id'],))[0]['c']
return jsonify({'ok': True, 'data': rows})
@app.route('/api/projects/<int:pid>', methods=['GET', 'PUT', 'DELETE'])
@require_auth
def project_detail(pid):
if request.method == 'GET':
p = db.q('SELECT * FROM projects WHERE id=?', (pid,), one=True)
if not p:
return jsonify({'ok': False, 'error': '项目不存在'}), 404
return jsonify({'ok': True, 'data': p})
if request.method == 'DELETE':
db.w('DELETE FROM tasks WHERE project_id=?', (pid,))
db.w('DELETE FROM cost_records WHERE project_id=?', (pid,))
db.w('DELETE FROM projects WHERE id=?', (pid,))
return jsonify({'ok': True})
d = request.get_json(force=True)
fields = ['name', 'description', 'objective', 'acceptance_criteria', 'status', 'budget_limit']
sets, args = [], []
for f in fields:
if f in d:
sets.append(f'{f}=?')
args.append(d[f])
if sets:
args.append(db.now())
db.w(f'UPDATE projects SET {", ".join(sets)}, updated_at=? WHERE id=?', (*args, pid))
return jsonify({'ok': True})
# ---------------------------------------------------------------------------
# Worker
# ---------------------------------------------------------------------------
@app.route('/api/workers', methods=['GET', 'POST'])
@require_auth
def workers():
if request.method == 'POST':
d = request.get_json(force=True)
wid = db.w(
'INSERT INTO workers (name, description, provider, model, base_url, api_key, '
'system_prompt, temperature, max_tokens, task_cost_limit, monthly_cost_limit, '
'status, created_at, updated_at) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?)',
(d.get('name', '').strip(), d.get('description', ''), d.get('provider', ''),
d.get('model', ''), d.get('base_url', ''), d.get('api_key', ''),
d.get('system_prompt', ''), float(d.get('temperature', 0.7)),
int(d.get('max_tokens', 2000)), float(d.get('task_cost_limit') or 0),
float(d.get('monthly_cost_limit') or 0), d.get('status', 'enabled'),
db.now(), db.now()))
return jsonify({'ok': True, 'id': wid})
rows = db.q('SELECT * FROM workers ORDER BY id DESC')
for r in rows:
r['month_cost'] = round(db.monthly_worker_cost(r['id']), 6)
return jsonify({'ok': True, 'data': rows})
@app.route('/api/workers/<int:wid>', methods=['GET', 'PUT', 'DELETE'])
@require_auth
def worker_detail(wid):
if request.method == 'GET':
w = db.q('SELECT * FROM workers WHERE id=?', (wid,), one=True)
return jsonify({'ok': True, 'data': w}) if w else (jsonify({'ok': False, 'error': '不存在'}), 404)
if request.method == 'DELETE':
db.w('UPDATE tasks SET worker_id=NULL WHERE worker_id=?', (wid,))
db.w('DELETE FROM workers WHERE id=?', (wid,))
return jsonify({'ok': True})
d = request.get_json(force=True)
fields = ['name', 'description', 'provider', 'model', 'base_url', 'api_key',
'system_prompt', 'temperature', 'max_tokens', 'task_cost_limit',
'monthly_cost_limit', 'status']
sets, args = [], []
for f in fields:
if f in d:
sets.append(f'{f}=?')
args.append(d[f])
if sets:
args.append(db.now())
db.w(f'UPDATE workers SET {", ".join(sets)}, updated_at=? WHERE id=?', (*args, wid))
return jsonify({'ok': True})
@app.route('/api/workers/<int:wid>/test', methods=['POST'])
@require_auth
def worker_test(wid):
w = db.q('SELECT * FROM workers WHERE id=?', (wid,), one=True)
if not w:
return jsonify({'ok': False, 'error': '不存在'}), 404
try:
r = llm_gateway.test_connection(w['provider'], w['model'],
base_url=w['base_url'] or None,
api_key=w['api_key'] or None)
return jsonify({'ok': True, 'data': r})
except Exception as e:
return jsonify({'ok': False, 'error': str(e)})
@app.route('/api/providers')
@require_auth
def providers():
data = []
for k, v in config.PROVIDERS.items():
data.append({
'id': k, 'name': v['name'], 'base_url': v['base_url'],
'has_key': bool(v['api_key']),
'models': [m for m in config.MODEL_PRICING.keys() if m.startswith(
{'doubao': 'doubao', 'deepseek': 'deepseek', 'openai': 'gpt',
'qwen': 'qwen', 'vllm': ''}.get(k, '__none__'))],
})
return jsonify({'ok': True, 'data': data})
# ---------------------------------------------------------------------------
# 任务
# ---------------------------------------------------------------------------
@app.route('/api/tasks', methods=['GET', 'POST'])
@require_auth
def tasks():
if request.method == 'POST':
d = request.get_json(force=True)
tid = db.w(
'INSERT INTO tasks (project_id, worker_id, title, description, priority, '
'review_required, deadline, created_at, updated_at) VALUES (?,?,?,?,?,?,?,?,?)',
(d.get('project_id'), d.get('worker_id'), d.get('title', '').strip(),
d.get('description', ''), d.get('priority', 'medium'),
1 if d.get('review_required', True) else 0, d.get('deadline', ''),
db.now(), db.now()))
db.w('INSERT INTO task_logs (task_id, level, message, created_at) VALUES (?,?,?,?)',
(tid, 'info', f'任务创建:{d.get("title","")}', db.now()))
return jsonify({'ok': True, 'id': tid})
pid = request.args.get('project_id')
if pid:
rows = db.q('SELECT * FROM tasks WHERE project_id=? ORDER BY id DESC', (int(pid),))
else:
rows = db.q('SELECT * FROM tasks ORDER BY id DESC LIMIT 200')
return jsonify({'ok': True, 'data': [db.serialize_task(t) for t in rows]})
@app.route('/api/tasks/<int:tid>', methods=['GET', 'PUT', 'DELETE'])
@require_auth
def task_detail(tid):
t = db.q('SELECT * FROM tasks WHERE id=?', (tid,), one=True)
if not t:
return jsonify({'ok': False, 'error': '任务不存在'}), 404
if request.method == 'GET':
t = db.serialize_task(t)
t['logs'] = db.q('SELECT * FROM task_logs WHERE task_id=? ORDER BY id', (tid,))
t['costs'] = db.q('SELECT * FROM cost_records WHERE task_id=? ORDER BY id', (tid,))
if t['worker_id']:
t['worker'] = db.q('SELECT id,name,provider,model FROM workers WHERE id=?',
(t['worker_id'],), one=True)
return jsonify({'ok': True, 'data': t})
if request.method == 'DELETE':
db.w('DELETE FROM task_logs WHERE task_id=?', (tid,))
db.w('DELETE FROM cost_records WHERE task_id=?', (tid,))
db.w('DELETE FROM tasks WHERE id=?', (tid,))
return jsonify({'ok': True})
d = request.get_json(force=True)
# 只允许在非运行中修改基础字段
if t['status'] == 'running':
return jsonify({'ok': False, 'error': '任务执行中,禁止修改'}), 400
fields = ['title', 'description', 'worker_id', 'priority', 'review_required',
'deadline', 'status']
sets, args = [], []
for f in fields:
if f in d:
sets.append(f'{f}=?')
args.append(d[f])
if sets:
args.append(db.now())
db.w(f'UPDATE tasks SET {", ".join(sets)}, updated_at=? WHERE id=?', (*args, tid))
return jsonify({'ok': True})
@app.route('/api/tasks/<int:tid>/run', methods=['POST'])
@require_auth
def task_run(tid):
t = db.q('SELECT * FROM tasks WHERE id=?', (tid,), one=True)
if not t:
return jsonify({'ok': False, 'error': '任务不存在'}), 404
if t['status'] == 'running':
return jsonify({'ok': False, 'error': '任务已在执行中'}), 400
if engine.runner.submit(tid):
return jsonify({'ok': True})
return jsonify({'ok': False, 'error': '任务已在执行中'}), 400
@app.route('/api/tasks/<int:tid>/cancel', methods=['POST'])
@require_auth
def task_cancel(tid):
t = db.q('SELECT * FROM tasks WHERE id=?', (tid,), one=True)
if not t:
return jsonify({'ok': False, 'error': '任务不存在'}), 404
if t['status'] != 'running':
return jsonify({'ok': False, 'error': '仅执行中的任务可取消'}), 400
# MVP:标记取消(线程无法强杀,完成后会回到 review;这里直接置 cancelled 并忽略结果)
db.w('UPDATE tasks SET status="cancelled", updated_at=?, finished_at=? WHERE id=?',
(db.now(), db.now(), tid))
db.w('INSERT INTO task_logs (task_id, level, message, created_at) VALUES (?,?,?,?)',
(tid, 'warn', '任务已由人工取消(引擎线程将在后台自然结束)', db.now()))
return jsonify({'ok': True})
@app.route('/api/tasks/<int:tid>/review', methods=['POST'])
@require_auth
def task_review(tid):
t = db.q('SELECT * FROM tasks WHERE id=?', (tid,), one=True)
if not t:
return jsonify({'ok': False, 'error': '任务不存在'}), 404
if t['status'] != 'review':
return jsonify({'ok': False, 'error': '仅待审核状态可审核'}), 400
d = request.get_json(force=True)
action = d.get('action')
reason = (d.get('reason') or '').strip()
if action == 'approve':
db.w('UPDATE tasks SET status="done", updated_at=? WHERE id=?', (db.now(), tid))
db.w('INSERT INTO task_logs (task_id, level, message, created_at) VALUES (?,?,?,?)',
(tid, 'success', '✅ 人工审核通过,任务完成', db.now()))
return jsonify({'ok': True})
if action == 'reject':
if not reason:
return jsonify({'ok': False, 'error': '打回必须填写原因'}), 400
db.w('UPDATE tasks SET status="todo", rejection_count=rejection_count+1, '
'output_text="", updated_at=? WHERE id=?', (db.now(), tid))
db.w('INSERT INTO task_logs (task_id, level, message, created_at) VALUES (?,?,?,?)',
(tid, 'warn', f'⛔ 人工打回:{reason}(等待返工重跑)', db.now()))
return jsonify({'ok': True})
return jsonify({'ok': False, 'error': '未知操作'}), 400
# ---------------------------------------------------------------------------
# 报表 / 统计
# ---------------------------------------------------------------------------
@app.route('/api/reports/cost')
@require_auth
def report_cost():
group = request.args.get('group', 'project')
if group == 'worker':
rows = db.q(
'SELECT worker_id, provider, model, COUNT(*) runs, SUM(total_tokens) tokens, '
'SUM(cost) cost FROM cost_records GROUP BY worker_id ORDER BY cost DESC')
for r in rows:
w = db.q('SELECT name FROM workers WHERE id=?', (r['worker_id'],), one=True)
r['worker_name'] = w['name'] if w else f'#{r["worker_id"]}'
elif group == 'model':
rows = db.q(
'SELECT provider, model, COUNT(*) runs, SUM(total_tokens) tokens, '
'SUM(cost) cost FROM cost_records GROUP BY model ORDER BY cost DESC')
else:
rows = db.q(
'SELECT project_id, COUNT(*) runs, SUM(total_tokens) tokens, '
'SUM(cost) cost FROM cost_records GROUP BY project_id ORDER BY cost DESC')
for r in rows:
p = db.q('SELECT name FROM projects WHERE id=?', (r['project_id']), one=True)
r['project_name'] = p['name'] if p else f'#{r["project_id"]}'
return jsonify({'ok': True, 'data': rows})
@app.route('/api/stats')
@require_auth
def stats():
out = {'projects': 0, 'tasks': 0, 'workers': 0, 'total_cost': 0, 'total_tokens': 0,
'by_status': {}, 'recent': [], 'daily_cost': []}
out['projects'] = db.q('SELECT COUNT(*) c FROM projects')[0]['c']
out['workers'] = db.q('SELECT COUNT(*) c FROM workers')[0]['c']
out['tasks'] = db.q('SELECT COUNT(*) c FROM tasks')[0]['c']
for r in db.q('SELECT status, COUNT(*) c FROM tasks GROUP BY status'):
out['by_status'][r['status']] = r['c']
c = db.q('SELECT COALESCE(SUM(cost),0) cost, COALESCE(SUM(total_tokens),0) tokens '
'FROM cost_records')[0]
out['total_cost'], out['total_tokens'] = round(c['cost'], 4), c['tokens']
# 一次通过率:done 且 rejection_count=0
done = out['by_status'].get('done', 0)
clean = db.q('SELECT COUNT(*) c FROM tasks WHERE status="done" AND rejection_count=0')[0]['c']
out['one_pass_rate'] = round(clean / done * 100, 1) if done else None
out['rework_count'] = db.q('SELECT COALESCE(SUM(rejection_count),0) c FROM tasks')[0]['c']
out['recent'] = db.q(
'SELECT t.id, t.title, t.status, t.project_id, p.name AS project_name, t.updated_at '
'FROM tasks t LEFT JOIN projects p ON p.id=t.project_id '
'ORDER BY t.updated_at DESC LIMIT 10')
# 近 7 天成本
import datetime
rows = db.q('SELECT created_at, cost FROM cost_records ORDER BY created_at')
day_map = {}
for r in rows:
d = datetime.datetime.fromtimestamp(r['created_at']).strftime('%m-%d')
day_map[d] = round(day_map.get(d, 0) + r['cost'], 4)
for i in range(6, -1, -1):
d = (datetime.datetime.now() - datetime.timedelta(days=i)).strftime('%m-%d')
out['daily_cost'].append({'day': d, 'cost': day_map.get(d, 0)})
return jsonify({'ok': True, 'data': out})
@app.route('/api/logs')
@require_auth
def logs():
limit = min(int(request.args.get('limit', 100)), 500)
rows = db.q(
'SELECT l.*, t.title AS task_title FROM task_logs l '
'LEFT JOIN tasks t ON t.id=l.task_id ORDER BY l.id DESC LIMIT ?', (limit,))
return jsonify({'ok': True, 'data': rows})
# ---------------------------------------------------------------------------
# 前端
# ---------------------------------------------------------------------------
@app.route('/')
def index():
return send_from_directory(app.static_folder, 'index.html')
if __name__ == '__main__':
print(f'AI Worker 项目管理平台启动: http://0.0.0.0:{config.PORT}')
app.run(host=config.HOST, port=config.PORT, debug=False, threaded=True)