|
|
@@ -14,6 +14,7 @@ from agent_lab.application.runtime import DebugRuntime
|
|
|
from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
|
|
|
from agent_lab.domain.events import ToolCallEvent
|
|
|
from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
|
|
|
+from agent_lab.infrastructure.chat_client import OpenAICompatibleChatClient
|
|
|
from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
|
|
|
|
|
|
|
|
|
@@ -1610,6 +1611,24 @@ class ScriptedChatClient:
|
|
|
yield item
|
|
|
|
|
|
|
|
|
+class TranscriptValidatingScriptedChatClient(ScriptedChatClient):
|
|
|
+ async def stream_chat(
|
|
|
+ self,
|
|
|
+ messages: list[ChatMessage],
|
|
|
+ tools: list[dict],
|
|
|
+ params: AgentParams,
|
|
|
+ tool_choice: dict[str, Any] | None = None,
|
|
|
+ ) -> AsyncIterator[StreamItem]:
|
|
|
+ OpenAICompatibleChatClient._validate_tool_transcript(self, messages)
|
|
|
+ async for item in super().stream_chat(
|
|
|
+ messages,
|
|
|
+ tools,
|
|
|
+ params,
|
|
|
+ tool_choice,
|
|
|
+ ):
|
|
|
+ yield item
|
|
|
+
|
|
|
+
|
|
|
def _direct_request(
|
|
|
*,
|
|
|
enabled_tools: list[str],
|
|
|
@@ -1969,6 +1988,27 @@ async def test_direct_mode_session_turns_reset_budget_and_snapshot_mode(tmp_path
|
|
|
session = store.get_session("direct-session")
|
|
|
assert session is not None
|
|
|
assert session["config"]["tool_invocation_mode"] == "chat_agent_tools"
|
|
|
+ persisted_messages = store.list_messages("direct-session")
|
|
|
+ assert [message["role"] for message in persisted_messages] == [
|
|
|
+ "user",
|
|
|
+ "assistant",
|
|
|
+ "tool",
|
|
|
+ "assistant",
|
|
|
+ "user",
|
|
|
+ "assistant",
|
|
|
+ "tool",
|
|
|
+ "assistant",
|
|
|
+ ]
|
|
|
+ assert [
|
|
|
+ call["id"]
|
|
|
+ for message in persisted_messages
|
|
|
+ for call in message["tool_calls"]
|
|
|
+ ] == ["turn-1-call", "turn-2-call"]
|
|
|
+ assert [
|
|
|
+ message["tool_call_id"]
|
|
|
+ for message in persisted_messages
|
|
|
+ if message["role"] == "tool"
|
|
|
+ ] == ["turn-1-call", "turn-2-call"]
|
|
|
request_audits = [
|
|
|
audit
|
|
|
for audit in store.list_audit_logs("direct-session")
|
|
|
@@ -1981,6 +2021,139 @@ async def test_direct_mode_session_turns_reset_budget_and_snapshot_mode(tmp_path
|
|
|
)
|
|
|
|
|
|
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_direct_mode_reconnect_restores_provider_valid_session_transcript(tmp_path):
|
|
|
+ store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
|
|
|
+ session_id = store.create_session(
|
|
|
+ title="restored direct",
|
|
|
+ config={"tool_invocation_mode": "chat_agent_tools"},
|
|
|
+ session_id="restored-direct",
|
|
|
+ )
|
|
|
+ store.start_turn(session_id, turn_index=1, user_message="past user")
|
|
|
+ previous_call = ToolCallEvent(
|
|
|
+ id="past-call",
|
|
|
+ name="safe_tool",
|
|
|
+ arguments={"value": 1},
|
|
|
+ raw_arguments='{"value":1}',
|
|
|
+ )
|
|
|
+ for message in [
|
|
|
+ ChatMessage(role="user", content="past user"),
|
|
|
+ ChatMessage(role="assistant", content="", tool_calls=[previous_call]),
|
|
|
+ ChatMessage(
|
|
|
+ role="tool",
|
|
|
+ content='{"ok":true}',
|
|
|
+ name="safe_tool",
|
|
|
+ tool_call_id="past-call",
|
|
|
+ ),
|
|
|
+ ChatMessage(role="assistant", content="past answer"),
|
|
|
+ ]:
|
|
|
+ store.append_message(session_id, turn_index=1, message=message)
|
|
|
+ store.complete_turn(session_id, turn_index=1)
|
|
|
+ client = TranscriptValidatingScriptedChatClient(
|
|
|
+ [[StreamItem.message_delta("current answer")]]
|
|
|
+ )
|
|
|
+ runtime = DebugRuntime(
|
|
|
+ client,
|
|
|
+ registry=_single_tool_registry(lambda event: {"ok": True}),
|
|
|
+ session_store=store,
|
|
|
+ )
|
|
|
+ request = _direct_request(
|
|
|
+ enabled_tools=["safe_tool"],
|
|
|
+ session_id=session_id,
|
|
|
+ ).model_copy(
|
|
|
+ update={
|
|
|
+ "user_message": "current user",
|
|
|
+ "pre_messages": [ChatMessage(role="user", content="configured pre")],
|
|
|
+ }
|
|
|
+ )
|
|
|
+
|
|
|
+ queues = runtime.start_session(request)
|
|
|
+ outputs: list[dict[str, Any]] = []
|
|
|
+ while not any(message["type"] == "turn_completed" for message in outputs):
|
|
|
+ outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
|
|
|
+ await runtime.aclose()
|
|
|
+
|
|
|
+ transcript = client.messages_by_call[0]
|
|
|
+ assert [(message.role, message.content) for message in transcript] == [
|
|
|
+ ("system", "You are a debugger."),
|
|
|
+ ("user", "configured pre"),
|
|
|
+ ("user", "past user"),
|
|
|
+ ("assistant", ""),
|
|
|
+ ("tool", '{"ok":true}'),
|
|
|
+ ("assistant", "past answer"),
|
|
|
+ ("user", "current user"),
|
|
|
+ ]
|
|
|
+ assert transcript[3].tool_calls == [previous_call]
|
|
|
+ assert sum(
|
|
|
+ message.role == "user" and message.content == "current user"
|
|
|
+ for message in transcript
|
|
|
+ ) == 1
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_dual_mode_reconnect_restores_visible_history_without_tool_protocol(tmp_path):
|
|
|
+ store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
|
|
|
+ session_id = store.create_session(
|
|
|
+ title="restored dual",
|
|
|
+ config={"tool_invocation_mode": "dual_agent"},
|
|
|
+ session_id="restored-dual",
|
|
|
+ )
|
|
|
+ store.start_turn(session_id, turn_index=1, user_message="past user")
|
|
|
+ store.append_message(
|
|
|
+ session_id,
|
|
|
+ turn_index=1,
|
|
|
+ message=ChatMessage(role="user", content="past user"),
|
|
|
+ )
|
|
|
+ store.append_message(
|
|
|
+ session_id,
|
|
|
+ turn_index=1,
|
|
|
+ message=ChatMessage(role="assistant", content="past answer"),
|
|
|
+ )
|
|
|
+ store.complete_turn(session_id, turn_index=1)
|
|
|
+ registry = _single_tool_registry(lambda event: {"ok": True})
|
|
|
+ client = ScriptedChatClient(
|
|
|
+ [
|
|
|
+ [StreamItem.text_event(_tool_call("dual-call"))],
|
|
|
+ [StreamItem.message_delta("current answer")],
|
|
|
+ ]
|
|
|
+ )
|
|
|
+ runtime = DebugRuntime(client, registry=registry, session_store=store)
|
|
|
+ request = DebugRunRequest(
|
|
|
+ session_id=session_id,
|
|
|
+ user_message="current user",
|
|
|
+ system_prompts=["system context"],
|
|
|
+ pre_messages=[ChatMessage(role="assistant", content="configured pre")],
|
|
|
+ chat_agent=AgentParams(model="chat-model"),
|
|
|
+ event_agent=EventAgentParams(
|
|
|
+ enabled_tools=["safe_tool"],
|
|
|
+ max_event_loops=1,
|
|
|
+ ),
|
|
|
+ tool_invocation_mode="dual_agent",
|
|
|
+ )
|
|
|
+
|
|
|
+ queues = runtime.start_session(request)
|
|
|
+ outputs: list[dict[str, Any]] = []
|
|
|
+ while not any(message["type"] == "turn_completed" for message in outputs):
|
|
|
+ outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
|
|
|
+ await runtime.aclose()
|
|
|
+
|
|
|
+ first_request = client.messages_by_call[0]
|
|
|
+ assert [(message.role, message.content) for message in first_request if message.role != "system"] == [
|
|
|
+ ("assistant", "configured pre"),
|
|
|
+ ("user", "past user"),
|
|
|
+ ("assistant", "past answer"),
|
|
|
+ ("user", "current user"),
|
|
|
+ ]
|
|
|
+ assert sum(
|
|
|
+ message.role == "user" and message.content == "current user"
|
|
|
+ for message in first_request
|
|
|
+ ) == 1
|
|
|
+ persisted_messages = store.list_messages(session_id)
|
|
|
+ assert not any(message["role"] == "tool" for message in persisted_messages)
|
|
|
+ assert not any(message["tool_calls"] for message in persisted_messages)
|
|
|
+ assert not any(message["content"] == "" for message in persisted_messages)
|
|
|
+
|
|
|
+
|
|
|
def _single_tool_registry(
|
|
|
handler: Any,
|
|
|
*,
|