Files
Video-Tree-TRM5/core/evolution/protocols.py
T

153 lines
3.9 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.
"""core/evolution/ 子包的持久化 Protocol 定义。
SkillStore / PromptStore / RunLog 为只读——core/ 返回结果 dataclass
读取由 app/ 落盘的资源。SkillStore / PromptStore 同步(文件读取量小且快),
RunLog 异步(隔离 SQLite 查询,core/ 不写 SQL)。
DiagnosisSignalStore 兼具读写:逐题 upsert 诊断信号并支持断点续跑查询,
同样隔离 SQLite 实现,app/core 不写裸 SQL。
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable
if TYPE_CHECKING:
from core.evolution.types import DiagnosisSignalRow
@runtime_checkable
class SkillStore(Protocol):
"""版本化技能读取端口。
实现方解析 manifest 指针,core/ 不感知版本号。
"""
def read_skill(self, filename: str) -> str:
"""读取指定 skill 文件的全文内容。
参数:
filename: skill 文件名,如 'temporal-reasoning.md'。
返回:
文件全文内容。
"""
...
def list_skill_files(self) -> list[str]:
"""列出当前版本所有 skill 文件名。
返回:
文件名列表。
"""
...
@runtime_checkable
class PromptStore(Protocol):
"""版本化提示词读取端口。
覆盖 system.md 和 tool extract/verify 文件。
"""
def read_prompt(self, filename: str) -> str:
"""读取指定 prompt 文件的全文内容。
参数:
filename: prompt 文件名,如 'system.md'。
返回:
文件全文内容。
"""
...
def list_prompt_files(self) -> list[str]:
"""列出当前版本所有 prompt 文件名。
返回:
文件名列表。
"""
...
@runtime_checkable
class RunLog(Protocol):
"""实验日志查询端口。
隔离 SQLite 实现细节,core/ 不写 SQL。
"""
async def get_predictions(
self,
run_id: str,
*,
question_ids: list[str] | None = None,
) -> list[dict[str, Any]]:
"""查询指定 run 的预测记录。
参数:
run_id: 运行标识。
question_ids: 可选的题目 ID 过滤列表。
返回:
预测记录字典列表。
"""
...
async def get_traces(
self,
run_id: str,
*,
question_ids: list[str] | None = None,
) -> list[dict[str, Any]]:
"""查询指定 run 的推理轨迹。
参数:
run_id: 运行标识。
question_ids: 可选的题目 ID 过滤列表。
返回:
轨迹记录字典列表。
"""
...
@runtime_checkable
class DiagnosisSignalStore(Protocol):
"""逐题诊断信号存储端口。
隔离 SQLite 实现细节,app/core 不写裸 SQL。逐题 upsert 落盘、
支持断点续跑(done_question_ids 查已完成集合)。
"""
def upsert(self, row: DiagnosisSignalRow) -> None:
"""写入或覆盖单题诊断信号(按主键幂等)。
参数:
row: 待持久化的诊断信号行。
"""
...
def done_question_ids(self, baseline_run_id: str, diag_fingerprint: str) -> set[str]:
"""查询指定 run 与诊断指纹下已完成的 question_id 集合。
参数:
baseline_run_id: baseline run 标识。
diag_fingerprint: 诊断口径指纹。
返回:
已落盘信号的 question_id 集合,用于断点续跑跳过。
"""
...
def load(self, baseline_run_id: str, diag_fingerprint: str) -> list[DiagnosisSignalRow]:
"""加载指定 run 与诊断指纹下的全部诊断信号行。
参数:
baseline_run_id: baseline run 标识。
diag_fingerprint: 诊断口径指纹。
返回:
还原后的 DiagnosisSignalRow 列表。
"""
...