259 lines
11 KiB
Python
259 lines
11 KiB
Python
"""会话: 内存中的消息列表 + meta + 落 PG `messages` 表。
|
|
|
|
§7 B Step 2:消息走 ORM(append-only, idx 严格递增,payload jsonb)。
|
|
|
|
system prompt **不入库** —— 每次 build_agent 重建拼到 messages[0](§3.7
|
|
"memory 演化即时生效")。Session 内存里仍维持 [system, user_1, assistant_1, ...]
|
|
全列表;DB idx 从 0 开始数第一条非 system 消息。
|
|
|
|
保留 `atomic_write_text` 给 skill 产物 / 其他 .md 文件写入使用。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from pathlib import Path
|
|
from typing import Any, Dict, List, Optional
|
|
from uuid import UUID
|
|
|
|
from sqlalchemy import delete, func, select
|
|
|
|
from .storage import session_scope
|
|
from .storage.models import Message, Task, sanitize_jsonb_nul
|
|
from .file_store import atomic_write_text
|
|
from .markdown_guard import normalize_markdown_fences
|
|
|
|
|
|
def _to_dict(msg: Any) -> Any:
|
|
if isinstance(msg, dict):
|
|
return msg
|
|
if hasattr(msg, "model_dump"):
|
|
return msg.model_dump(exclude_none=True)
|
|
if hasattr(msg, "dict"):
|
|
return msg.dict(exclude_none=True)
|
|
return msg
|
|
|
|
|
|
class Session:
|
|
"""消息列表 anchored on task_id。
|
|
|
|
Lazy-persist: 构造时不动 DB,第一条非 system 消息 append 时:
|
|
1) 调 ensure_task_row 保证 tasks 行存在(Step 2 用占位值,Step 3 由 TaskState 提供完整值)
|
|
2) INSERT 一行 messages
|
|
|
|
系统 reset 走 DB DELETE 该 task 全部 messages。
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
task_id: UUID,
|
|
system_prompt: str = "",
|
|
meta: Optional[dict] = None,
|
|
) -> None:
|
|
self.task_id: UUID = task_id
|
|
self.messages: List[dict] = []
|
|
self.meta: Dict[str, Any] = dict(meta or {})
|
|
self._db_idx: int = 0 # 下一条要写 DB 的 idx
|
|
# 上下文窗口元数据(0019/0021):_base_idx = 窗口起点的 DB idx;_n_head = 内存
|
|
# 头部"不在 DB 里"的消息条数(system + 可选注入的前情摘要)。
|
|
# 映射不变量:messages[i] 的 DB idx = _base_idx + (i - _n_head),i >= _n_head。
|
|
self._base_idx: int = 0
|
|
self._n_head: int = 0
|
|
if system_prompt:
|
|
self.messages.append({"role": "system", "content": system_prompt})
|
|
self._n_head = 1
|
|
|
|
def append(
|
|
self,
|
|
msg: Any,
|
|
*,
|
|
artifact_refs: Optional[list[dict]] = None,
|
|
) -> Optional[UUID]:
|
|
"""追加消息;非 system 落 DB,system 仅内存。返回新落库行的 message_id。
|
|
|
|
前置条件:tasks 行已由 web 入口(`POST /v1/tasks` → `ensure_local_task_row`)写入;
|
|
Session 不再做 idempotent ensure(无 user 上下文,且 task 必先在,多余)。
|
|
|
|
返回值:非 system → 新 row 的 message_id(供 loop 给 usage_events 关联用);
|
|
system 消息不入库,返 None。旧调用方忽略返回值不影响行为。
|
|
"""
|
|
# 与 Message.payload 的 ORM 守门保持一致,使当前 run 的内存上下文和落库值
|
|
# 完全相同;外部工具提取文本偶尔会携带 PostgreSQL JSONB 不支持的 NUL。
|
|
msg_dict = sanitize_jsonb_nul(_to_dict(msg))
|
|
if msg_dict.get("role") == "assistant" and isinstance(msg_dict.get("content"), str):
|
|
fence_result = normalize_markdown_fences(msg_dict["content"])
|
|
msg_dict["content"] = fence_result.text
|
|
if fence_result.repairs:
|
|
print(
|
|
f"[markdown:fence-repair] task={self.task_id} "
|
|
f"repairs={fence_result.repairs}",
|
|
flush=True,
|
|
)
|
|
elif fence_result.unclosed_fence:
|
|
print(
|
|
f"[markdown:fence-warning] task={self.task_id} unclosed=1",
|
|
flush=True,
|
|
)
|
|
self.messages.append(msg_dict)
|
|
if msg_dict.get("role") == "system":
|
|
return None
|
|
|
|
with session_scope() as s:
|
|
row = Message(
|
|
task_id=self.task_id,
|
|
idx=self._db_idx,
|
|
payload=msg_dict,
|
|
artifact_refs=artifact_refs,
|
|
)
|
|
s.add(row)
|
|
s.flush() # 触发 INSERT 拿到 server-default 生成的 message_id
|
|
msg_id = row.message_id
|
|
self._db_idx += 1
|
|
return msg_id
|
|
|
|
def last_measured_usage(self) -> Optional[tuple]:
|
|
"""窗口内最后一条带 provider 实报 usage 的 assistant 消息,返回
|
|
(内存 pos, tokens_in, tokens_out);没有 / 查询失败一律返 None。
|
|
|
|
供上下文体量估算(context.estimate_window_tokens)做实测校准 —— 校准信号
|
|
而非正确性数据,契约是 **best-effort 绝不抛**:无 DB(测试 / CLI 冷路径)、
|
|
映射失配(理论不可达)等一切异常都吞掉走 None,调用方自然回退 chars 估算。
|
|
"""
|
|
try:
|
|
with session_scope() as s:
|
|
row = s.execute(
|
|
select(Message.idx, Message.tokens_in, Message.tokens_out)
|
|
.where(
|
|
Message.task_id == self.task_id,
|
|
Message.idx >= self._base_idx,
|
|
Message.tokens_in.isnot(None),
|
|
Message.tokens_in > 0,
|
|
)
|
|
.order_by(Message.idx.desc())
|
|
.limit(1)
|
|
).first()
|
|
if row is None:
|
|
return None
|
|
# DB idx → 内存 pos(映射不变量见 __init__);越界/角色不符视为映射失配,弃用。
|
|
pos = self._n_head + (row.idx - self._base_idx)
|
|
if not (self._n_head <= pos < len(self.messages)):
|
|
return None
|
|
if self.messages[pos].get("role") != "assistant":
|
|
return None
|
|
return pos, int(row.tokens_in), int(row.tokens_out or 0)
|
|
except Exception:
|
|
return None
|
|
|
|
@property
|
|
def context_head_len(self) -> int:
|
|
"""内存头部不落 DB 的消息条数(system + 可选前情摘要),窗口 idx 映射用。"""
|
|
return self._n_head
|
|
|
|
@property
|
|
def context_base(self) -> int:
|
|
"""当前窗口起点的 DB idx(tasks.context_base_idx 的内存镜像)。"""
|
|
return self._base_idx
|
|
|
|
def apply_fold(self, cutoff: int, summary: str) -> None:
|
|
"""折叠内存窗口(§8.8 Phase 2):裁掉 head 与 cutoff 之间的消息,注入新前情摘要。
|
|
|
|
只动内存;DB 持久化(tasks.context_summary + context_base_idx)由
|
|
context_fold.persist_fold 先行完成。cutoff 为内存 index,必须指向 user 消息
|
|
(find_cutoff 保证),折叠后窗口 = [system?] + [前情摘要] + messages[cutoff:]。
|
|
_db_idx 不变(总条数没变,append 续号不受影响)。
|
|
"""
|
|
from .context_fold import build_summary_note
|
|
|
|
new_base = self._base_idx + (cutoff - self._n_head)
|
|
head = (
|
|
[self.messages[0]]
|
|
if self._n_head and self.messages and self.messages[0].get("role") == "system"
|
|
else []
|
|
)
|
|
self.messages = head + [build_summary_note(summary)] + self.messages[cutoff:]
|
|
self._n_head = len(head) + 1
|
|
self._base_idx = new_base
|
|
|
|
def reset(self, keep_system: bool = True) -> None:
|
|
"""清空消息。keep_system 仅影响内存(system 本来就不在 DB)。"""
|
|
if keep_system and self.messages and self.messages[0].get("role") == "system":
|
|
self.messages = [self.messages[0]]
|
|
self._n_head = 1
|
|
else:
|
|
self.messages = []
|
|
self._n_head = 0
|
|
with session_scope() as s:
|
|
s.execute(delete(Message).where(Message.task_id == self.task_id))
|
|
self._db_idx = 0
|
|
self._base_idx = 0
|
|
|
|
@classmethod
|
|
def load(
|
|
cls,
|
|
task_id: UUID,
|
|
system_prompt: str = "",
|
|
meta: Optional[dict] = None,
|
|
) -> "Session":
|
|
"""从 DB 读历史 messages。system_prompt 由调用方注入(memory 演化即时生效)。
|
|
|
|
若 task_id 在 DB 不存在,返回空 Session(messages 只含 system,_db_idx=0);
|
|
调用方判断该不该报错。
|
|
|
|
只把 idx >= tasks.context_base_idx 的消息装进 LLM 上下文(channel 长会话软重置,
|
|
0019)。base 之前的历史仍全量留 messages 表(web `/messages` 不 gate,照旧翻得到)。
|
|
**关键**:`_db_idx` 必须取 DB 真实总条数(下一条 append 的 idx),不能用 len(rows)
|
|
—— 否则下次 append 会复用已存在的 idx,撞 uq_messages_task_idx / 覆盖历史。
|
|
"""
|
|
sess = cls(task_id=task_id, system_prompt=system_prompt, meta=meta)
|
|
with session_scope() as s:
|
|
task_row = s.execute(
|
|
select(Task.context_base_idx, Task.context_summary)
|
|
.where(Task.task_id == task_id)
|
|
).first()
|
|
base = (task_row.context_base_idx if task_row else 0) or 0
|
|
summary = (task_row.context_summary if task_row else "") or ""
|
|
# 前情摘要(0021,§8.8 Phase 2):仅内存注入在 system 之后,不占 DB idx。
|
|
if summary:
|
|
from .context_fold import build_summary_note
|
|
sess.messages.append(build_summary_note(summary))
|
|
sess._n_head += 1
|
|
sess._base_idx = base
|
|
rows = s.execute(
|
|
select(Message)
|
|
.where(Message.task_id == task_id, Message.idx >= base)
|
|
.order_by(Message.idx)
|
|
).scalars().all()
|
|
for row in rows:
|
|
payload = dict(row.payload)
|
|
if payload.get("role") == "assistant" and isinstance(payload.get("content"), str):
|
|
# 历史行不回写生产库;只在重建 LLM 上下文时应用同一窄修复,
|
|
# 与 Web 展示层保持一致,避免旧坏围栏继续污染后续轮次。
|
|
payload["content"] = normalize_markdown_fences(payload["content"]).text
|
|
sess.messages.append(payload)
|
|
# 真实总条数(含 base 之前的归档历史),保证 append 续号不撞 idx。
|
|
sess._db_idx = s.execute(
|
|
select(func.count())
|
|
.select_from(Message)
|
|
.where(Message.task_id == task_id)
|
|
).scalar_one()
|
|
return sess
|
|
|
|
@classmethod
|
|
def task_exists(cls, task_id: UUID) -> bool:
|
|
"""tasks 行 + messages 至少 1 条 → 该 task 真存在(不是 lazy 占位)。"""
|
|
with session_scope() as s:
|
|
row = s.execute(
|
|
select(Task.task_id).where(Task.task_id == task_id)
|
|
).scalar_one_or_none()
|
|
if row is None:
|
|
return False
|
|
cnt = s.execute(
|
|
select(Message.message_id)
|
|
.where(Message.task_id == task_id)
|
|
.limit(1)
|
|
).scalar_one_or_none()
|
|
return cnt is not None
|
|
|
|
def n_user_msgs(self) -> int:
|
|
"""内存里 user 消息数,用于 _cleanup_if_empty 守门(避免回 DB)。
|
|
跳过头部注入消息 —— 前情摘要也是 user 角色,不算真实用户发言。"""
|
|
return sum(1 for m in self.messages[self._n_head:] if m.get("role") == "user")
|