Files
ai-worker-platform/engine.py
T

290 lines
12 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.
# -*- coding: utf-8 -*-
"""
任务执行引擎 V1
- DAG 依赖校验 + 下游自动触发(串行/并行)
- RAG 知识库上下文注入
- 预算告警(项目预算使用率阈值)
- 事件通知(待审核/完成/失败)与告警记录
"""
import json
import threading
import traceback
import db
import llm_gateway
import config
import rag
import notify
def _log(task_id, level, message):
try:
db.w('INSERT INTO task_logs (task_id, level, message, created_at) VALUES (?,?,?,?)',
(task_id, level, message, db.now()))
except Exception:
pass
def _set_task(task_id, **fields):
if not fields:
return
fields['updated_at'] = db.now()
sets = ', '.join(f'{k}=?' for k in fields)
db.w(f'UPDATE tasks SET {sets} WHERE id=?', (*fields.values(), task_id))
def _cost_record(task, worker, usage):
db.w(
'INSERT INTO cost_records (task_id, project_id, worker_id, provider, model, '
'prompt_tokens, completion_tokens, total_tokens, cost, created_at) '
'VALUES (?,?,?,?,?,?,?,?,?,?)',
(task['id'], task['project_id'], worker['id'], worker['provider'],
worker['model'], usage['prompt_tokens'], usage['completion_tokens'],
usage['total_tokens'], usage['cost'], db.now()))
def _deps(task):
try:
return json.loads(task.get('depends_on') or '[]')
except Exception:
return []
def check_dependencies(task):
"""DAG 依赖检查:返回 (ok, blockers)"""
blockers = []
for dep_id in _deps(task):
dep = db.q('SELECT id, title, status FROM tasks WHERE id=?', (dep_id,), one=True)
if not dep:
blockers.append(f'# {dep_id}(已删除)')
elif dep['status'] != 'done':
blockers.append(f'「{dep["title"]}」(#{dep_id}) {dep["status"]}')
return (not blockers), blockers
def pick_worker_auto(task):
"""自动路由:按模型输入单价升序挑选 enabled Worker"""
rows = db.q('SELECT * FROM workers WHERE status="enabled" ORDER BY id')
if not rows:
return None
best, best_price = None, None
for r in rows:
pin, pout = llm_gateway.model_price(r['model'])
price = pin + pout * 0.5
if best_price is None or price < best_price:
best, best_price = r, price
return best
def check_worker_limits(worker):
"""成本上限预检:返回 (ok, reason)"""
if worker['task_cost_limit'] and worker['task_cost_limit'] > 0:
used = db.task_worker_cost(worker['id'])
if used >= worker['task_cost_limit']:
return False, f'该 Worker 累计成本 {used:.4f} 元已达单任务上限 {worker["task_cost_limit"]} 元'
if worker['monthly_cost_limit'] and worker['monthly_cost_limit'] > 0:
used = db.monthly_worker_cost(worker['id'])
if used >= worker['monthly_cost_limit']:
return False, f'该 Worker 本月成本 {used:.4f} 元已达月度上限 {worker["monthly_cost_limit"]} 元'
return True, ''
def _project_cost(project_id):
rows = db.q('SELECT COALESCE(SUM(cost),0) AS t FROM cost_records WHERE project_id=?',
(project_id,))
return rows[0]['t'] if rows else 0.0
def _check_project_budget(task):
proj = db.q('SELECT * FROM projects WHERE id=?', (task['project_id'],), one=True)
if proj and proj['budget_limit'] and proj['budget_limit'] > 0:
used = _project_cost(task['project_id'])
if used >= proj['budget_limit']:
return False, f'项目预算已用完({used:.2f}/{proj["budget_limit"]:.2f} 元)'
return True, ''
def _budget_alert(project_id):
"""预算使用率告警(每次任务完成后检查,避免重复刷屏)"""
proj = db.q('SELECT * FROM projects WHERE id=?', (project_id,), one=True)
if not proj or not proj['budget_limit'] or proj['budget_limit'] <= 0:
return
used = _project_cost(project_id)
ratio = used / proj['budget_limit']
if ratio >= config.BUDGET_ALERT_RATIO:
# 同项目 1 小时内只告警一次,避免刷屏
dup = db.q('SELECT COUNT(*) c FROM alerts WHERE type="budget" AND detail LIKE ? '
'AND created_at > ?', (f'项目「{proj["name"]}」%', db.now() - 3600))
if dup[0]['c'] == 0:
notify.notify('budget_alert',
f'预算告警:项目「{proj["name"]}」已使用 {ratio*100:.0f}%',
f'已花费 ¥{used:.2f} / 预算 ¥{proj["budget_limit"]:.2f}'
f'超过阈值 {config.BUDGET_ALERT_RATIO*100:.0f}%,请关注成本控制。',
save_alert=True, level='warn', atype='budget')
def _build_messages(task, worker):
"""构造提示词:任务指令 + RAG 知识库上下文"""
messages = []
if worker['system_prompt']:
messages.append({'role': 'system', 'content': worker['system_prompt']})
user_text = task['description'] or task['title']
ctx, refs = rag.build_context(task['project_id'], user_text)
if ctx:
user_text = f'{ctx}\n\n----\n\n任务指令:{user_text}'
_log(task['id'], 'info', f'📚 RAG 知识库命中 {len(refs)} 个片段:' + ''.join(refs[:5]))
messages.append({'role': 'user', 'content': user_text})
return messages
def _trigger_downstream(task):
"""DAG:任务完成后自动触发所有就绪的下游任务"""
rows = db.q('SELECT * FROM tasks WHERE status IN ("todo","failed")')
triggered = []
for t in rows:
deps = _deps(t)
if task['id'] not in deps:
continue
ok, blockers = check_dependencies(t)
if ok:
if runner.submit(t['id']):
triggered.append(t)
else:
_log(t['id'], 'info', f'⏳ 等待前置任务完成:' + '、'.join(blockers))
return triggered
def run_task(task_id):
"""在后台线程中执行任务"""
task = db.q('SELECT * FROM tasks WHERE id=?', (task_id,), one=True)
if not task:
return
if task['status'] == 'running':
return
# DAG 依赖检查
ok, blockers = check_dependencies(task)
if not ok:
_set_task(task_id, status='failed', error='前置任务未完成:' + '、'.join(blockers),
finished_at=db.now())
_log(task_id, 'error', '❌ 依赖未满足,无法执行:' + '、'.join(blockers))
notify.notify('task_failed', f'任务失败:{task["title"]}',
f'项目 #{task["project_id"]} 任务「{task["title"]}」因依赖未完成被拒绝执行:'
+ '、'.join(blockers), save_alert=True, level='warn', atype='task_failed')
return
# 确定 Worker
worker = None
if task['worker_id']:
worker = db.q('SELECT * FROM workers WHERE id=?', (task['worker_id'],), one=True)
if not worker or worker['status'] != 'enabled':
_set_task(task_id, status='failed', error='指定 Worker 不存在或已停用',
finished_at=db.now())
_log(task_id, 'error', '指定 Worker 不存在或已停用')
notify.notify('worker_alert', f'Worker 异常:任务「{task["title"]}」',
f'指定 Worker #{task["worker_id"]} 不存在或已停用', save_alert=True,
level='warn', atype='worker_alert')
return
else:
worker = pick_worker_auto(task)
if not worker:
_set_task(task_id, status='failed', error='无可用 Worker(自动路由失败)',
finished_at=db.now())
_log(task_id, 'error', '自动路由失败:无可用 Worker')
notify.notify('worker_alert', f'Worker 异常:任务「{task["title"]}」',
'自动路由失败:没有可用的 Worker', save_alert=True,
level='warn', atype='worker_alert')
return
_set_task(task_id, worker_id=worker['id'])
_log(task_id, 'info', f'自动路由 → Worker「{worker["name"]}」({worker["provider"]}/{worker["model"]}')
# 预算/成本预检
ok, reason = check_worker_limits(worker)
if not ok:
_set_task(task_id, status='failed', error=reason, finished_at=db.now())
_log(task_id, 'error', reason)
notify.notify('budget_alert', f'成本上限拦截:任务「{task["title"]}」', reason,
save_alert=True, level='warn', atype='budget')
return
ok, reason = _check_project_budget(task)
if not ok:
_set_task(task_id, status='failed', error=reason, finished_at=db.now())
_log(task_id, 'error', reason)
notify.notify('budget_alert', f'预算拦截:任务「{task["title"]}」', reason,
save_alert=True, level='warn', atype='budget')
return
_set_task(task_id, status='running', started_at=db.now(), error='')
_log(task_id, 'info', f'开始执行:Worker「{worker["name"]}」 模型 {worker["provider"]}/{worker["model"]}')
try:
usage = llm_gateway.chat(
worker['provider'], worker['model'], _build_messages(task, worker),
temperature=worker['temperature'], max_tokens=worker['max_tokens'],
base_url=worker['base_url'] or None, api_key=worker['api_key'] or None)
except Exception as e:
_set_task(task_id, status='failed', error=str(e), finished_at=db.now())
_log(task_id, 'error', f'执行失败: {e}')
notify.notify('task_failed', f'任务失败:{task["title"]}',
f'项目 #{task["project_id"]} 任务「{task["title"]}」执行出错:{str(e)[:300]}',
save_alert=True, level='warn', atype='task_failed')
return
_cost_record(task, worker, usage)
_log(task_id, 'success',
f'执行完成:{usage["total_tokens"]} tokens(输入 {usage["prompt_tokens"]} / 输出 {usage["completion_tokens"]}),'
f'成本 ¥{usage["cost"]:.6f}')
new_status = 'review' if task['review_required'] else 'done'
fields = {
'status': new_status,
'output_text': usage['text'],
'output_version': task['output_version'] + 1,
'finished_at': db.now(),
}
_set_task(task_id, **fields)
if new_status == 'review':
_log(task_id, 'info', '产出已提交,等待人工审核(HITL)')
notify.notify('task_review', f'任务待审核:{task["title"]}',
f'项目 #{task["project_id"]} 任务「{task["title"]}」已完成,等待人工验收。\n'
f'模型 {worker["provider"]}/{worker["model"]} · {usage["total_tokens"]} tokens · ¥{usage["cost"]:.4f}',
save_alert=True, level='info', atype='task_review')
else:
_log(task_id, 'success', '任务完成(无需审核)')
notify.notify('task_done', f'任务完成:{task["title"]}',
f'项目 #{task["project_id"]} 任务「{task["title"]}」执行完毕,'
f'成本 ¥{usage["cost"]:.4f}tokens {usage["total_tokens"]}',
save_alert=False)
_budget_alert(task['project_id'])
# DAG:触发下游就绪任务
downstream = _trigger_downstream(task)
for t in downstream:
_log(t['id'], 'info', f'🔗 前置任务「{task["title"]}」已完成,自动触发执行')
class TaskRunner:
def __init__(self):
self._threads = {}
def submit(self, task_id):
if task_id in self._threads and self._threads[task_id].is_alive():
return False
t = threading.Thread(target=self._safe_run, args=(task_id,), daemon=True)
self._threads[task_id] = t
t.start()
return True
def _safe_run(self, task_id):
try:
run_task(task_id)
except Exception:
_log(task_id, 'error', '引擎异常: ' + traceback.format_exc())
try:
_set_task(task_id, status='failed', error='引擎异常,见日志', finished_at=db.now())
except Exception:
pass
runner = TaskRunner()