pipeline-service/pipeline_service/task_capability.py

190 lines
8.8 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.

"""任务能力 — pipeline_tasks 的状态机语义化迁移(通用、产线无关)。
定位storage.py 是数据 CRUD 层create_task/list_tasks/get_task/update_task_state 低层改状态),
本模块是状态机语义化操作层 —— 每个迁移 CAS 原子(防并发/防越权覆盖)+ 租户隔离 + 审计。
任务的状态机/流转规则在 task skill 里LLM 读 skill 判断合法性);本模块只固化操作原语。
角色规范role 参数用 agent.{role}(无前缀自动补 agent.);人角色用 {orgtype}.{role}
"""
import logging
from sqlor.dbpools import DBPools
from appPublic.uniqueID import getID
from .audit import record_audit
DBNAME = "pipeline"
logger = logging.getLogger("pipeline.task_capability")
# 任务状态SDLC 默认,状态机语义见 task skill
S_SUBMITTED = "submitted" # 待认领
S_RUNNING = "running" # 角色 agent 执行中
S_REVIEW = "review" # 已提交待审核
S_APPROVED = "approved" # 审核通过
S_COMPLETED = "completed" # 全流程完成
S_WAITING = "waiting" # 挂起等回答
S_FAILED = "failed" # 失败
def _get_db():
db = DBPools()
if not db.databases:
from appPublic.jsonConfig import getConfig
config = getConfig()
if config.databases:
db.databases = config.databases
return db, DBNAME
def _normalize_role(role):
"""角色规范agent 角色补 agent. 前缀;人角色 {orgtype}.{role} 保留原样。"""
role = (role or "").strip()
if not role:
return ""
if "." in role:
return role
return f"agent.{role}"
async def _transition(task_id, tenant_id, from_state, to_state, action,
who=None, agent_id=None, detail=None):
"""CAS 状态迁移(清 claimed_by+ 审计。返回 (ok, message)。"""
if not task_id or not tenant_id:
return False, "缺少 task_id 或 tenant_id"
db, dbname = _get_db()
async with db.sqlorContext(dbname) as sor:
await sor.sqlExe(
"UPDATE pipeline_tasks SET state=${to}$, claimed_by=NULL, updated_at=NOW() "
"WHERE id=${tid}$ AND tenant_id=${tn}$ AND state=${from}$",
{"to": to_state, "tid": task_id, "tn": tenant_id, "from": from_state})
recs = await sor.R('pipeline_tasks', {'id': task_id})
await sor.sqlExe("COMMIT", {})
if not recs:
return False, "任务不存在"
cur = getattr(recs[0], 'state', '')
if cur != to_state:
return False, f"状态迁移失败(CAS): 期望 from={from_state} 实际 state={cur}"
await record_audit(tenant_id, 'pipeline_tasks', task_id, action,
from_state=from_state, to_state=to_state,
who=who, agent_id=agent_id, detail=detail, sor=sor)
return True, to_state
async def claim_task(tenant_id, role, agent_id, from_state=S_SUBMITTED,
to_state=S_RUNNING):
"""认领任务CAS claimed_by IS NULL 保证并发认领原子;设置 claimed_by=agent_id。
返回 (ok, task_id_or_message)。
"""
role_norm = _normalize_role(role) # 审计记录用规范名;查询用原样(兼容存量裸名数据)
db, dbname = _get_db()
async with db.sqlorContext(dbname) as sor:
claim_token = agent_id or getID()
recs = await sor.sqlExe(
"SELECT id FROM pipeline_tasks "
"WHERE tenant_id=${tn}$ AND state=${st}$ AND role=${role}$ "
"AND (claimed_by IS NULL OR claimed_by='') LIMIT 1",
{"tn": tenant_id, "st": from_state, "role": role})
await sor.sqlExe("COMMIT", {})
if not recs:
return False, f"没有可认领的 {role} 任务(state={from_state})"
task_id = getattr(recs[0], 'id', '')
# CASclaimed_by IS NULL 保证原子(并发 poller 不会双重认领)
await sor.sqlExe(
"UPDATE pipeline_tasks SET state=${to}$, claimed_by=${cb}$, updated_at=NOW() "
"WHERE id=${tid}$ AND tenant_id=${tn}$ AND state=${from}$ AND claimed_by IS NULL",
{"to": to_state, "cb": claim_token, "tid": task_id, "tn": tenant_id, "from": from_state})
chk = await sor.R('pipeline_tasks', {'id': task_id})
await sor.sqlExe("COMMIT", {})
if not chk or getattr(chk[0], 'state', '') != to_state:
return False, "认领竞争失败(已被他人认领)"
await record_audit(tenant_id, 'pipeline_tasks', task_id, 'claim',
from_state=from_state, to_state=to_state,
who=role_norm, agent_id=agent_id, sor=sor)
return True, task_id
async def submit_task(task_id, tenant_id, who=None, agent_id=None):
"""提交产出running → review清 claimed_by"""
return await _transition(task_id, tenant_id, S_RUNNING, S_REVIEW, 'submit',
who=who, agent_id=agent_id)
async def approve_task(task_id, tenant_id, who=None, agent_id=None, comment=None):
"""审核通过review → approved。"""
return await _transition(task_id, tenant_id, S_REVIEW, S_APPROVED, 'approve',
who=who, agent_id=agent_id, detail=comment)
async def reject_task(task_id, tenant_id, who=None, agent_id=None, comment=None):
"""审核退回review → submitted清 claimed_by被退角色重新认领"""
return await _transition(task_id, tenant_id, S_REVIEW, S_SUBMITTED, 'reject',
who=who, agent_id=agent_id, detail=comment)
async def complete_task(task_id, tenant_id, who=None, agent_id=None):
"""标记完成approved → completed也兼容 review → completed"""
ok, msg = await _transition(task_id, tenant_id, S_APPROVED, S_COMPLETED, 'complete',
who=who, agent_id=agent_id)
if ok:
return ok, msg
return await _transition(task_id, tenant_id, S_REVIEW, S_COMPLETED, 'complete',
who=who, agent_id=agent_id)
async def mark_failed(task_id, tenant_id, who=None, agent_id=None, error=None):
"""标记失败(任意状态 → failed清 claimed_by"""
db, dbname = _get_db()
async with db.sqlorContext(dbname) as sor:
recs = await sor.R('pipeline_tasks', {'id': task_id})
if not recs:
return False, "任务不存在"
from_state = getattr(recs[0], 'state', '')
await sor.sqlExe(
"UPDATE pipeline_tasks SET state='failed', last_error=${e}$, "
"claimed_by=NULL, updated_at=NOW() WHERE id=${tid}$ AND tenant_id=${tn}$",
{"e": (error or '')[:4000], "tid": task_id, "tn": tenant_id})
await sor.sqlExe("COMMIT", {})
await record_audit(tenant_id, 'pipeline_tasks', task_id, 'fail',
from_state=from_state, to_state=S_FAILED,
who=who, agent_id=agent_id, detail=error, sor=sor)
return True, S_FAILED
async def retry_task(task_id, tenant_id, who=None, agent_id=None):
"""重试failed → submittedretry_count+1清 claimed_by"""
db, dbname = _get_db()
async with db.sqlorContext(dbname) as sor:
await sor.sqlExe(
"UPDATE pipeline_tasks SET retry_count=retry_count+1, state='submitted', "
"claimed_by=NULL, last_error=NULL, updated_at=NOW() "
"WHERE id=${tid}$ AND tenant_id=${tn}$ AND state='failed'",
{"tid": task_id, "tn": tenant_id})
chk = await sor.R('pipeline_tasks', {'id': task_id})
await sor.sqlExe("COMMIT", {})
if not chk or getattr(chk[0], 'state', '') != S_SUBMITTED:
return False, "重试失败(任务不在 failed 状态)"
await record_audit(tenant_id, 'pipeline_tasks', task_id, 'retry',
from_state=S_FAILED, to_state=S_SUBMITTED,
who=who, agent_id=agent_id, sor=sor)
return True, S_SUBMITTED
async def suspend_task(task_id, tenant_id, who=None, agent_id=None):
"""挂起等回答running → waiting清 claimed_by"""
return await _transition(task_id, tenant_id, S_RUNNING, S_WAITING, 'suspend',
who=who, agent_id=agent_id)
async def revive_task(task_id, tenant_id, who=None, agent_id=None):
"""恢复认领waiting → submitted清 claimed_by重新认领"""
return await _transition(task_id, tenant_id, S_WAITING, S_SUBMITTED, 'revive',
who=who, agent_id=agent_id)
async def set_task_state(task_id, tenant_id, from_state, to_state,
who=None, agent_id=None, detail=None):
"""通用 CAS 状态迁移兜底(跨产线自定义状态机用)。"""
return await _transition(task_id, tenant_id, from_state, to_state, 'set_state',
who=who, agent_id=agent_id, detail=detail)