From 8e207b82e5e41ee2ccf8d9e7d2ae95b41d0cd390 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 21:31:03 -0700 Subject: [PATCH] refactor(mini_swe_runner): dispatch env classes by name, flatten assistant-turn branches, drop body blanks; -33 LOC --- mini_swe_runner.py | 93 +++++++++++++++------------------------------- 1 file changed, 30 insertions(+), 63 deletions(-) diff --git a/mini_swe_runner.py b/mini_swe_runner.py index fc7dc5099d..62fb8179a8 100644 --- a/mini_swe_runner.py +++ b/mini_swe_runner.py @@ -12,6 +12,7 @@ Usage: python mini_swe_runner.py --prompts_file prompts.jsonl --output_file trajectories.jsonl --env docker """ +import importlib import json import logging import os @@ -100,13 +101,10 @@ def create_environment(env_type: str = "local", image: str = "python:3.11-slim", if env_type == "local": from tools.environments.local import LocalEnvironment return LocalEnvironment(cwd=cwd, timeout=timeout) - if env_type == "docker": - from tools.environments.docker import DockerEnvironment - return DockerEnvironment(image=image, cwd=cwd, timeout=timeout, **kwargs) - if env_type == "modal": - from tools.environments.modal import ModalEnvironment - return ModalEnvironment(image=image, cwd=cwd, timeout=timeout, **kwargs) - raise ValueError(f"Unknown environment type: {env_type}. Use 'local', 'docker', or 'modal'") + if env_type not in ("docker", "modal"): + raise ValueError(f"Unknown environment type: {env_type}. Use 'local', 'docker', or 'modal'") + module = importlib.import_module(f"tools.environments.{env_type}") + return getattr(module, f"{env_type.capitalize()}Environment")(image=image, cwd=cwd, timeout=timeout, **kwargs) def _parse_json_args(raw: Any) -> Any: @@ -121,8 +119,7 @@ def _parse_json_args(raw: Any) -> Any: def _gpt_content(msg: Dict[str, Any], content: str) -> str: """Prefix ``content`` with a ```` block when the message carries reasoning.""" - think = f"{msg['reasoning']}" if msg.get("reasoning") else "" - return think + content + return (f"{msg['reasoning']}" if msg.get("reasoning") else "") + content class MiniSWERunner: @@ -137,7 +134,6 @@ class MiniSWERunner: self.client = self._init_client(base_url, api_key) self.env = None # created per-task self.tools = [TERMINAL_TOOL_DEFINITION] - print("šŸ¤– Mini-SWE Runner initialized") print(f" Model: {self.model}") print(f" Environment: {self.env_type}") @@ -167,10 +163,9 @@ class MiniSWERunner: def _cleanup_env(self): if self.env is not None: - if hasattr(self.env, 'cleanup'): - self.env.cleanup() - elif hasattr(self.env, 'stop'): - self.env.stop() + stop = getattr(self.env, 'cleanup', None) or getattr(self.env, 'stop', None) + if stop: + stop() self.env = None def _execute_command(self, command: str, timeout: int = None) -> Dict[str, Any]: @@ -212,36 +207,30 @@ class MiniSWERunner: "content": tool_content}, ensure_ascii=False) tool_responses.append(f"\n{body}\n") j += 1 - if not tool_responses: - return None, i - return "\n".join(tool_responses), j - 1 + return ("\n".join(tool_responses), j - 1) if tool_responses else (None, i) def _convert_to_hermes_format(self, messages: List[Dict[str, Any]], user_query: str) -> List[Dict[str, Any]]: """Convert the OpenAI-style message list to the Hermes trajectory format used by batch_runner.py.""" system_msg = HERMES_SYSTEM_PREFIX + f"\n{self._format_tools_for_system_message()}\n\n" + HERMES_SYSTEM_SUFFIX trajectory = [{"from": "system", "value": system_msg}, {"from": "human", "value": user_query}] - i = 1 # first user message already added while i < len(messages): msg = messages[i] - if msg["role"] == "assistant": - if msg.get("tool_calls"): - content = (msg["content"] + "\n") if msg.get("content") else "" - for tool_call in msg["tool_calls"]: - if not tool_call or not isinstance(tool_call, dict): - continue + if msg["role"] == "user": + trajectory.append({"from": "human", "value": msg["content"]}) + elif msg["role"] == "assistant" and not msg.get("tool_calls"): + trajectory.append({"from": "gpt", "value": _gpt_content(msg, msg.get("content") or "")}) + elif msg["role"] == "assistant": + content = (msg["content"] + "\n") if msg.get("content") else "" + for tool_call in msg["tool_calls"]: + if isinstance(tool_call, dict) and tool_call: tool_call_json = {"name": tool_call["function"]["name"], "arguments": _parse_json_args(tool_call["function"]["arguments"])} content += f"\n{json.dumps(tool_call_json, ensure_ascii=False)}\n\n" - trajectory.append({"from": "gpt", "value": _gpt_content(msg, content).rstrip()}) - tool_value, i = self._tool_response_turn(messages, i) - if tool_value is not None: - trajectory.append({"from": "tool", "value": tool_value}) - else: - trajectory.append({"from": "gpt", "value": _gpt_content(msg, msg.get("content") or "")}) - elif msg["role"] == "user": - trajectory.append({"from": "human", "value": msg["content"]}) + trajectory.append({"from": "gpt", "value": _gpt_content(msg, content).rstrip()}) + tool_value, i = self._tool_response_turn(messages, i) + if tool_value is not None: + trajectory.append({"from": "tool", "value": tool_value}) i += 1 - return trajectory def _call_model(self, messages: List[Dict[str, Any]]): @@ -256,7 +245,6 @@ class MiniSWERunner: return self.client.chat.completions.create(**api_kwargs).choices[0].message except Exception as e: self.logger.error("API call failed: %s", e) - return None def _run_tool_calls(self, assistant_message, messages: List[Dict[str, Any]]) -> bool: """Record the assistant turn, execute each terminal call, append results; True if the completion signal fired.""" @@ -270,15 +258,12 @@ class MiniSWERunner: for tc in assistant_message.tool_calls: args = _parse_json_args(tc.function.arguments) command = args.get("command", "echo 'No command provided'") - timeout = args.get("timeout", self.command_timeout) print(f" šŸ“ž terminal: {command[:60]}...") - - result = self._execute_command(command, timeout) - result_json = json.dumps({"content": {"output": result["output"], "exit_code": result["exit_code"], "error": result["error"]}}, ensure_ascii=False) + result = self._execute_command(command, args.get("timeout", self.command_timeout)) if "MINI_SWE_AGENT_FINAL_OUTPUT" in result["output"]: print(" āœ… Task completion signal detected!") completed = True - messages.append(make_tool_result_message(tc.function.name, result_json, tc.id)) + messages.append(make_tool_result_message(tc.function.name, json.dumps({"content": result}, ensure_ascii=False), tc.id)) print(f" āœ… exit_code={result['exit_code']}, output={len(result['output'])} chars") return completed @@ -287,58 +272,44 @@ class MiniSWERunner: print(f"\n{'='*60}") print(f"šŸ“ Task: {task[:80]}{'...' if len(task) > 80 else ''}") print(f"{'='*60}") - self._create_env() messages = [{"role": "user", "content": task}] api_call_count = 0 completed = False - try: while api_call_count < self.max_iterations: api_call_count += 1 print(f"\nšŸ”„ API call #{api_call_count}/{self.max_iterations}") - assistant_message = self._call_model(messages) if assistant_message is None: break if assistant_message.content: print(f"šŸ¤– Assistant: {assistant_message.content[:100]}...") - - if assistant_message.tool_calls: - if self._run_tool_calls(assistant_message, messages): - completed = True - break - else: + if not assistant_message.tool_calls: messages.append({"role": "assistant", "content": assistant_message.content or ""}) completed = True print("šŸŽ‰ Agent finished (no more tool calls)") break - + if self._run_tool_calls(assistant_message, messages): + completed = True + break if api_call_count >= self.max_iterations: print(f"āš ļø Reached max iterations ({self.max_iterations})") finally: self._cleanup_env() - - return { - "conversations": self._convert_to_hermes_format(messages, task), - "completed": completed, - "api_calls": api_call_count, - "metadata": {"model": self.model, "env_type": self.env_type, "timestamp": datetime.now().isoformat()}, - } + return {"conversations": self._convert_to_hermes_format(messages, task), "completed": completed, "api_calls": api_call_count, + "metadata": {"model": self.model, "env_type": self.env_type, "timestamp": datetime.now().isoformat()}} def run_batch(self, prompts: List[str], output_file: str) -> List[Dict[str, Any]]: """Run every prompt, appending each result to ``output_file`` as it finishes.""" results = [] - print(f"\nšŸ“¦ Running batch of {len(prompts)} tasks") print(f"šŸ“ Output: {output_file}") - with open(output_file, 'w', encoding='utf-8') as f: for i, prompt in enumerate(prompts, 1): print(f"\n{'='*60}") print(f"šŸ“‹ Task {i}/{len(prompts)}") print(f"{'='*60}") - try: result = self.run_task(prompt) print(f"āœ… Task {i} completed (api_calls={result['api_calls']})") @@ -349,7 +320,6 @@ class MiniSWERunner: results.append(result) f.write(json.dumps(result, ensure_ascii=False) + "\n") f.flush() - print(f"\nāœ… Batch complete! {len(results)} trajectories saved to {output_file}") return results @@ -403,14 +373,11 @@ def main( """ print("šŸš€ Mini-SWE Runner with Hermes Trajectory Format") print("=" * 60) - # Configure root logging at the entry point (not in library __init__). logging.basicConfig(level=logging.DEBUG if verbose else logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s', datefmt='%H:%M:%S') - runner = MiniSWERunner(model=model, base_url=base_url, api_key=api_key, env_type=env, image=image, cwd=cwd, max_iterations=max_iterations, command_timeout=timeout, verbose=verbose) - if task: result = runner.run_task(task) with open(output_file, 'w', encoding='utf-8') as f: