From f7e8273434076d491e7865b7266feb37722689fd Mon Sep 17 00:00:00 2001 From: tzt <14718231+flying-travel@user.noreply.gitee.com> Date: Sat, 19 Sep 2026 09:29:57 +0800 Subject: [PATCH] =?UTF-8?q?feat(v2):=20T-M1=20=E9=87=87=E7=BA=B3=20pi=20?= =?UTF-8?q?=E4=BC=9A=E8=AF=9D=E6=A0=91=E8=AE=BE=E8=AE=A1=E2=80=94=E2=80=94?= =?UTF-8?q?AgentSession=20fork=20=E5=88=86=E5=8F=89=EF=BC=88parent=5Fid/fo?= =?UTF-8?q?rk=5Fpoint=20=E8=A1=80=E7=BB=9F=20+=20=E5=80=BC=E6=8B=B7?= =?UTF-8?q?=E8=B4=9D=E9=9A=94=E7=A6=BB=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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) --- gateway/agent.py | 31 ++++++++++++ gateway/api.py | 18 +++++++ tests/test_session_fork.py | 98 ++++++++++++++++++++++++++++++++++++++ 3 files changed, 147 insertions(+) create mode 100644 tests/test_session_fork.py diff --git a/gateway/agent.py b/gateway/agent.py index 99bd8ef..bf6958e 100644 --- a/gateway/agent.py +++ b/gateway/agent.py @@ -15,6 +15,7 @@ from __future__ import annotations import asyncio +import copy import json import time import uuid @@ -889,6 +890,36 @@ class SessionStore: self.save(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: self._cache[sess.data["id"]] = sess self._save(sess) diff --git a/gateway/api.py b/gateway/api.py index 59788db..9a2fc1f 100644 --- a/gateway/api.py +++ b/gateway/api.py @@ -740,6 +740,24 @@ try: raise HTTPException(status_code=404, detail=f"会话不存在: {sid}") 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"]) async def agent_session_delete(sid: str): _check_id(sid, "会话 ID") diff --git a/tests/test_session_fork.py b/tests/test_session_fork.py new file mode 100644 index 0000000..a05ba52 --- /dev/null +++ b/tests/test_session_fork.py @@ -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