From 882b107070a28bbff38e807d96d7c3ccf1ac00e5 Mon Sep 17 00:00:00 2001 From: luijait Date: Mon, 16 Jun 2025 15:10:22 +0200 Subject: [PATCH] Improve chronological merge strategy --- src/cai/repl/commands/parallel.py | 115 +++++++++++++++++++++++------- uv.lock | 16 ++++- 2 files changed, 103 insertions(+), 28 deletions(-) diff --git a/src/cai/repl/commands/parallel.py b/src/cai/repl/commands/parallel.py index f0c05bdc..ee734736 100644 --- a/src/cai/repl/commands/parallel.py +++ b/src/cai/repl/commands/parallel.py @@ -951,36 +951,91 @@ class ParallelCommand(Command): def _merge_chronological( self, all_histories: dict[str, list], agents_to_merge: list[str] ) -> list[dict[str, Any]]: - """Merge histories chronologically by timestamp.""" - # Collect all messages with agent source - all_messages = [] - - # Create a global message counter for better chronological ordering - global_idx = 0 - - for agent_idx, agent_name in enumerate(agents_to_merge): + """Merge histories chronologically by interleaving messages based on conversation flow.""" + # Collect all messages with agent source and their indices + agent_messages = {} + for agent_name in agents_to_merge: history = all_histories.get(agent_name, []) - for local_idx, msg in enumerate(history): - # Create a copy of the message to avoid modifying original + agent_messages[agent_name] = [] + for idx, msg in enumerate(history): msg_copy = msg.copy() - # Add metadata about source agent and order msg_copy["_source_agent"] = agent_name - msg_copy["_original_index"] = local_idx - msg_copy["_agent_index"] = agent_idx - # Create a unique timestamp that considers both agent and message order - # This ensures messages from different agents are properly interleaved - if "_timestamp" not in msg_copy: - # Use global index to maintain proper chronological order - msg_copy["_timestamp"] = global_idx - global_idx += 1 - all_messages.append(msg_copy) + msg_copy["_original_index"] = idx + agent_messages[agent_name].append(msg_copy) - # Sort by timestamp to maintain chronological order - all_messages.sort(key=lambda x: x.get("_timestamp", 0)) + # Create indices to track position in each agent's history + indices = {agent: 0 for agent in agents_to_merge} + + # Process messages in an intelligent interleaved fashion + all_messages = [] + + while any(indices[agent] < len(agent_messages[agent]) for agent in agents_to_merge): + # Look for the next user message across all agents + next_user_msgs = [] + for agent in agents_to_merge: + if indices[agent] < len(agent_messages[agent]): + msg = agent_messages[agent][indices[agent]] + if msg.get("role") == "user": + next_user_msgs.append((agent, msg)) + + if next_user_msgs: + # Process the first user message found (they should be similar across agents) + chosen_agent, user_msg = next_user_msgs[0] + all_messages.append(user_msg) + indices[chosen_agent] += 1 + + # Skip duplicate user messages from other agents + for agent, msg in next_user_msgs[1:]: + if msg.get("content") == user_msg.get("content"): + indices[agent] += 1 + + # Now collect all responses to this user message from all agents + responses_collected = True + while responses_collected: + responses_collected = False + + for agent in agents_to_merge: + if indices[agent] < len(agent_messages[agent]): + msg = agent_messages[agent][indices[agent]] + + # Collect assistant responses and tool interactions until next user message + if msg.get("role") in ["assistant", "tool", "system"]: + all_messages.append(msg) + indices[agent] += 1 + responses_collected = True + + # If this is a tool call, look for the corresponding tool response + if msg.get("role") == "assistant" and msg.get("tool_calls"): + tool_call_ids = [tc.get("id") for tc in msg.get("tool_calls", [])] + + # Look ahead for tool responses + temp_idx = indices[agent] + while temp_idx < len(agent_messages[agent]): + next_msg = agent_messages[agent][temp_idx] + if next_msg.get("role") == "tool" and next_msg.get("tool_call_id") in tool_call_ids: + all_messages.append(next_msg) + indices[agent] = temp_idx + 1 + break + elif next_msg.get("role") == "user": + # Stop if we hit another user message + break + temp_idx += 1 + elif msg.get("role") == "user": + # Don't process user messages here - they'll be handled in the next iteration + break + else: + # No more user messages, collect any remaining messages + for agent in agents_to_merge: + if indices[agent] < len(agent_messages[agent]): + msg = agent_messages[agent][indices[agent]] + all_messages.append(msg) + indices[agent] += 1 + break # Process one at a time to maintain some order # Process messages to create the merged history merged = [] seen_tool_calls = {} # Track tool calls by ID to avoid duplicates + seen_messages = set() # Track message signatures to avoid duplicates # Debug: show total messages collected console.print(f"[dim]Total messages collected from all agents: {len(all_messages)}[/dim]") @@ -995,18 +1050,22 @@ class ParallelCommand(Command): for msg in all_messages: should_add = True + msg_sig = self._get_message_signature(msg) - # Check for duplicate messages - if msg.get("role") == "user": + # Check if we've already seen this exact message + if msg_sig and msg_sig in seen_messages: + should_add = False + + # Additional checks for specific message types + if should_add and msg.get("role") == "user": # For user messages, check if the same content was just added - # This helps when both agents received the same user input if ( merged and merged[-1].get("role") == "user" and merged[-1].get("content") == msg.get("content") ): should_add = False - elif msg.get("role") == "assistant" and msg.get("tool_calls"): + elif should_add and msg.get("role") == "assistant" and msg.get("tool_calls"): # For tool calls, track by tool call ID for tool_call in msg.get("tool_calls", []): tool_id = tool_call.get("id") @@ -1015,7 +1074,7 @@ class ParallelCommand(Command): should_add = False break seen_tool_calls[tool_id] = msg.get("_source_agent") - elif msg.get("role") == "tool": + elif should_add and msg.get("role") == "tool": # Tool responses should match their tool calls tool_call_id = msg.get("tool_call_id") if tool_call_id and tool_call_id in seen_tool_calls: @@ -1027,6 +1086,8 @@ class ParallelCommand(Command): # Clean up internal metadata before adding clean_msg = {k: v for k, v in msg.items() if not k.startswith("_")} merged.append(clean_msg) + if msg_sig: + seen_messages.add(msg_sig) return merged diff --git a/uv.lock b/uv.lock index 9e3898ab..a62e3c6f 100644 --- a/uv.lock +++ b/uv.lock @@ -281,7 +281,7 @@ wheels = [ [[package]] name = "cai-framework" -version = "0.4.0" +version = "0.5.0" source = { editable = "." } dependencies = [ { name = "dnspython" }, @@ -307,6 +307,7 @@ dependencies = [ { name = "paramiko" }, { name = "prompt-toolkit" }, { name = "pydantic" }, + { name = "pypdf2" }, { name = "requests" }, { name = "rich" }, { name = "types-requests" }, @@ -370,6 +371,7 @@ requires-dist = [ { name = "paramiko", specifier = ">=3.5.1" }, { name = "prompt-toolkit", specifier = ">=3.0.39" }, { name = "pydantic", specifier = ">=2.10,<3" }, + { name = "pypdf2" }, { name = "requests", specifier = ">=2.0,<3" }, { name = "rich", specifier = ">=13.9.4" }, { name = "types-requests", specifier = ">=2.0,<3" }, @@ -3180,6 +3182,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/05/e7/df2285f3d08fee213f2d041540fa4fc9ca6c2d44cf36d3a035bf2a8d2bcc/pyparsing-3.2.3-py3-none-any.whl", hash = "sha256:a749938e02d6fd0b59b356ca504a24982314bb090c383e3cf201c95ef7e2bfcf", size = 111120 }, ] +[[package]] +name = "pypdf2" +version = "3.0.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions", marker = "python_full_version < '3.10'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/9f/bb/18dc3062d37db6c491392007dfd1a7f524bb95886eb956569ac38a23a784/PyPDF2-3.0.1.tar.gz", hash = "sha256:a74408f69ba6271f71b9352ef4ed03dc53a31aa404d29b5d31f53bfecfee1440", size = 227419 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/8e/5e/c86a5643653825d3c913719e788e41386bee415c2b87b4f955432f2de6b2/pypdf2-3.0.1-py3-none-any.whl", hash = "sha256:d16e4205cfee272fbdc0568b68d82be796540b1537508cef59388f839c191928", size = 232572 }, +] + [[package]] name = "pytest" version = "8.3.5"