"""会话树 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