diff --git a/slime/agent/trajectory.py b/slime/agent/trajectory.py index edb130c648..606eb4157a 100644 --- a/slime/agent/trajectory.py +++ b/slime/agent/trajectory.py @@ -19,6 +19,17 @@ logger = logging.getLogger(__name__) +# Prefix Claude Code injects as the first user message of a post-compaction +# request (context ran out -> history replaced by a summary). Used as the +# content signal for ``is_compact_start`` (see ``_classify_segment``). +COMPACT_SUMMARY_PREFIX = "This session is being continued from a previous conversation" + +# Tool-call names that spawn a sub-agent with a fresh context. The wire records +# the sub-agent's task text under ``input.prompt``; the manager sees it again as +# the first user message of the sub-agent's own segment, which is how a sub-agent +# is linked back to its caller (see ``_build_agent_prompt_index``). +SUBAGENT_TOOL_NAMES = ("Agent", "Task") + # =========================================================================== # TurnRecord @@ -76,6 +87,11 @@ def __init__( self.children: list[MessageNode] = [] self.turn: TurnRecord | None = None # the generated TurnRecord, else None (routing-only) self.turn_index: int | None = None + # Per-sid monotonic id assigned at mount time to EVERY node (routing-only + # and generated alike). Stable within one sid; reused as the unit a + # Sample's ``identity.node_list`` references and a sub-agent's + # ``caller_node_ids`` points at. The dummy root keeps ``None``. + self.node_id: int | None = None # Shared by sibling leaf paths; the first to reach it trains on it, the rest # re-emit it as loss_mask=0 context -- so each response is trained exactly once. self.response_trained: bool = False @@ -164,6 +180,11 @@ def __init__(self, fork_threshold: int) -> None: self.logprobs: list[float] = [] self.last_response_start_idx: int | None = None self.leading_prompt_len: int = 0 + # Generated assistant nodes packed into this builder, in append order. + # "Include = record": a node re-emitted as loss=0 context (claimed by an + # earlier sibling leaf) is still listed here, so ``node_list`` reflects + # what this Sample spans; reverse lookup picks the trainer (loss=1). + self.nodes: list[MessageNode] = [] def classify_token_drift(self, turn: TurnRecord) -> DriftKind: """Decide how this builder should absorb ``turn``'s prompt. @@ -212,6 +233,10 @@ def append_turn(self, turn: TurnRecord, kind: DriftKind, *, trained: bool = True if is_first_turn: self.leading_prompt_len = len(turn.prompt_ids) + def add_node(self, node: MessageNode) -> None: + """Record the generated node whose turn was just appended (for node_list).""" + self.nodes.append(node) + def _align_to_prompt(self, prompt_ids: list[int]) -> None: """Heal REALIGN drift by overwriting the most-recent response span with ``prompt_ids`` as loss_mask=0: the drifted tokens carry no signal, and re-appending @@ -230,10 +255,34 @@ def _append_tokens(self, ids: list[int], *, loss_mask: int, logprobs: list[float def has_trained_response(self) -> bool: return any(self.loss_mask[self.leading_prompt_len :]) + def _identity_metadata(self) -> dict[str, Any]: + """Build this Sample's ``identity`` block from the generated nodes it spans. + + The origin / compact / caller facts were resolved once per segment in + ``_classify_segments`` and stamped onto each generated node's + ``identity_segment``; here we just read the start node's segment and list + the node ids this builder packed. A builder always begins at a segment + start (a fork opens a fresh builder), so ``nodes[0]`` carries the + authoritative identity for the whole Sample. + """ + start = self.nodes[0] + seg = start.metadata.get("identity_segment", {}) + return { + "origin": seg.get("origin", "main"), + "is_compact_start": bool(seg.get("is_compact_start", False)), + "node_list": [n.node_id for n in self.nodes], + "start_node_id": start.node_id, + "caller_node_ids": list(seg.get("caller_node_ids", [])), + "match_kind": seg.get("match_kind"), + } + def to_sample(self, base_sample: Sample, extra_metadata: dict[str, Any] | None) -> Sample: """Emit the accumulated tokens as one ``Sample``, stripping the first-turn prompt so loss_mask / logprobs cover only the response region.""" start = self.leading_prompt_len # first-turn prompt stripped; response region starts here + metadata = dict(extra_metadata or {}) + if self.nodes: + metadata["identity"] = self._identity_metadata() return Sample( index=base_sample.index, group_index=base_sample.group_index, @@ -246,7 +295,7 @@ def to_sample(self, base_sample: Sample, extra_metadata: dict[str, Any] | None) rollout_log_probs=self.logprobs[start:], reward=0.0, status=Sample.Status.COMPLETED, - metadata=dict(extra_metadata or {}), + metadata=metadata, ) @@ -260,6 +309,7 @@ def __init__(self, *, fork_threshold_tokens: int | None = None) -> None: self._fork_threshold: int = 1024 if fork_threshold_tokens is None else fork_threshold_tokens self._trees: dict[str, MessageNode] = {} self._turn_count: dict[str, int] = {} + self._node_count: dict[str, int] = {} # per-sid monotonic node_id allocator # -------------------- public ------------------------------------------ @@ -290,7 +340,7 @@ def record_turn( node, depth = self._find_mount_point(root, prompt_messages) node, depth = self._try_merge_assistant_rewrite(sid, node, prompt_messages, depth) - node = self._mount_prompt_messages(node, prompt_messages[depth:]) + node = self._mount_prompt_messages(node, prompt_messages[depth:], sid=sid) self._attach_assistant_leaf(sid, node, turn=turn, response_message=response_message, metadata=metadata) def get_trajectory( @@ -312,6 +362,11 @@ def get_trajectory( if root is None: return [] + # Resolve per-segment identity (main / sub_agent / compact-start + caller) + # onto the generated nodes before draining, so each emitted Sample can read + # its start node's identity. Pure annotation -- never affects routing. + self._classify_segments(root) + samples: list[Sample] = [] for routing_leaf in root.leaves(): if routing_leaf.is_root: @@ -326,6 +381,7 @@ def get_trajectory( self._trees.pop(sid, None) self._turn_count.pop(sid, None) + self._node_count.pop(sid, None) return samples # -------------------- internals ---------------------------------------- @@ -406,13 +462,23 @@ def _try_merge_assistant_rewrite( rewritten_node.message = prompt_messages[depth] return rewritten_node, depth + 1 + def _next_node_id(self, sid: str) -> int: + """Allocate the next per-sid node_id (monotonic from 0).""" + nid = self._node_count.get(sid, 0) + self._node_count[sid] = nid + 1 + return nid + def _mount_prompt_messages( self, node: MessageNode, remaining_messages: list[dict[str, Any]], + *, + sid: str, ) -> MessageNode: for m in remaining_messages: - node = node.add_child(MessageNode(role=m.get("role"), message=m)) + child = MessageNode(role=m.get("role"), message=m) + child.node_id = self._next_node_id(sid) + node = node.add_child(child) return node def _attach_assistant_leaf( @@ -429,6 +495,7 @@ def _attach_assistant_leaf( message=response_message, metadata=dict(metadata or {}), ) + asst.node_id = self._next_node_id(sid) asst.turn = turn asst.turn_index = self._turn_count.get(sid, 0) + 1 node.add_child(asst) @@ -455,6 +522,7 @@ def _split_chain_into_builders(self, chain: list[MessageNode]) -> list[_SampleBu builders[-1].append_turn(asst_node.turn, DriftKind.CLEAN, trained=trained) else: builders[-1].append_turn(asst_node.turn, kind, trained=trained) + builders[-1].add_node(asst_node) return builders def _chain_to_samples( @@ -470,6 +538,173 @@ def _chain_to_samples( if builder.has_trained_response() ] + # -------------------- identity (segment classification) ---------------- + + def _classify_segments(self, root: MessageNode) -> None: + """Stamp each generated assistant node with its identity. + + Identity has three orthogonal facets, each from a distinct content signal + (the manager never sees SDK markers like ``parent_tool_use_id`` -- only the + translated prompt messages), resolved per generated node: + + * **origin** (``main`` / ``sub_agent``) -- from the SYSTEM prompt. The main + agent keeps one system prompt for the whole run; a sub-agent re-roots with + its own (e.g. an Explore specialist prompt). A node whose leading system + content differs from the main agent's (the earliest generated turn's) is + a sub-agent. This is the only origin signal that survives a compaction, + which severs the token path back to the sub-agent's first turn. + * **is_compact_start** -- True iff a user message opening this node's turn + starts with ``COMPACT_SUMMARY_PREFIX`` (context ran out, history replaced + by a summary). Per-node, not propagated: only the genuine post-compact + restart turn is a compact start, not later turns. + * **caller_node_ids** / **match_kind** -- for sub-agents, the node(s) whose + ``Agent``/``Task`` tool call spawned this run, matched by prompt text. The + spawning turn's first user message equals the agent-call prompt; later + turns inherit the caller from their nearest generated ancestor (so a + drift-fork mid-run keeps the link). A compaction breaks that path, so a + post-compact sub-agent turn keeps ``origin=sub_agent`` but loses the + caller (empty list). + + Pure annotation written under ``metadata['identity_segment']``; the routing + tree is untouched. Nodes are visited parent-before-child (pre-order) so + caller inheritance can read an already-resolved ancestor. + """ + gens = self._generated_nodes(root) # pre-order: parents before children + if not gens: + return + prompt_index = self._build_agent_prompt_index(gens) + main_system = self._system_content(min(gens, key=lambda n: n.turn_index or 0)) + + for gen in gens: + origin = "sub_agent" if self._system_content(gen) != main_system else "main" + is_compact_start = False + own_callers: list[int] | None = None + own_kind: str | None = None + for text in self._lead_in_user_texts(gen): + if text.startswith(COMPACT_SUMMARY_PREFIX): + is_compact_start = True + continue + callers, kind = self._match_agent_prompt(text, prompt_index) + if callers is not None: + own_callers, own_kind = callers, kind + + caller_node_ids: list[int] = [] + match_kind: str | None = None + if origin == "sub_agent": + if own_callers is not None: + caller_node_ids, match_kind = own_callers, own_kind + else: # inherit the caller from the nearest generated ancestor + parent_seg = self._parent_gen_segment(gen) + if parent_seg and parent_seg.get("origin") == "sub_agent": + caller_node_ids = list(parent_seg.get("caller_node_ids", [])) + match_kind = parent_seg.get("match_kind") + + gen.metadata["identity_segment"] = { + "origin": origin, + "is_compact_start": is_compact_start, + "caller_node_ids": caller_node_ids, + "match_kind": match_kind, + } + + @staticmethod + def _generated_nodes(root: MessageNode) -> list[MessageNode]: + """All generated assistant nodes (turn set) under root, pre-order.""" + out: list[MessageNode] = [] + stack = list(reversed(root.children)) + while stack: + n = stack.pop() + if n.role == "assistant" and n.turn is not None: + out.append(n) + stack.extend(reversed(n.children)) + return out + + @staticmethod + def _system_content(node: MessageNode) -> str: + """Leading system-message content on ``node``'s path (``""`` if none).""" + for n in node.path_from_root(): + if n.role == "system": + content = (n.message or {}).get("content") + return content if isinstance(content, str) else "" + return "" + + @staticmethod + def _parent_gen_segment(gen: MessageNode) -> dict[str, Any] | None: + """``identity_segment`` of the nearest generated ancestor, if resolved.""" + node = gen.parent + while node is not None and not node.is_root: + if node.role == "assistant" and node.turn is not None: + return node.metadata.get("identity_segment") + node = node.parent + return None + + def _build_agent_prompt_index(self, gens: list[MessageNode]) -> dict[str, list[int]]: + """Map each prior agent-call prompt text -> the node_ids that issued it. + + Scans every generated assistant node's ``message['tool_calls']`` for + ``Agent``/``Task`` calls and indexes their ``arguments['prompt']``. A + prompt may map to several caller node_ids (parallel fan-out reuses the + same prompt, or one node issues several agent calls), so the value is a + list -- the full candidate set, deduped, in node_id order. + """ + index: dict[str, list[int]] = {} + for node in gens: + msg = node.message or {} + for tc in msg.get("tool_calls") or []: + fn = tc.get("function") if isinstance(tc, dict) else None + if not isinstance(fn, dict) or fn.get("name") not in SUBAGENT_TOOL_NAMES: + continue + args = fn.get("arguments") + prompt = args.get("prompt") if isinstance(args, dict) else None + if not isinstance(prompt, str) or not prompt: + continue + callers = index.setdefault(prompt, []) + if node.node_id not in callers: + callers.append(node.node_id) + return index + + @staticmethod + def _lead_in_user_texts(gen: MessageNode) -> list[str]: + """User-message texts between ``gen`` and the previous generated node. + + These routing-only user nodes are what opened ``gen``'s turn -- the + compaction summary or the sub-agent task prompt land here. Walks up from + ``gen`` collecting user contents until it hits another generated assistant + (the previous turn's tail) or the root. + """ + texts: list[str] = [] + node = gen.parent + while node is not None and not node.is_root: + if node.role == "assistant" and node.turn is not None: + break + if node.role == "user": + content = (node.message or {}).get("content") + if isinstance(content, str) and content: + texts.append(content) + node = node.parent + return texts + + @staticmethod + def _match_agent_prompt(text: str, prompt_index: dict[str, list[int]]) -> tuple[list[int] | None, str | None]: + """Match a lead-in user text against the agent-call prompt index. + + Exact (byte-equal) match wins and is reported as ``"exact"``. Failing + that, a relaxed pass tolerates whitespace drift and prefix wrapping (cc + prepends a ```` block or trims trailing space): the + index prompt is accepted if it equals the text after stripping, or the + text starts with / contains the stripped prompt. Relaxed hits report + ``"approx"``. Returns ``(None, None)`` when nothing matches. + """ + if text in prompt_index: + return list(prompt_index[text]), "exact" + stripped = text.strip() + for prompt, callers in prompt_index.items(): + p = prompt.strip() + if not p: + continue + if stripped == p or stripped.startswith(p) or p in stripped: + return list(callers), "approx" + return None, None + __all__ = [ "TrajectoryManager", diff --git a/tests/test_agent/test_trajectory_manager_identity.py b/tests/test_agent/test_trajectory_manager_identity.py new file mode 100644 index 0000000000..664b9aeca9 --- /dev/null +++ b/tests/test_agent/test_trajectory_manager_identity.py @@ -0,0 +1,290 @@ +"""Identity-metadata tests for TrajectoryManager (origin / compact / caller). + +Standalone -- does NOT import the e2e dump helpers (absent on this branch). Each +test drives ``record_turn`` / ``get_trajectory`` directly and asserts on the +``sample.metadata['identity']`` block: + + {origin, is_compact_start, node_list, start_node_id, caller_node_ids, match_kind} + +Token ids are kept minimal; routing is driven by message-dict equality, and +linearization just needs each turn's prompt to prefix-extend the prior held +tokens (so segments stay CLEAN unless a test deliberately drifts). +""" + +from __future__ import annotations + +from slime.agent.trajectory import COMPACT_SUMMARY_PREFIX, TrajectoryManager, TurnRecord +from slime.utils.types import Sample + + +def _base() -> Sample: + return Sample(index=0, group_index=0, prompt="p", label="l") + + +def _agent_call(prompt: str) -> dict: + """An assistant message issuing an Agent tool call with ``prompt``.""" + return { + "role": "assistant", + "content": "", + "tool_calls": [{"type": "function", "function": {"name": "Agent", "arguments": {"prompt": prompt}}}], + } + + +def _rec(mgr, sid, prompt_messages, response_message, *, prompt_ids, output_ids): + mgr.record_turn( + sid, + turn=TurnRecord(prompt_ids=list(prompt_ids), output_ids=list(output_ids), finish_reason="stop"), + prompt_messages=prompt_messages, + response_message=response_message, + ) + + +def _identities(samples): + return [s.metadata["identity"] for s in samples] + + +# --------------------------------------------------------------------------- +# main-only +# --------------------------------------------------------------------------- + + +def test_main_only_single_segment(): + mgr = TrajectoryManager() + sid = "main" + sysm = {"role": "system", "content": "MAIN-SYS"} + _rec( + mgr, + sid, + [sysm, {"role": "user", "content": "do x"}], + {"role": "assistant", "content": "a1"}, + prompt_ids=[1, 2, 3], + output_ids=[10], + ) + _rec( + mgr, + sid, + [ + sysm, + {"role": "user", "content": "do x"}, + {"role": "assistant", "content": "a1"}, + {"role": "tool", "content": "r"}, + ], + {"role": "assistant", "content": "a2"}, + prompt_ids=[1, 2, 3, 10, 4], + output_ids=[11], + ) + samples = mgr.get_trajectory(sid, base_sample=_base()) + assert len(samples) == 1 + idn = samples[0].metadata["identity"] + assert idn["origin"] == "main" + assert idn["is_compact_start"] is False + assert idn["caller_node_ids"] == [] + assert idn["match_kind"] is None + assert idn["node_list"] == idn["node_list"] # present + assert idn["start_node_id"] == idn["node_list"][0] + assert len(idn["node_list"]) == 2 # two generated turns in one segment + + +# --------------------------------------------------------------------------- +# sub-agent (exact prompt match) + caller linkage +# --------------------------------------------------------------------------- + + +def test_sub_agent_exact_match_and_caller(): + mgr = TrajectoryManager() + sid = "sub" + main_sys = {"role": "system", "content": "MAIN-SYS"} + P = "Explore the repo layout" + + # turn 1: main agent issues an Agent tool call with prompt P + _rec(mgr, sid, [main_sys, {"role": "user", "content": "go"}], _agent_call(P), prompt_ids=[1, 2], output_ids=[10]) + # turn 2: sub-agent starts -- fresh system + user==P + sub_sys = {"role": "system", "content": "SUB-SYS-explore"} + _rec( + mgr, + sid, + [sub_sys, {"role": "user", "content": P}], + {"role": "assistant", "content": "s1"}, + prompt_ids=[20, 21], + output_ids=[30], + ) + + samples = mgr.get_trajectory(sid, base_sample=_base()) + idns = _identities(samples) + mains = [i for i in idns if i["origin"] == "main"] + subs = [i for i in idns if i["origin"] == "sub_agent"] + assert len(mains) == 1 and len(subs) == 1 + caller_node_id = mains[0]["node_list"][0] + sub = subs[0] + assert sub["match_kind"] == "exact" + assert sub["caller_node_ids"] == [caller_node_id] + assert sub["is_compact_start"] is False + + +# --------------------------------------------------------------------------- +# compact start +# --------------------------------------------------------------------------- + + +def test_compact_start_main(): + mgr = TrajectoryManager() + sid = "compact" + main_sys = {"role": "system", "content": "MAIN-SYS"} + _rec( + mgr, + sid, + [main_sys, {"role": "user", "content": "go"}], + {"role": "assistant", "content": "a1"}, + prompt_ids=[1, 2], + output_ids=[10], + ) + # compaction: fresh history, same system, user begins with the summary prefix + summary = COMPACT_SUMMARY_PREFIX + " ... summary body" + _rec( + mgr, + sid, + [main_sys, {"role": "user", "content": summary}], + {"role": "assistant", "content": "a2"}, + prompt_ids=[1, 50, 51], + output_ids=[60], + ) + + samples = mgr.get_trajectory(sid, base_sample=_base()) + idns = _identities(samples) + compacts = [i for i in idns if i["is_compact_start"]] + assert len(compacts) == 1 + assert compacts[0]["origin"] == "main" + # the pre-compact turn is NOT a compact start + assert any((not i["is_compact_start"]) and i["origin"] == "main" for i in idns) + + +# --------------------------------------------------------------------------- +# sub-agent that itself compacts (orthogonal) -- the turn25 shape +# --------------------------------------------------------------------------- + + +def test_sub_agent_internal_compact_orthogonal(): + mgr = TrajectoryManager() + sid = "subcompact" + main_sys = {"role": "system", "content": "MAIN-SYS"} + P = "deep dive task" + _rec(mgr, sid, [main_sys, {"role": "user", "content": "go"}], _agent_call(P), prompt_ids=[1, 2], output_ids=[10]) + sub_sys = {"role": "system", "content": "SUB-SYS-explore"} + # sub-agent post-compact restart: sub system + user begins with summary prefix. + summary = COMPACT_SUMMARY_PREFIX + " sub summary" + _rec( + mgr, + sid, + [sub_sys, {"role": "user", "content": summary}], + {"role": "assistant", "content": "s1"}, + prompt_ids=[20, 70, 71], + output_ids=[80], + ) + + samples = mgr.get_trajectory(sid, base_sample=_base()) + idns = _identities(samples) + sub = next(i for i in idns if i["origin"] == "sub_agent") + assert sub["origin"] == "sub_agent" + assert sub["is_compact_start"] is True + # compaction severs the token path -> caller link lost (empty), per spec + assert sub["caller_node_ids"] == [] + + +# --------------------------------------------------------------------------- +# parallel fan-out: same prompt issued by two caller nodes +# --------------------------------------------------------------------------- + + +def test_fanout_same_prompt_records_all_callers(): + mgr = TrajectoryManager() + sid = "fanout" + main_sys = {"role": "system", "content": "MAIN-SYS"} + P = "same fanout prompt" + # two main turns, each issuing the SAME agent prompt + _rec(mgr, sid, [main_sys, {"role": "user", "content": "go"}], _agent_call(P), prompt_ids=[1, 2], output_ids=[10]) + _rec( + mgr, + sid, + [main_sys, {"role": "user", "content": "go"}, _agent_call(P), {"role": "tool", "content": "r"}], + _agent_call(P), + prompt_ids=[1, 2, 10, 3], + output_ids=[11], + ) + # one sub-agent with prompt P -> should list BOTH caller node ids + sub_sys = {"role": "system", "content": "SUB-SYS"} + _rec( + mgr, + sid, + [sub_sys, {"role": "user", "content": P}], + {"role": "assistant", "content": "s"}, + prompt_ids=[20, 21], + output_ids=[30], + ) + + samples = mgr.get_trajectory(sid, base_sample=_base()) + sub = next(i for i in _identities(samples) if i["origin"] == "sub_agent") + assert sub["match_kind"] == "exact" + assert len(sub["caller_node_ids"]) == 2 + + +# --------------------------------------------------------------------------- +# approx fallback: trailing whitespace on the replayed prompt +# --------------------------------------------------------------------------- + + +def test_sub_agent_approx_match_whitespace(): + mgr = TrajectoryManager() + sid = "approx" + main_sys = {"role": "system", "content": "MAIN-SYS"} + P = "task body" + _rec(mgr, sid, [main_sys, {"role": "user", "content": "go"}], _agent_call(P), prompt_ids=[1, 2], output_ids=[10]) + sub_sys = {"role": "system", "content": "SUB-SYS"} + _rec( + mgr, + sid, + [sub_sys, {"role": "user", "content": P + " "}], + {"role": "assistant", "content": "s"}, + prompt_ids=[20, 21], + output_ids=[30], + ) + + samples = mgr.get_trajectory(sid, base_sample=_base()) + sub = next(i for i in _identities(samples) if i["origin"] == "sub_agent") + assert sub["match_kind"] == "approx" + assert len(sub["caller_node_ids"]) == 1 + + +# --------------------------------------------------------------------------- +# node_id uniqueness + cross-sid isolation +# --------------------------------------------------------------------------- + + +def test_node_ids_unique_and_per_sid(): + mgr = TrajectoryManager() + for sid in ("A", "B"): + sysm = {"role": "system", "content": "S"} + _rec( + mgr, + sid, + [sysm, {"role": "user", "content": "u"}], + {"role": "assistant", "content": "a"}, + prompt_ids=[1, 2], + output_ids=[10], + ) + # node_ids assigned monotonically from 0 within each sid + root = mgr._trees[sid] + ids = [n.node_id for n in _iter(root) if n.node_id is not None] + assert ids == sorted(ids) + assert ids[0] == 0 + assert len(set(ids)) == len(ids) + # draining A leaves B intact and B started its own node_id space at 0 + sa = mgr.get_trajectory("A", base_sample=_base()) + assert sa[0].metadata["identity"]["start_node_id"] == sa[0].metadata["identity"]["node_list"][0] + + +def _iter(root): + stack = list(root.children) + while stack: + n = stack.pop() + yield n + stack.extend(n.children) diff --git a/tests/test_agent/test_trajectory_manager_identity_replay.py b/tests/test_agent/test_trajectory_manager_identity_replay.py new file mode 100644 index 0000000000..db04a91223 --- /dev/null +++ b/tests/test_agent/test_trajectory_manager_identity_replay.py @@ -0,0 +1,100 @@ +"""Real-trajectory replay check for identity metadata. + +Replays the wire ``/v1/messages`` requests of a recorded SWE rollout +(``0610-e2e-test/runs/20260614_134732/0016``) through TrajectoryManager and +asserts the identity facets resolve as observed by hand-inspection: + +* the run mixes ``main`` and ``sub_agent`` origins; +* compaction is detected (``is_compact_start``) and is orthogonal to origin -- + some sub-agent turns are also post-compact restarts (the turn25 shape); +* at least one sub-agent turn links back to its caller node by exact prompt + match, and post-compact sub-agent turns keep ``origin=sub_agent`` but lose the + caller (compaction severs the token path). + +The recorded turns are REQUESTS only (no responses), so each turn's response is +reconstructed from the next request's replayed history -- imperfect (it triggers +benign rewrite-forks) but enough to exercise all three facets end-to-end. The +test skips if the run directory is absent so it stays green off the dev box. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +from slime.agent.adapters.anthropic import _fold_mid_list_system_into_user, _translate_messages +from slime.agent.adapters.common import tool_call_dict # noqa: F401 (kept: documents wire shape) +from slime.agent.trajectory import TrajectoryManager, TurnRecord +from slime.utils.types import Sample + +RUN = Path("/mnt/jingshenghang/code/slime_swe/0610-e2e-test/runs/20260614_134732/0016") + + +def _load_requests() -> list[list[dict]]: + """Translated chat-message lists for each real /v1/messages request.""" + out: list[list[dict]] = [] + for line in (RUN / "turns.jsonl").read_text().splitlines(): + t = json.loads(line) + p = t.get("payload") + if not isinstance(p, dict) or not (p.get("messages") or []): + continue + body = dict(p) + _fold_mid_list_system_into_user(body) + tr = _translate_messages(body.get("messages") or [], body.get("system")) + if tr: + out.append(tr) + return out + + +def _reply_revealed_by_next(prev: list[dict], cur: list[dict]) -> dict: + """The assistant reply to ``prev``'s turn, recovered from the next request.""" + n = 0 + while n < len(prev) and n < len(cur) and prev[n] == cur[n]: + n += 1 + for m in cur[n:]: + if m.get("role") == "assistant": + return m + for m in reversed(cur): + if m.get("role") == "assistant": + return m + return {"role": "assistant", "content": ""} + + +@pytest.mark.skipif(not RUN.exists(), reason="recorded run dir not present") +def test_replay_0016_identity_facets(): + reqs = _load_requests() + assert len(reqs) >= 10 + + mgr = TrajectoryManager() + sid = "cagent-0016" + for i, tr in enumerate(reqs): + rmsg = _reply_revealed_by_next(tr, reqs[i + 1]) if i + 1 < len(reqs) else {"role": "assistant", "content": ""} + prompt_ids = [1_000_000 + i * 1000 + j for j in range(len(tr) + 1)] + mgr.record_turn( + sid, + turn=TurnRecord(prompt_ids=prompt_ids, output_ids=[2_000_000 + i], finish_reason="stop"), + prompt_messages=tr, + response_message=rmsg, + ) + + samples = mgr.get_trajectory(sid, base_sample=Sample(index=0, group_index=0, prompt="p", label="l")) + idns = [s.metadata["identity"] for s in samples] + + origins = {i["origin"] for i in idns} + assert origins == {"main", "sub_agent"}, origins + + # compaction detected and orthogonal: some sub-agent turns are compact starts + assert any(i["is_compact_start"] for i in idns) + assert any(i["origin"] == "sub_agent" and i["is_compact_start"] for i in idns) + assert any(i["origin"] == "main" and not i["is_compact_start"] for i in idns) + + # caller linkage: at least one sub-agent turn matched its caller exactly + linked = [i for i in idns if i["origin"] == "sub_agent" and i["caller_node_ids"]] + assert linked, "no sub_agent sample linked to a caller node" + assert all(i["match_kind"] in ("exact", "approx") for i in linked) + + # post-compact sub-agent turns keep origin but drop the (severed) caller link + severed = [i for i in idns if i["origin"] == "sub_agent" and i["is_compact_start"]] + assert all(i["caller_node_ids"] == [] for i in severed)