feat(orchestrator): DataGate 机制化 + 任务级持久化(T14/T16 架构审查整改)
- T14 (OV5): 新建 src/genesis/orchestrator/datagate.py
- DataGate.load(source, selector) 机制化:子集加载 + 规模保护(超 500 行
无 selector 拒绝全量,1000 行 Excel 防上下文爆炸)+ token 预算(8000,
复用 CJK 保守估算)+ 未知表容错
- agent-runtime §4.2 原则→机制说明
- T16 (OV7): 新建 src/genesis/orchestrator/task_queue.py
- TaskQueue ABC + PersistentTaskQueue(SQLite 落盘)
- enqueue/poll/update_status/get/cancel/recover/close
- 幂等去重(§5.3 缓存结果)+ recover 将 running→failed、pending 保留
- api-design §5.2/5.3、agent-runtime §3.5/3.6、design §8.4.1 同步
- 新增 test_datagate.py(8 用例)+ test_task_queue.py(11 用例)
- TDD: RED(模块缺失)→ GREEN(聚焦 16 passed)→ 全量 218 passed / 100.00%(1140 stmts/278 br)
This commit is contained in:
@@ -0,0 +1,25 @@
|
||||
"""orchestrator 包:编排层组件(DataGate / TaskQueue 等)。
|
||||
|
||||
架构审查整改 Lane A:T14 DataGate 机制化(OV5)、T16 任务级持久化(OV7)。
|
||||
"""
|
||||
|
||||
from genesis.orchestrator.datagate import DataGate, DataGateError, DataSelector, DataGateResult
|
||||
from genesis.orchestrator.task_queue import (
|
||||
PersistentTaskQueue,
|
||||
TaskHandle,
|
||||
TaskQueue,
|
||||
TaskSpec,
|
||||
TaskStatus,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DataGate",
|
||||
"DataGateError",
|
||||
"DataSelector",
|
||||
"DataGateResult",
|
||||
"PersistentTaskQueue",
|
||||
"TaskHandle",
|
||||
"TaskQueue",
|
||||
"TaskSpec",
|
||||
"TaskStatus",
|
||||
]
|
||||
@@ -0,0 +1,121 @@
|
||||
"""DataGate 数据门(T14 机制化,OV5)。
|
||||
|
||||
背景:设计文档 §4.2 中 DataGate 仅是原则(「控制工作记忆 → 短时记忆的加载,
|
||||
避免上下文爆炸」)。OV5 裁定将其机制化:1000 行 Excel 等大源必须通过
|
||||
selector 限定子集才能进入 LLM 上下文,并提供 token 预算硬护栏。
|
||||
|
||||
机制:
|
||||
1. 子集加载 — selector.table_ids 指定要加载的表,不复制全量
|
||||
2. 规模保护 — 源总行数超过 max_total_rows 且未指定 selector → 拒绝(防上下文爆炸)
|
||||
3. token 预算 — 加载后估算 token(复用 inference/token 的 CJK 保守估算),
|
||||
超过 max_total_tokens → 拒绝
|
||||
4. 未知表容错 — selector 引用了不存在的表 → 返回空结果(不抛错)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from genesis.data_models import StructuredSource
|
||||
from genesis.inference.token import approximate_token_count
|
||||
|
||||
|
||||
class DataGateError(Exception):
|
||||
"""数据门拒绝加载(规模超限未限定 / token 超预算)。"""
|
||||
|
||||
|
||||
class DataSelector(BaseModel):
|
||||
"""加载子集描述:指定要进入上下文的表。
|
||||
|
||||
空 table_ids 等同未指定 → 走全量规模保护。
|
||||
"""
|
||||
|
||||
table_ids: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataGateResult:
|
||||
"""加载结果(供 prompt 组装方消费)。"""
|
||||
|
||||
loaded_tables: list[str]
|
||||
loaded_rows: int
|
||||
token_estimate: int
|
||||
# 未来可扩展:引用型数据(refs)与展开数据(content)分离
|
||||
|
||||
|
||||
class DataGate:
|
||||
"""控制「工作记忆 → 短时记忆」加载的机制化实现。
|
||||
|
||||
参数(可经 config/rag.yaml 或编排层注入调整):
|
||||
max_total_rows: 源总行数阈值;超过则必须提供 selector
|
||||
max_total_tokens: 加载结果 token 预算硬上限
|
||||
token_estimator: 估算函数(默认 CJK 保守估算,与 T9 一致)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_total_rows: int = 500,
|
||||
max_total_tokens: int = 8_000,
|
||||
token_estimator=approximate_token_count,
|
||||
) -> None:
|
||||
self.max_total_rows = max_total_rows
|
||||
self.max_total_tokens = max_total_tokens
|
||||
self._token_estimator = token_estimator
|
||||
|
||||
# ---------- 公共 API ----------
|
||||
|
||||
def load(self, source: StructuredSource, selector: DataSelector | None = None) -> DataGateResult:
|
||||
"""按 selector 从 StructuredSource 加载子集;无 selector 时全量(受规模保护)。
|
||||
|
||||
Raises:
|
||||
DataGateError: 规模超限未限定子集,或加载结果超 token 预算。
|
||||
"""
|
||||
total_rows = sum(len(t.rows) for t in source.tables)
|
||||
has_selector = selector is not None and bool(selector.table_ids)
|
||||
|
||||
if not has_selector and total_rows > self.max_total_rows:
|
||||
raise DataGateError(
|
||||
f"源数据 {total_rows} 行超过阈值 {self.max_total_rows},"
|
||||
"必须提供 selector 限定子集(如 DataSelector(table_ids=[...])),"
|
||||
"防上下文爆炸(OV5)。"
|
||||
)
|
||||
|
||||
tables = self._select_tables(source, selector)
|
||||
loaded_rows = sum(len(t.rows) for t in tables)
|
||||
token_estimate = self._estimate(tables)
|
||||
|
||||
if token_estimate > self.max_total_tokens:
|
||||
raise DataGateError(
|
||||
f"加载结果估算 {token_estimate} token 超过预算 {self.max_total_tokens},"
|
||||
"请缩小 selector 范围(如按表拆分加载)。"
|
||||
)
|
||||
|
||||
return DataGateResult(
|
||||
loaded_tables=[t.name for t in tables],
|
||||
loaded_rows=loaded_rows,
|
||||
token_estimate=token_estimate,
|
||||
)
|
||||
|
||||
# ---------- 内部 ----------
|
||||
|
||||
def _select_tables(self, source: StructuredSource, selector: DataSelector | None) -> list:
|
||||
if selector is None or not selector.table_ids:
|
||||
return list(source.tables)
|
||||
wanted = set(selector.table_ids)
|
||||
return [t for t in source.tables if t.name in wanted]
|
||||
|
||||
def _estimate(self, tables: list) -> int:
|
||||
"""估算表集合的 token 数:表头 + 每行单元格值文本。"""
|
||||
total = 0
|
||||
for table in tables:
|
||||
header_text = " ".join(str(h) for h in table.headers)
|
||||
total += self._token_estimator(header_text)
|
||||
for row in table.rows:
|
||||
row_text = " ".join(
|
||||
str(cell.value) if cell.value is not None else ""
|
||||
for cell in row.values()
|
||||
)
|
||||
total += self._token_estimator(row_text)
|
||||
return total
|
||||
@@ -0,0 +1,250 @@
|
||||
"""任务级持久化(T16,OV7)。
|
||||
|
||||
背景:api-design §5 定义 TaskQueue 抽象(v1 仅 InMemoryQueue),但崩溃恢复
|
||||
只到会话级(agent-runtime §3.6),任务层数据丢失。OV7 裁定任务级持久化。
|
||||
|
||||
实现:PersistentTaskQueue —— TaskQueue 抽象 + SQLite 落盘。
|
||||
- 任务状态 / payload / result 全部写入 SQLite(零外部依赖,标准库 sqlite3)
|
||||
- 幂等去重(§5.3):同 (session_id, step, chapter_id) 已完成 → 返回缓存结果
|
||||
- recover():重启后 running → failed(中断标记),pending 保留待执行
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class TaskStatus(str, Enum):
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
# 终态:不可再转移(cancel 仅对非终态生效)
|
||||
_TERMINAL = frozenset({TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED})
|
||||
|
||||
|
||||
class TaskSpec(BaseModel):
|
||||
"""任务投递规格(api-design §5.1)。"""
|
||||
|
||||
task_id: str
|
||||
session_id: str
|
||||
step: str
|
||||
chapter_id: str | None = None
|
||||
payload: dict[str, Any] = Field(default_factory=dict)
|
||||
idempotency_key: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskHandle:
|
||||
"""任务句柄(含状态与结果)。"""
|
||||
|
||||
task_id: str
|
||||
session_id: str
|
||||
step: str
|
||||
chapter_id: str | None = None
|
||||
payload: dict[str, Any] = field(default_factory=dict)
|
||||
idempotency_key: str = ""
|
||||
status: TaskStatus = TaskStatus.PENDING
|
||||
result: Any | None = None
|
||||
retry_count: int = 0
|
||||
created_at: str = ""
|
||||
updated_at: str = ""
|
||||
|
||||
|
||||
class TaskQueue(ABC):
|
||||
"""统一任务队列抽象(v1 仅 PersistentTaskQueue;Redis/Valkey 为 v2 预留)。"""
|
||||
|
||||
@abstractmethod
|
||||
def enqueue(self, task: TaskSpec) -> TaskHandle: ...
|
||||
|
||||
@abstractmethod
|
||||
def poll(self, session_id: str) -> list[TaskHandle]: ...
|
||||
|
||||
@abstractmethod
|
||||
def update_status(self, task_id: str, status: TaskStatus, result: Any = None) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def get(self, task_id: str) -> TaskHandle | None: ...
|
||||
|
||||
@abstractmethod
|
||||
def cancel(self, task_id: str) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def recover(self) -> list[TaskHandle]: ...
|
||||
|
||||
@abstractmethod
|
||||
def close(self) -> None: ...
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
class PersistentTaskQueue(TaskQueue):
|
||||
"""SQLite 持久化任务队列(T16)。
|
||||
|
||||
表结构 tasks:
|
||||
task_id PK | session_id | step | chapter_id | idempotency_key
|
||||
payload JSON | status | result JSON | created_at | updated_at
|
||||
"""
|
||||
|
||||
def __init__(self, db_path: Path | str) -> None:
|
||||
self._db_path = str(db_path)
|
||||
self._conn = sqlite3.connect(self._db_path)
|
||||
self._conn.row_factory = sqlite3.Row
|
||||
self._init_schema()
|
||||
|
||||
# ---------- 生命周期 ----------
|
||||
|
||||
def _init_schema(self) -> None:
|
||||
self._conn.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS tasks (
|
||||
task_id TEXT PRIMARY KEY,
|
||||
session_id TEXT NOT NULL,
|
||||
step TEXT NOT NULL,
|
||||
chapter_id TEXT,
|
||||
idempotency_key TEXT NOT NULL,
|
||||
payload TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
result TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)
|
||||
"""
|
||||
)
|
||||
self._conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_tasks_session ON tasks(session_id)"
|
||||
)
|
||||
self._conn.commit()
|
||||
|
||||
def close(self) -> None:
|
||||
self._conn.close()
|
||||
|
||||
# ---------- TaskQueue 接口 ----------
|
||||
|
||||
def enqueue(self, task: TaskSpec) -> TaskHandle:
|
||||
# 幂等去重(§5.3):同幂等键已存在 → 返回已有句柄(completed 带缓存结果)
|
||||
existing = self._find_by_idem(task.session_id, task.step, task.chapter_id, task.idempotency_key)
|
||||
if existing is not None:
|
||||
return existing
|
||||
|
||||
now = _now()
|
||||
self._conn.execute(
|
||||
"""
|
||||
INSERT INTO tasks (task_id, session_id, step, chapter_id, idempotency_key,
|
||||
payload, status, result, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
task.task_id,
|
||||
task.session_id,
|
||||
task.step,
|
||||
task.chapter_id,
|
||||
task.idempotency_key,
|
||||
json.dumps(task.payload, ensure_ascii=False),
|
||||
TaskStatus.PENDING.value,
|
||||
None,
|
||||
now,
|
||||
now,
|
||||
),
|
||||
)
|
||||
self._conn.commit()
|
||||
return self._row_to_handle(task.task_id)
|
||||
|
||||
def poll(self, session_id: str) -> list[TaskHandle]:
|
||||
rows = self._conn.execute(
|
||||
"SELECT * FROM tasks WHERE session_id = ? ORDER BY created_at",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [self._row_to_handle(row["task_id"], row=row) for row in rows]
|
||||
|
||||
def update_status(self, task_id: str, status: TaskStatus, result: Any = None) -> None:
|
||||
current = self.get(task_id)
|
||||
if current is None:
|
||||
raise KeyError(f"任务不存在: {task_id}")
|
||||
if current.status in _TERMINAL:
|
||||
raise ValueError(f"终态任务不可再转移: {task_id} ({current.status})")
|
||||
|
||||
self._conn.execute(
|
||||
"UPDATE tasks SET status = ?, result = ?, updated_at = ? WHERE task_id = ?",
|
||||
(
|
||||
status.value,
|
||||
json.dumps(result, ensure_ascii=False) if result is not None else None,
|
||||
_now(),
|
||||
task_id,
|
||||
),
|
||||
)
|
||||
self._conn.commit()
|
||||
|
||||
def get(self, task_id: str) -> TaskHandle | None:
|
||||
row = self._conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
||||
return self._row_to_handle(task_id, row=row) if row else None
|
||||
|
||||
def cancel(self, task_id: str) -> bool:
|
||||
current = self.get(task_id)
|
||||
if current is None or current.status in _TERMINAL:
|
||||
return False
|
||||
self.update_status(task_id, TaskStatus.CANCELLED)
|
||||
return True
|
||||
|
||||
def recover(self) -> list[TaskHandle]:
|
||||
"""崩溃恢复:running → failed(中断标记);pending 保留;返回全部未完成。"""
|
||||
rows = self._conn.execute("SELECT * FROM tasks WHERE status = ?", (TaskStatus.RUNNING.value,)).fetchall()
|
||||
for row in rows:
|
||||
self._conn.execute(
|
||||
"UPDATE tasks SET status = ?, updated_at = ? WHERE task_id = ?",
|
||||
(TaskStatus.FAILED.value, _now(), row["task_id"]),
|
||||
)
|
||||
self._conn.commit()
|
||||
incomplete = self._conn.execute(
|
||||
"SELECT * FROM tasks WHERE status IN (?, ?) ORDER BY created_at",
|
||||
(TaskStatus.PENDING.value, TaskStatus.FAILED.value),
|
||||
).fetchall()
|
||||
return [self._row_to_handle(row["task_id"], row=row) for row in incomplete]
|
||||
|
||||
# ---------- 内部 ----------
|
||||
|
||||
def _find_by_idem(self, session_id: str, step: str, chapter_id: str | None, idem: str) -> TaskHandle | None:
|
||||
row = self._conn.execute(
|
||||
"SELECT * FROM tasks WHERE session_id = ? AND step = ? AND idempotency_key = ? AND "
|
||||
"chapter_id IS ?",
|
||||
(session_id, step, idem, chapter_id),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
# chapter_id 可为 NULL(SQL 的 IS 处理);此处统一按精确匹配
|
||||
row = self._conn.execute(
|
||||
"SELECT * FROM tasks WHERE session_id = ? AND step = ? AND idempotency_key = ?",
|
||||
(session_id, step, idem),
|
||||
).fetchone()
|
||||
return self._row_to_handle(row["task_id"], row=row) if row else None
|
||||
|
||||
def _row_to_handle(self, task_id: str, row: sqlite3.Row | None = None) -> TaskHandle:
|
||||
if row is None:
|
||||
row = self._conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(f"任务不存在: {task_id}")
|
||||
return TaskHandle(
|
||||
task_id=row["task_id"],
|
||||
session_id=row["session_id"],
|
||||
step=row["step"],
|
||||
chapter_id=row["chapter_id"],
|
||||
payload=json.loads(row["payload"]),
|
||||
idempotency_key=row["idempotency_key"],
|
||||
status=TaskStatus(row["status"]),
|
||||
result=json.loads(row["result"]) if row["result"] else None,
|
||||
created_at=row["created_at"],
|
||||
updated_at=row["updated_at"],
|
||||
)
|
||||
Reference in New Issue
Block a user