feat(v2): T-M1 采纳 pi 会话树设计——AgentSession fork 分叉(parent_id/fork_point 血统 + 值拷贝隔离)
- SessionStore.fork(sid, turn_index, title):值拷贝 turns[:n+1] 生成新会话,
原会话不变;新会话记录 parent_id/fork_point(pi 的 id/parentId 树思想);
拷贝式而非共享存储引用——规避并发写冲突,会话体量小冗余可接受
- 旧会话数据无血统字段自然兼容(dict.get 语义,零迁移,对齐 pi 的版本迁移纪律)
- 新增 POST /agent/sessions/{sid}/fork 端点(404 = 会话不存在/分叉点越界)
- 新增 tests/test_session_fork.py 6 项(拷贝正确性/缺省语义/越界/深拷贝/
旧数据兼容/端点 200+404)
This commit is contained in:
@@ -15,6 +15,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import copy
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
@@ -889,6 +890,36 @@ class SessionStore:
|
|||||||
self.save(sess)
|
self.save(sess)
|
||||||
return sess
|
return sess
|
||||||
|
|
||||||
|
def fork(self, sid: str, turn_index: Optional[int] = None,
|
||||||
|
title: str = "") -> Optional[AgentSession]:
|
||||||
|
"""从既有会话分叉新会话(T-M1,采纳 pi 会话树 id/parentId 设计)。
|
||||||
|
|
||||||
|
值拷贝 turns[:turn_index+1](缺省 = 全部已完结轮次;-1 = 空会话壳),
|
||||||
|
原会话不变;新会话记录 parent_id / fork_point 保留血统。采用拷贝式
|
||||||
|
而非 pi 的共享存储引用——规避并发写冲突,代价是 fork 后磁盘冗余
|
||||||
|
(会话体量小,可接受)。turn_index 越界返回 None。
|
||||||
|
"""
|
||||||
|
src = self.get(sid)
|
||||||
|
if src is None:
|
||||||
|
return None
|
||||||
|
turns = src.data.get("turns", []) or []
|
||||||
|
n = len(turns) - 1 if turn_index is None else int(turn_index)
|
||||||
|
if n < -1 or n >= len(turns):
|
||||||
|
return None
|
||||||
|
data = copy.deepcopy(src.data)
|
||||||
|
data["id"] = "as" + uuid.uuid4().hex[:10]
|
||||||
|
base_title = (title or "").strip() or (src.data.get("title") or "会话") + "·分叉"
|
||||||
|
data["title"] = base_title[:24]
|
||||||
|
data["parent_id"] = sid
|
||||||
|
data["fork_point"] = n
|
||||||
|
data["created_at"] = time.time()
|
||||||
|
data["busy"] = False
|
||||||
|
data["turns"] = copy.deepcopy(turns[:n + 1]) if n >= 0 else []
|
||||||
|
sess = AgentSession(data)
|
||||||
|
self._cache[data["id"]] = sess
|
||||||
|
self._save(sess)
|
||||||
|
return sess
|
||||||
|
|
||||||
def save(self, sess: AgentSession) -> None:
|
def save(self, sess: AgentSession) -> None:
|
||||||
self._cache[sess.data["id"]] = sess
|
self._cache[sess.data["id"]] = sess
|
||||||
self._save(sess)
|
self._save(sess)
|
||||||
|
|||||||
@@ -740,6 +740,24 @@ try:
|
|||||||
raise HTTPException(status_code=404, detail=f"会话不存在: {sid}")
|
raise HTTPException(status_code=404, detail=f"会话不存在: {sid}")
|
||||||
return sess.view()
|
return sess.view()
|
||||||
|
|
||||||
|
@app.post("/agent/sessions/{sid}/fork", tags=["agent"])
|
||||||
|
async def agent_session_fork(sid: str, req: dict = None):
|
||||||
|
"""从会话分叉新会话:{"turn_index": 截止轮下标(缺省=全部轮), "title"}。
|
||||||
|
|
||||||
|
T-M1(采纳 pi 会话树设计):值拷贝历史轮次,原会话不变;
|
||||||
|
新会话带 parent_id/fork_point 血统字段。
|
||||||
|
"""
|
||||||
|
_check_id(sid, "会话 ID")
|
||||||
|
from gateway.agent import get_session_store
|
||||||
|
body = req or {}
|
||||||
|
turn_index = body.get("turn_index")
|
||||||
|
sess = get_session_store().fork(sid, turn_index,
|
||||||
|
str(body.get("title") or ""))
|
||||||
|
if sess is None:
|
||||||
|
raise HTTPException(status_code=404,
|
||||||
|
detail=f"会话不存在或分叉点越界: {sid}")
|
||||||
|
return sess.view()
|
||||||
|
|
||||||
@app.delete("/agent/sessions/{sid}", tags=["agent"])
|
@app.delete("/agent/sessions/{sid}", tags=["agent"])
|
||||||
async def agent_session_delete(sid: str):
|
async def agent_session_delete(sid: str):
|
||||||
_check_id(sid, "会话 ID")
|
_check_id(sid, "会话 ID")
|
||||||
|
|||||||
@@ -0,0 +1,98 @@
|
|||||||
|
"""会话树 fork 单元与端点测试(T-M1,采纳 pi 会话树 id/parentId 设计)。"""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
pytest.importorskip("fastapi")
|
||||||
|
pytest.importorskip("httpx")
|
||||||
|
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
import gateway.agent as ag
|
||||||
|
import gateway.api as ga
|
||||||
|
from gateway.agent import AgentSession, SessionStore
|
||||||
|
|
||||||
|
|
||||||
|
def _mk_turn(rid: str, task: str) -> dict:
|
||||||
|
return {"request_id": rid, "task": task, "response": f"ans-{rid}",
|
||||||
|
"state": "done", "tool_calls": [], "tokens": 10,
|
||||||
|
"error": None, "ts": 0.0}
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_copies_turns_and_keeps_lineage(tmp_path):
|
||||||
|
"""值拷贝历史轮次 + 记录 parent_id/fork_point 血统。"""
|
||||||
|
store = SessionStore(root=tmp_path / "sessions")
|
||||||
|
src = store.create("源会话", "ws")
|
||||||
|
src.data["turns"] = [_mk_turn("t1", "任务一"), _mk_turn("t2", "任务二")]
|
||||||
|
store.save(src)
|
||||||
|
|
||||||
|
forked = store.fork(src.data["id"], turn_index=0, title="分叉A")
|
||||||
|
assert forked is not None
|
||||||
|
assert forked.data["id"] != src.data["id"]
|
||||||
|
assert forked.data["parent_id"] == src.data["id"]
|
||||||
|
assert forked.data["fork_point"] == 0
|
||||||
|
assert [t["request_id"] for t in forked.data["turns"]] == ["t1"]
|
||||||
|
# 原会话不变
|
||||||
|
assert len(store.get(src.data["id"]).data["turns"]) == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_default_takes_all_turns(tmp_path):
|
||||||
|
"""缺省 turn_index = 全部已完结轮次。"""
|
||||||
|
store = SessionStore(root=tmp_path / "sessions")
|
||||||
|
src = store.create("源", "ws")
|
||||||
|
src.data["turns"] = [_mk_turn("t1", "a"), _mk_turn("t2", "b")]
|
||||||
|
store.save(src)
|
||||||
|
forked = store.fork(src.data["id"])
|
||||||
|
assert len(forked.data["turns"]) == 2
|
||||||
|
assert forked.data["fork_point"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_out_of_range_returns_none(tmp_path):
|
||||||
|
"""turn_index 越界 / 源会话不存在 -> None(端点映射 404)。"""
|
||||||
|
store = SessionStore(root=tmp_path / "sessions")
|
||||||
|
src = store.create("源", "ws")
|
||||||
|
src.data["turns"] = [_mk_turn("t1", "a")]
|
||||||
|
store.save(src)
|
||||||
|
assert store.fork(src.data["id"], turn_index=5) is None
|
||||||
|
assert store.fork(src.data["id"], turn_index=-2) is None
|
||||||
|
assert store.fork("as-not-exist") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_is_deep_copy(tmp_path):
|
||||||
|
"""fork 后修改新会话不影响原会话(值拷贝语义)。"""
|
||||||
|
store = SessionStore(root=tmp_path / "sessions")
|
||||||
|
src = store.create("源", "ws")
|
||||||
|
src.data["turns"] = [_mk_turn("t1", "a")]
|
||||||
|
store.save(src)
|
||||||
|
forked = store.fork(src.data["id"])
|
||||||
|
forked.data["turns"][0]["task"] = "被篡改"
|
||||||
|
forked.data["turns"].append(_mk_turn("t2", "新轮"))
|
||||||
|
assert store.get(src.data["id"]).data["turns"][0]["task"] == "a"
|
||||||
|
assert len(store.get(src.data["id"]).data["turns"]) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_old_sessions_without_lineage_fields_compatible(tmp_path):
|
||||||
|
"""旧数据无 parent_id/fork_point 字段自然兼容(dict.get 语义,零迁移)。"""
|
||||||
|
store = SessionStore(root=tmp_path / "sessions")
|
||||||
|
sess = AgentSession({"id": "as-legacy", "title": "旧", "turns": []})
|
||||||
|
store.save(sess)
|
||||||
|
assert store.get("as-legacy").data.get("parent_id", "") == ""
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_endpoint(tmp_path, monkeypatch):
|
||||||
|
"""POST /agent/sessions/{sid}/fork:端到端 200 + 404。"""
|
||||||
|
ag.reset_session_store()
|
||||||
|
monkeypatch.setattr(ag, "_session_store", SessionStore(root=tmp_path / "sessions"))
|
||||||
|
store = ag._session_store
|
||||||
|
src = store.create("源", "ws")
|
||||||
|
src.data["turns"] = [_mk_turn("t1", "a")]
|
||||||
|
store.save(src)
|
||||||
|
|
||||||
|
client = TestClient(ga.app)
|
||||||
|
r = client.post(f"/agent/sessions/{src.data['id']}/fork",
|
||||||
|
json={"turn_index": 0, "title": "端点分叉"})
|
||||||
|
assert r.status_code == 200
|
||||||
|
data = r.json()
|
||||||
|
assert data["parent_id"] == src.data["id"]
|
||||||
|
assert data["title"] == "端点分叉"
|
||||||
|
|
||||||
|
r404 = client.post("/agent/sessions/as-not-exist/fork", json={})
|
||||||
|
assert r404.status_code == 404
|
||||||
Reference in New Issue
Block a user