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:
lhl
2026-08-12 15:43:10 +08:00
parent 3decdfc31c
commit 6a54580ca4
9 changed files with 734 additions and 5 deletions
+25
View File
@@ -0,0 +1,25 @@
"""orchestrator 包:编排层组件(DataGate / TaskQueue 等)。
架构审查整改 Lane AT14 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",
]
+121
View File
@@ -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
+250
View File
@@ -0,0 +1,250 @@
"""任务级持久化(T16OV7)。
背景: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 仅 PersistentTaskQueueRedis/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 可为 NULLSQL 的 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"],
)