diff --git a/hermes_state.py b/hermes_state.py index eab7011f2e..c8238a74de 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -6001,6 +6001,117 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) return None + # INSERT column list shared by append_message and append_messages_batch so + # the row shape can never drift between the single and batched writers. + _MESSAGE_INSERT_SQL = ( + "INSERT INTO messages (session_id, role, content, tool_call_id, " + "tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, " + "reasoning, reasoning_content, reasoning_details, codex_reasoning_items, " + "codex_message_items, platform_message_id, observed, active, api_content, display_kind, display_metadata) " + "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)" + ) + + def _prepare_message_row( + self, + *, + session_id: str, + role: str, + content: Any = None, + tool_name: Optional[str] = None, + tool_calls: Any = None, + tool_call_id: Optional[str] = None, + token_count: Optional[int] = None, + finish_reason: Optional[str] = None, + reasoning: Optional[str] = None, + reasoning_content: Optional[str] = None, + reasoning_details: Any = None, + codex_reasoning_items: Any = None, + codex_message_items: Any = None, + platform_message_id: Optional[str] = None, + observed: bool = False, + effect_disposition: Optional[str] = None, + timestamp: Any = None, + api_content: Optional[str] = None, + display_kind: Optional[str] = None, + display_metadata: Optional[Dict[str, Any]] = None, + ) -> Tuple[tuple, int]: + """Serialize one message into ``_MESSAGE_INSERT_SQL`` bind params. + + Runs entirely OUTSIDE the write transaction (JSON encoding, surrogate + scrubbing, timestamp coercion), so batched flushes keep the BEGIN + IMMEDIATE window as short as possible. Returns ``(params, + num_tool_calls)`` where ``num_tool_calls`` feeds the session counter + update. + """ + # Display metadata is presentation-only and never changes the model + # context role/content replayed to providers. + display_metadata_json = self._encode_display_metadata(display_metadata) + # Serialize structured fields to JSON before entering the write txn + reasoning_details_json = ( + json.dumps(reasoning_details) + if reasoning_details else None + ) + codex_items_json = ( + json.dumps(codex_reasoning_items) + if codex_reasoning_items else None + ) + codex_message_items_json = ( + json.dumps(codex_message_items) + if codex_message_items else None + ) + # tool_calls may arrive as a Python list (from the live agent) or + # as a JSON string (from import/export). Parse first to avoid + # double-encoding. + if isinstance(tool_calls, str): + try: + tool_calls = json.loads(tool_calls) + except (json.JSONDecodeError, TypeError): + tool_calls = [] + tool_calls_json = json.dumps(tool_calls) if tool_calls else None + # Multimodal content (list of parts) must be JSON-encoded: sqlite3 + # cannot bind list/dict parameters directly. + stored_content = self._encode_content(content) + + message_timestamp = time.time() + if timestamp is not None: + try: + if hasattr(timestamp, "timestamp"): + message_timestamp = float(timestamp.timestamp()) + else: + message_timestamp = float(timestamp) + except (TypeError, ValueError): + logger.debug("Ignoring invalid explicit message timestamp: %r", timestamp) + + # Pre-compute tool call count + num_tool_calls = 0 + if tool_calls is not None: + num_tool_calls = len(tool_calls) if isinstance(tool_calls, list) else 1 + + params = ( + session_id, + role, + stored_content, + tool_call_id, + tool_calls_json, + _scrub_surrogates(tool_name), + effect_disposition, + message_timestamp, + token_count, + finish_reason, + _scrub_surrogates(reasoning), + _scrub_surrogates(reasoning_content), + reasoning_details_json, + codex_items_json, + codex_message_items_json, + platform_message_id, + 1 if observed else 0, + 1, + _scrub_surrogates(api_content) if isinstance(api_content, str) else None, + _scrub_surrogates(display_kind) if isinstance(display_kind, str) else None, + display_metadata_json, + ) + return params, num_tool_calls + @staticmethod def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]: """Decode a ``display_metadata`` column into the dict every reader expects. @@ -6068,49 +6179,28 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) from every outgoing payload anyway, so the scrubbed form IS the wire bytes). """ - # Display metadata is presentation-only and never changes the model - # context role/content replayed to providers. - display_metadata_json = self._encode_display_metadata(display_metadata) - # Serialize structured fields to JSON before entering the write txn - reasoning_details_json = ( - json.dumps(reasoning_details) - if reasoning_details else None + row_params, num_tool_calls = self._prepare_message_row( + session_id=session_id, + role=role, + content=content, + tool_name=tool_name, + tool_calls=tool_calls, + tool_call_id=tool_call_id, + token_count=token_count, + finish_reason=finish_reason, + reasoning=reasoning, + reasoning_content=reasoning_content, + reasoning_details=reasoning_details, + codex_reasoning_items=codex_reasoning_items, + codex_message_items=codex_message_items, + platform_message_id=platform_message_id, + observed=observed, + effect_disposition=effect_disposition, + timestamp=timestamp, + api_content=api_content, + display_kind=display_kind, + display_metadata=display_metadata, ) - codex_items_json = ( - json.dumps(codex_reasoning_items) - if codex_reasoning_items else None - ) - codex_message_items_json = ( - json.dumps(codex_message_items) - if codex_message_items else None - ) - # tool_calls may arrive as a Python list (from the live agent) or - # as a JSON string (from import/export). Parse first to avoid - # double-encoding. - if isinstance(tool_calls, str): - try: - tool_calls = json.loads(tool_calls) - except (json.JSONDecodeError, TypeError): - tool_calls = [] - tool_calls_json = json.dumps(tool_calls) if tool_calls else None - # Multimodal content (list of parts) must be JSON-encoded: sqlite3 - # cannot bind list/dict parameters directly. - stored_content = self._encode_content(content) - - message_timestamp = time.time() - if timestamp is not None: - try: - if hasattr(timestamp, "timestamp"): - message_timestamp = float(timestamp.timestamp()) - else: - message_timestamp = float(timestamp) - except (TypeError, ValueError): - logger.debug("Ignoring invalid explicit message timestamp: %r", timestamp) - - # Pre-compute tool call count - num_tool_calls = 0 - if tool_calls is not None: - num_tool_calls = len(tool_calls) if isinstance(tool_calls, list) else 1 def _do(conn): active_lock = conn.execute( @@ -6135,36 +6225,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) and session["end_reason"] == "compression" ): raise CompressionSessionClosedError(session_id) - cursor = conn.execute( - """INSERT INTO messages (session_id, role, content, tool_call_id, - tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, - reasoning, reasoning_content, reasoning_details, codex_reasoning_items, - codex_message_items, platform_message_id, observed, active, api_content, display_kind, display_metadata) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", - ( - session_id, - role, - stored_content, - tool_call_id, - tool_calls_json, - _scrub_surrogates(tool_name), - effect_disposition, - message_timestamp, - token_count, - finish_reason, - _scrub_surrogates(reasoning), - _scrub_surrogates(reasoning_content), - reasoning_details_json, - codex_items_json, - codex_message_items_json, - platform_message_id, - 1 if observed else 0, - 1, - _scrub_surrogates(api_content) if isinstance(api_content, str) else None, - _scrub_surrogates(display_kind) if isinstance(display_kind, str) else None, - display_metadata_json, - ), - ) + cursor = conn.execute(self._MESSAGE_INSERT_SQL, row_params) msg_id = cursor.lastrowid # Update counters @@ -6190,6 +6251,113 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) _do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S ) + def append_messages_batch( + self, + session_id: str, + messages: List[Dict[str, Any]], + compression_lock_holder: Optional[str] = None, + ) -> List[int]: + """Append multiple messages atomically in ONE write transaction. + + ``messages`` is a list of dicts whose keys mirror + :meth:`append_message`'s keyword arguments (role, content, tool_name, + tool_calls, tool_call_id, finish_reason, reasoning, reasoning_content, + reasoning_details, codex_reasoning_items, codex_message_items, + timestamp, api_content, display_kind, display_metadata, ...). + + A turn-boundary flush writes the whole turn (user + assistant + tool + rows, typically 3-8 messages) as one BEGIN IMMEDIATE / commit pair + instead of one transaction (and, off WAL, one fsync) per row. Row + serialization happens OUTSIDE the transaction via + ``_prepare_message_row`` — the same helper ``append_message`` uses, + so the row shape cannot drift between the two writers. + + Atomicity contract: all rows land or none do (the caller re-flushes + unstamped messages on the next attempt). The compression-lock and + compression-closed guards from ``append_message`` run once for the + batch — same session, same instant. Returns the inserted row IDs in + input order. + """ + if not messages: + return [] + + prepared: List[tuple] = [] + total_tool_calls = 0 + for msg in messages: + role = msg.get("role", "unknown") + params, num_tc = self._prepare_message_row( + session_id=session_id, + role=role, + content=msg.get("content"), + tool_name=msg.get("tool_name"), + tool_calls=msg.get("tool_calls"), + tool_call_id=msg.get("tool_call_id"), + token_count=msg.get("token_count"), + finish_reason=msg.get("finish_reason"), + reasoning=msg.get("reasoning") if role == "assistant" else None, + reasoning_content=msg.get("reasoning_content") if role == "assistant" else None, + reasoning_details=msg.get("reasoning_details") if role == "assistant" else None, + codex_reasoning_items=msg.get("codex_reasoning_items") if role == "assistant" else None, + codex_message_items=msg.get("codex_message_items") if role == "assistant" else None, + platform_message_id=msg.get("platform_message_id"), + observed=bool(msg.get("observed")), + effect_disposition=msg.get("effect_disposition"), + timestamp=msg.get("timestamp"), + api_content=msg.get("api_content"), + display_kind=msg.get("display_kind"), + display_metadata=msg.get("display_metadata"), + ) + prepared.append(params) + total_tool_calls += num_tc + + def _do(conn): + active_lock = conn.execute( + "SELECT holder FROM compression_locks " + "WHERE session_id = ? AND expires_at > ?", + (session_id, time.time()), + ).fetchone() + if ( + active_lock is not None + and active_lock["holder"] != compression_lock_holder + ): + raise SessionCompressionInProgressError( + f"Session {session_id!r} is being compressed by another writer" + ) + session = conn.execute( + "SELECT ended_at, end_reason FROM sessions WHERE id = ?", + (session_id,), + ).fetchone() + if ( + session is not None + and session["ended_at"] is not None + and session["end_reason"] == "compression" + ): + raise CompressionSessionClosedError(session_id) + + row_ids: List[int] = [] + for params in prepared: + cursor = conn.execute(self._MESSAGE_INSERT_SQL, params) + row_ids.append(cursor.lastrowid) + + # One aggregated counter update for the whole batch. + if total_tool_calls > 0: + conn.execute( + """UPDATE sessions SET message_count = message_count + ?, + tool_call_count = tool_call_count + ? WHERE id = ?""", + (len(prepared), total_tool_calls, session_id), + ) + else: + conn.execute( + "UPDATE sessions SET message_count = message_count + ? WHERE id = ?", + (len(prepared), session_id), + ) + return row_ids + + # Same criticality as append_message: this IS the turn's transcript. + return self._execute_write( + _do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S + ) + def set_latest_matching_message_display_kind( self, session_id: str, *, role: str, content: str, display_kind: str, display_metadata: Optional[Dict[str, Any]] = None, diff --git a/run_agent.py b/run_agent.py index 953a127323..11e3d732d9 100644 --- a/run_agent.py +++ b/run_agent.py @@ -2096,6 +2096,10 @@ class AIAgent: ): _scan_start += 1 + # Collect this flush's new rows and write them in ONE transaction + # at the end of the scan (see append_messages_batch). + _batch_rows: List[Dict[str, Any]] = [] + _batch_msgs: List[Dict] = [] for _msg_idx in range(_scan_start, len(messages)): msg = messages[_msg_idx] if not isinstance(msg, dict): @@ -2214,33 +2218,46 @@ class AIAgent: ] elif isinstance(msg.get("tool_calls"), list): tool_calls_data = msg["tool_calls"] - self._session_db.append_message( - session_id=self.session_id, - role=role, - content=content, - tool_name=msg.get("tool_name"), - tool_calls=tool_calls_data, - tool_call_id=msg.get("tool_call_id"), - finish_reason=msg.get("finish_reason"), - reasoning=msg.get("reasoning") if role == "assistant" else None, - reasoning_content=msg.get("reasoning_content") if role == "assistant" else None, - reasoning_details=msg.get("reasoning_details") if role == "assistant" else None, - codex_reasoning_items=msg.get("codex_reasoning_items") if role == "assistant" else None, - codex_message_items=msg.get("codex_message_items") if role == "assistant" else None, - timestamp=_row_timestamp, - api_content=_row_api_content, - display_kind=( + _batch_rows.append({ + "role": role, + "content": content, + "tool_name": msg.get("tool_name"), + "tool_calls": tool_calls_data, + "tool_call_id": msg.get("tool_call_id"), + "finish_reason": msg.get("finish_reason"), + "reasoning": msg.get("reasoning") if role == "assistant" else None, + "reasoning_content": msg.get("reasoning_content") if role == "assistant" else None, + "reasoning_details": msg.get("reasoning_details") if role == "assistant" else None, + "codex_reasoning_items": msg.get("codex_reasoning_items") if role == "assistant" else None, + "codex_message_items": msg.get("codex_message_items") if role == "assistant" else None, + "timestamp": _row_timestamp, + "api_content": _row_api_content, + "display_kind": ( "hidden" if msg.get(COMPRESSED_SUMMARY_METADATA_KEY) and not msg.get("_compressed_summary_has_user_turn") else msg.get("display_kind") ), - display_metadata=msg.get("display_metadata"), + "display_metadata": msg.get("display_metadata"), + }) + _batch_msgs.append(msg) + # One transaction for the whole turn's new rows (typically 3-8 + # messages): one BEGIN IMMEDIATE / commit — and, off WAL, one + # fsync — instead of one per row. All-or-nothing pairs exactly + # with the marker stamping below: on failure NO rows landed and + # NO markers were stamped, so the next flush re-scans and + # re-writes the whole tail (same recovery contract as before, + # minus the partial-prefix case that could double-pay counters). + if _batch_rows: + self._session_db.append_messages_batch( + session_id=self.session_id, + messages=_batch_rows, compression_lock_holder=getattr( self, "_active_compression_lock_holder", None ), ) - msg[_DB_PERSISTED_MARKER] = True + for _written in _batch_msgs: + _written[_DB_PERSISTED_MARKER] = True # The intrinsic markers are now the sole source of truth. Reset the # one-shot seed so no id() outlives this flush to alias a message # allocated next turn at a recycled address. diff --git a/tests/agent/test_cursor_optimizations_parity.py b/tests/agent/test_cursor_optimizations_parity.py index c08b012aca..afb8e967bc 100644 --- a/tests/agent/test_cursor_optimizations_parity.py +++ b/tests/agent/test_cursor_optimizations_parity.py @@ -146,6 +146,12 @@ def test_parity_persist_bounded_scan(): self.rows = [] def append_message(self, **kw): self.rows.append({k: copy.deepcopy(v) for k, v in kw.items()}) + def append_messages_batch(self, session_id, messages, **kw): + for m in messages: + row = {k: copy.deepcopy(v) for k, v in m.items()} + row["session_id"] = session_id + self.rows.append(row) + return list(range(1, len(messages) + 1)) def make_agent(bounded): a = ra.AIAgent.__new__(ra.AIAgent) diff --git a/tests/hermes_state/test_append_messages_batch.py b/tests/hermes_state/test_append_messages_batch.py new file mode 100644 index 0000000000..aaff6884c8 --- /dev/null +++ b/tests/hermes_state/test_append_messages_batch.py @@ -0,0 +1,180 @@ +"""Tests for SessionDB.append_messages_batch (#23254 salvage). + +The batch writer must be row-shape-identical to append_message (shared +_prepare_message_row + _MESSAGE_INSERT_SQL), atomic (all rows or none), +and must aggregate the session counters in one UPDATE. +""" + +import json +import sqlite3 + +import pytest + +from hermes_state import ( + CompressionSessionClosedError, + SessionDB, +) + + +@pytest.fixture() +def db(tmp_path): + d = SessionDB(db_path=tmp_path / "state.db") + d.create_session("sess-batch", source="cli") + yield d + d.close() + + +def _turn_messages(): + return [ + {"role": "user", "content": "question"}, + { + "role": "assistant", + "content": "let me check", + "tool_calls": [{"name": "terminal", "arguments": "{}"}], + "reasoning_content": "thinking...", + "finish_reason": "tool_calls", + }, + { + "role": "tool", + "content": "tool output", + "tool_name": "terminal", + "tool_call_id": "call_1", + }, + {"role": "assistant", "content": "answer", "finish_reason": "stop"}, + ] + + +class TestAppendMessagesBatch: + def test_batch_rows_identical_to_single_appends(self, db, tmp_path): + """The batch writer stores the same bytes append_message would.""" + db2 = SessionDB(db_path=tmp_path / "state2.db") + db2.create_session("sess-batch", source="cli") + try: + msgs = _turn_messages() + db.append_messages_batch("sess-batch", msgs) + for m in msgs: + role = m["role"] + db2.append_message( + session_id="sess-batch", + role=role, + content=m.get("content"), + tool_name=m.get("tool_name"), + tool_calls=m.get("tool_calls"), + tool_call_id=m.get("tool_call_id"), + finish_reason=m.get("finish_reason"), + reasoning_content=( + m.get("reasoning_content") if role == "assistant" else None + ), + ) + cols = ( + "role, content, tool_call_id, tool_calls, tool_name, " + "finish_reason, reasoning_content, observed, active" + ) + rows_a = db._conn.execute( + f"SELECT {cols} FROM messages ORDER BY id" + ).fetchall() + rows_b = db2._conn.execute( + f"SELECT {cols} FROM messages ORDER BY id" + ).fetchall() + assert [tuple(r) for r in rows_a] == [tuple(r) for r in rows_b] + finally: + db2.close() + + def test_counters_aggregate_once(self, db): + db.append_messages_batch("sess-batch", _turn_messages()) + row = db._conn.execute( + "SELECT message_count, tool_call_count FROM sessions WHERE id = ?", + ("sess-batch",), + ).fetchone() + assert row["message_count"] == 4 + assert row["tool_call_count"] == 1 + + def test_returns_row_ids_in_input_order(self, db): + ids = db.append_messages_batch("sess-batch", _turn_messages()) + assert ids == sorted(ids) + assert len(ids) == 4 + + def test_empty_batch_is_noop(self, db): + assert db.append_messages_batch("sess-batch", []) == [] + row = db._conn.execute( + "SELECT message_count FROM sessions WHERE id = ?", ("sess-batch",) + ).fetchone() + assert row["message_count"] == 0 + + def test_atomicity_all_or_nothing(self, db, monkeypatch): + """A failure mid-batch leaves ZERO rows and untouched counters.""" + real_execute_write = db._execute_write + original_insert = SessionDB._MESSAGE_INSERT_SQL + + calls = {"n": 0} + + def _do_wrapper(fn, **kwargs): + def failing(conn): + real_conn_execute = conn.execute + + def exec_counting(sql, *args): + if sql == original_insert: + calls["n"] += 1 + if calls["n"] == 3: + raise sqlite3.OperationalError("boom mid-batch") + return real_conn_execute(sql, *args) + + conn.execute = exec_counting + try: + return fn(conn) + finally: + conn.execute = real_conn_execute + + return real_execute_write(failing, **kwargs) + + monkeypatch.setattr(db, "_execute_write", _do_wrapper) + with pytest.raises(sqlite3.OperationalError): + db.append_messages_batch("sess-batch", _turn_messages()) + monkeypatch.undo() + + count = db._conn.execute("SELECT COUNT(*) FROM messages").fetchone()[0] + assert count == 0 + row = db._conn.execute( + "SELECT message_count, tool_call_count FROM sessions WHERE id = ?", + ("sess-batch",), + ).fetchone() + assert row["message_count"] == 0 + assert row["tool_call_count"] == 0 + + def test_compression_closed_session_rejected(self, db): + db._conn.execute( + "UPDATE sessions SET ended_at = 1.0, end_reason = 'compression' " + "WHERE id = ?", + ("sess-batch",), + ) + db._conn.commit() + with pytest.raises(CompressionSessionClosedError): + db.append_messages_batch("sess-batch", _turn_messages()) + + def test_multimodal_content_encoded(self, db): + msgs = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + {"type": "image_url", "image_url": {"url": "data:x"}}, + ], + } + ] + db.append_messages_batch("sess-batch", msgs) + raw = db._conn.execute("SELECT content FROM messages").fetchone()[0] + # encoded via _encode_content — same sentinel prefix as append_message + loaded = db.get_messages("sess-batch") + assert loaded, raw + + def test_tool_calls_json_string_not_double_encoded(self, db): + msgs = [ + { + "role": "assistant", + "content": "x", + "tool_calls": json.dumps([{"name": "t", "arguments": "{}"}]), + } + ] + db.append_messages_batch("sess-batch", msgs) + raw = db._conn.execute("SELECT tool_calls FROM messages").fetchone()[0] + assert json.loads(raw) == [{"name": "t", "arguments": "{}"}] diff --git a/tests/run_agent/test_empty_response_recovery_persistence.py b/tests/run_agent/test_empty_response_recovery_persistence.py index ff8b9cf312..70e262fbf2 100644 --- a/tests/run_agent/test_empty_response_recovery_persistence.py +++ b/tests/run_agent/test_empty_response_recovery_persistence.py @@ -13,6 +13,12 @@ class _CapturingSessionDB: self.rows.append({"role": role, "content": content}) return len(self.rows) + def append_messages_batch(self, session_id, messages, **kwargs): + # Mirror the real batch writer: same rows, one call. + for m in messages: + self.rows.append({"role": m.get("role"), "content": m.get("content")}) + return list(range(len(self.rows) - len(messages) + 1, len(self.rows) + 1)) + def _agent_with_capturing_db(): agent = AIAgent.__new__(AIAgent) diff --git a/tests/run_agent/test_message_sequence_repair.py b/tests/run_agent/test_message_sequence_repair.py index b1a40e78dc..e169f1e068 100644 --- a/tests/run_agent/test_message_sequence_repair.py +++ b/tests/run_agent/test_message_sequence_repair.py @@ -227,6 +227,11 @@ def test_flush_guard_clamps_overshooting_cursor(): def append_message(self, **kw): self.rows.append(kw) + def append_messages_batch(self, session_id, messages, **kw): + for m in messages: + self.rows.append(dict(m, session_id=session_id)) + return list(range(1, len(messages) + 1)) + agent = _bare_agent() agent._session_db = _DB() agent._session_db_created = True diff --git a/tests/run_agent/test_tool_name_db_persistence.py b/tests/run_agent/test_tool_name_db_persistence.py index 3fcf7f33c3..29596c04dd 100644 --- a/tests/run_agent/test_tool_name_db_persistence.py +++ b/tests/run_agent/test_tool_name_db_persistence.py @@ -27,7 +27,8 @@ def _make_agent(session_db): def test_tool_name_persisted_to_session_db(): """tool_name set by make_tool_result_message must be passed through to - append_message so the column is populated on first flush to the session DB.""" + the batched flush so the column is populated on first write to the + session DB.""" session_db = MagicMock() agent = _make_agent(session_db) @@ -37,9 +38,8 @@ def test_tool_name_persisted_to_session_db(): ] agent._flush_messages_to_session_db(messages) - tool_appends = [ - c for c in session_db.append_message.call_args_list - if c.kwargs.get("role") == "tool" - ] - assert len(tool_appends) == 1 - assert tool_appends[0].kwargs["tool_name"] == "terminal" + assert session_db.append_messages_batch.call_count == 1 + batch = session_db.append_messages_batch.call_args.kwargs["messages"] + tool_rows = [m for m in batch if m.get("role") == "tool"] + assert len(tool_rows) == 1 + assert tool_rows[0]["tool_name"] == "terminal"