zcbot/core/session.py

253 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) -> 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,
)
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")