refactor(mini_swe_runner): dispatch env classes by name, flatten assistant-turn branches, drop body blanks; -33 LOC
This commit is contained in:
@@ -12,6 +12,7 @@ Usage:
|
|||||||
python mini_swe_runner.py --prompts_file prompts.jsonl --output_file trajectories.jsonl --env docker
|
python mini_swe_runner.py --prompts_file prompts.jsonl --output_file trajectories.jsonl --env docker
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import importlib
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -100,13 +101,10 @@ def create_environment(env_type: str = "local", image: str = "python:3.11-slim",
|
|||||||
if env_type == "local":
|
if env_type == "local":
|
||||||
from tools.environments.local import LocalEnvironment
|
from tools.environments.local import LocalEnvironment
|
||||||
return LocalEnvironment(cwd=cwd, timeout=timeout)
|
return LocalEnvironment(cwd=cwd, timeout=timeout)
|
||||||
if env_type == "docker":
|
if env_type not in ("docker", "modal"):
|
||||||
from tools.environments.docker import DockerEnvironment
|
raise ValueError(f"Unknown environment type: {env_type}. Use 'local', 'docker', or 'modal'")
|
||||||
return DockerEnvironment(image=image, cwd=cwd, timeout=timeout, **kwargs)
|
module = importlib.import_module(f"tools.environments.{env_type}")
|
||||||
if env_type == "modal":
|
return getattr(module, f"{env_type.capitalize()}Environment")(image=image, cwd=cwd, timeout=timeout, **kwargs)
|
||||||
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'")
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_json_args(raw: Any) -> Any:
|
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:
|
def _gpt_content(msg: Dict[str, Any], content: str) -> str:
|
||||||
"""Prefix ``content`` with a ``<think>`` block when the message carries reasoning."""
|
"""Prefix ``content`` with a ``<think>`` block when the message carries reasoning."""
|
||||||
think = f"<think>{msg['reasoning']}</think>" if msg.get("reasoning") else ""
|
return (f"<think>{msg['reasoning']}</think>" if msg.get("reasoning") else "") + content
|
||||||
return think + content
|
|
||||||
|
|
||||||
|
|
||||||
class MiniSWERunner:
|
class MiniSWERunner:
|
||||||
@@ -137,7 +134,6 @@ class MiniSWERunner:
|
|||||||
self.client = self._init_client(base_url, api_key)
|
self.client = self._init_client(base_url, api_key)
|
||||||
self.env = None # created per-task
|
self.env = None # created per-task
|
||||||
self.tools = [TERMINAL_TOOL_DEFINITION]
|
self.tools = [TERMINAL_TOOL_DEFINITION]
|
||||||
|
|
||||||
print("🤖 Mini-SWE Runner initialized")
|
print("🤖 Mini-SWE Runner initialized")
|
||||||
print(f" Model: {self.model}")
|
print(f" Model: {self.model}")
|
||||||
print(f" Environment: {self.env_type}")
|
print(f" Environment: {self.env_type}")
|
||||||
@@ -167,10 +163,9 @@ class MiniSWERunner:
|
|||||||
|
|
||||||
def _cleanup_env(self):
|
def _cleanup_env(self):
|
||||||
if self.env is not None:
|
if self.env is not None:
|
||||||
if hasattr(self.env, 'cleanup'):
|
stop = getattr(self.env, 'cleanup', None) or getattr(self.env, 'stop', None)
|
||||||
self.env.cleanup()
|
if stop:
|
||||||
elif hasattr(self.env, 'stop'):
|
stop()
|
||||||
self.env.stop()
|
|
||||||
self.env = None
|
self.env = None
|
||||||
|
|
||||||
def _execute_command(self, command: str, timeout: int = None) -> Dict[str, Any]:
|
def _execute_command(self, command: str, timeout: int = None) -> Dict[str, Any]:
|
||||||
@@ -212,36 +207,30 @@ class MiniSWERunner:
|
|||||||
"content": tool_content}, ensure_ascii=False)
|
"content": tool_content}, ensure_ascii=False)
|
||||||
tool_responses.append(f"<tool_response>\n{body}\n</tool_response>")
|
tool_responses.append(f"<tool_response>\n{body}\n</tool_response>")
|
||||||
j += 1
|
j += 1
|
||||||
if not tool_responses:
|
return ("\n".join(tool_responses), j - 1) if tool_responses else (None, i)
|
||||||
return None, i
|
|
||||||
return "\n".join(tool_responses), j - 1
|
|
||||||
|
|
||||||
def _convert_to_hermes_format(self, messages: List[Dict[str, Any]], user_query: str) -> List[Dict[str, Any]]:
|
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."""
|
"""Convert the OpenAI-style message list to the Hermes trajectory format used by batch_runner.py."""
|
||||||
system_msg = HERMES_SYSTEM_PREFIX + f"<tools>\n{self._format_tools_for_system_message()}\n</tools>\n" + HERMES_SYSTEM_SUFFIX
|
system_msg = HERMES_SYSTEM_PREFIX + f"<tools>\n{self._format_tools_for_system_message()}\n</tools>\n" + HERMES_SYSTEM_SUFFIX
|
||||||
trajectory = [{"from": "system", "value": system_msg}, {"from": "human", "value": user_query}]
|
trajectory = [{"from": "system", "value": system_msg}, {"from": "human", "value": user_query}]
|
||||||
|
|
||||||
i = 1 # first user message already added
|
i = 1 # first user message already added
|
||||||
while i < len(messages):
|
while i < len(messages):
|
||||||
msg = messages[i]
|
msg = messages[i]
|
||||||
if msg["role"] == "assistant":
|
if msg["role"] == "user":
|
||||||
if msg.get("tool_calls"):
|
trajectory.append({"from": "human", "value": msg["content"]})
|
||||||
content = (msg["content"] + "\n") if msg.get("content") else ""
|
elif msg["role"] == "assistant" and not msg.get("tool_calls"):
|
||||||
for tool_call in msg["tool_calls"]:
|
trajectory.append({"from": "gpt", "value": _gpt_content(msg, msg.get("content") or "")})
|
||||||
if not tool_call or not isinstance(tool_call, dict):
|
elif msg["role"] == "assistant":
|
||||||
continue
|
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"])}
|
tool_call_json = {"name": tool_call["function"]["name"], "arguments": _parse_json_args(tool_call["function"]["arguments"])}
|
||||||
content += f"<tool_call>\n{json.dumps(tool_call_json, ensure_ascii=False)}\n</tool_call>\n"
|
content += f"<tool_call>\n{json.dumps(tool_call_json, ensure_ascii=False)}\n</tool_call>\n"
|
||||||
trajectory.append({"from": "gpt", "value": _gpt_content(msg, content).rstrip()})
|
trajectory.append({"from": "gpt", "value": _gpt_content(msg, content).rstrip()})
|
||||||
tool_value, i = self._tool_response_turn(messages, i)
|
tool_value, i = self._tool_response_turn(messages, i)
|
||||||
if tool_value is not None:
|
if tool_value is not None:
|
||||||
trajectory.append({"from": "tool", "value": tool_value})
|
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"]})
|
|
||||||
i += 1
|
i += 1
|
||||||
|
|
||||||
return trajectory
|
return trajectory
|
||||||
|
|
||||||
def _call_model(self, messages: List[Dict[str, Any]]):
|
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
|
return self.client.chat.completions.create(**api_kwargs).choices[0].message
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error("API call failed: %s", e)
|
self.logger.error("API call failed: %s", e)
|
||||||
return None
|
|
||||||
|
|
||||||
def _run_tool_calls(self, assistant_message, messages: List[Dict[str, Any]]) -> bool:
|
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."""
|
"""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:
|
for tc in assistant_message.tool_calls:
|
||||||
args = _parse_json_args(tc.function.arguments)
|
args = _parse_json_args(tc.function.arguments)
|
||||||
command = args.get("command", "echo 'No command provided'")
|
command = args.get("command", "echo 'No command provided'")
|
||||||
timeout = args.get("timeout", self.command_timeout)
|
|
||||||
print(f" 📞 terminal: {command[:60]}...")
|
print(f" 📞 terminal: {command[:60]}...")
|
||||||
|
result = self._execute_command(command, args.get("timeout", self.command_timeout))
|
||||||
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)
|
|
||||||
if "MINI_SWE_AGENT_FINAL_OUTPUT" in result["output"]:
|
if "MINI_SWE_AGENT_FINAL_OUTPUT" in result["output"]:
|
||||||
print(" ✅ Task completion signal detected!")
|
print(" ✅ Task completion signal detected!")
|
||||||
completed = True
|
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")
|
print(f" ✅ exit_code={result['exit_code']}, output={len(result['output'])} chars")
|
||||||
return completed
|
return completed
|
||||||
|
|
||||||
@@ -287,58 +272,44 @@ class MiniSWERunner:
|
|||||||
print(f"\n{'='*60}")
|
print(f"\n{'='*60}")
|
||||||
print(f"📝 Task: {task[:80]}{'...' if len(task) > 80 else ''}")
|
print(f"📝 Task: {task[:80]}{'...' if len(task) > 80 else ''}")
|
||||||
print(f"{'='*60}")
|
print(f"{'='*60}")
|
||||||
|
|
||||||
self._create_env()
|
self._create_env()
|
||||||
messages = [{"role": "user", "content": task}]
|
messages = [{"role": "user", "content": task}]
|
||||||
api_call_count = 0
|
api_call_count = 0
|
||||||
completed = False
|
completed = False
|
||||||
|
|
||||||
try:
|
try:
|
||||||
while api_call_count < self.max_iterations:
|
while api_call_count < self.max_iterations:
|
||||||
api_call_count += 1
|
api_call_count += 1
|
||||||
print(f"\n🔄 API call #{api_call_count}/{self.max_iterations}")
|
print(f"\n🔄 API call #{api_call_count}/{self.max_iterations}")
|
||||||
|
|
||||||
assistant_message = self._call_model(messages)
|
assistant_message = self._call_model(messages)
|
||||||
if assistant_message is None:
|
if assistant_message is None:
|
||||||
break
|
break
|
||||||
if assistant_message.content:
|
if assistant_message.content:
|
||||||
print(f"🤖 Assistant: {assistant_message.content[:100]}...")
|
print(f"🤖 Assistant: {assistant_message.content[:100]}...")
|
||||||
|
if not assistant_message.tool_calls:
|
||||||
if assistant_message.tool_calls:
|
|
||||||
if self._run_tool_calls(assistant_message, messages):
|
|
||||||
completed = True
|
|
||||||
break
|
|
||||||
else:
|
|
||||||
messages.append({"role": "assistant", "content": assistant_message.content or ""})
|
messages.append({"role": "assistant", "content": assistant_message.content or ""})
|
||||||
completed = True
|
completed = True
|
||||||
print("🎉 Agent finished (no more tool calls)")
|
print("🎉 Agent finished (no more tool calls)")
|
||||||
break
|
break
|
||||||
|
if self._run_tool_calls(assistant_message, messages):
|
||||||
|
completed = True
|
||||||
|
break
|
||||||
if api_call_count >= self.max_iterations:
|
if api_call_count >= self.max_iterations:
|
||||||
print(f"⚠️ Reached max iterations ({self.max_iterations})")
|
print(f"⚠️ Reached max iterations ({self.max_iterations})")
|
||||||
finally:
|
finally:
|
||||||
self._cleanup_env()
|
self._cleanup_env()
|
||||||
|
return {"conversations": self._convert_to_hermes_format(messages, task), "completed": completed, "api_calls": api_call_count,
|
||||||
return {
|
"metadata": {"model": self.model, "env_type": self.env_type, "timestamp": datetime.now().isoformat()}}
|
||||||
"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]]:
|
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."""
|
"""Run every prompt, appending each result to ``output_file`` as it finishes."""
|
||||||
results = []
|
results = []
|
||||||
|
|
||||||
print(f"\n📦 Running batch of {len(prompts)} tasks")
|
print(f"\n📦 Running batch of {len(prompts)} tasks")
|
||||||
print(f"📁 Output: {output_file}")
|
print(f"📁 Output: {output_file}")
|
||||||
|
|
||||||
with open(output_file, 'w', encoding='utf-8') as f:
|
with open(output_file, 'w', encoding='utf-8') as f:
|
||||||
for i, prompt in enumerate(prompts, 1):
|
for i, prompt in enumerate(prompts, 1):
|
||||||
print(f"\n{'='*60}")
|
print(f"\n{'='*60}")
|
||||||
print(f"📋 Task {i}/{len(prompts)}")
|
print(f"📋 Task {i}/{len(prompts)}")
|
||||||
print(f"{'='*60}")
|
print(f"{'='*60}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
result = self.run_task(prompt)
|
result = self.run_task(prompt)
|
||||||
print(f"✅ Task {i} completed (api_calls={result['api_calls']})")
|
print(f"✅ Task {i} completed (api_calls={result['api_calls']})")
|
||||||
@@ -349,7 +320,6 @@ class MiniSWERunner:
|
|||||||
results.append(result)
|
results.append(result)
|
||||||
f.write(json.dumps(result, ensure_ascii=False) + "\n")
|
f.write(json.dumps(result, ensure_ascii=False) + "\n")
|
||||||
f.flush()
|
f.flush()
|
||||||
|
|
||||||
print(f"\n✅ Batch complete! {len(results)} trajectories saved to {output_file}")
|
print(f"\n✅ Batch complete! {len(results)} trajectories saved to {output_file}")
|
||||||
return results
|
return results
|
||||||
|
|
||||||
@@ -403,14 +373,11 @@ def main(
|
|||||||
"""
|
"""
|
||||||
print("🚀 Mini-SWE Runner with Hermes Trajectory Format")
|
print("🚀 Mini-SWE Runner with Hermes Trajectory Format")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
# Configure root logging at the entry point (not in library __init__).
|
# Configure root logging at the entry point (not in library __init__).
|
||||||
logging.basicConfig(level=logging.DEBUG if verbose else logging.INFO,
|
logging.basicConfig(level=logging.DEBUG if verbose else logging.INFO,
|
||||||
format='%(asctime)s - %(levelname)s - %(message)s', datefmt='%H:%M:%S')
|
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,
|
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)
|
max_iterations=max_iterations, command_timeout=timeout, verbose=verbose)
|
||||||
|
|
||||||
if task:
|
if task:
|
||||||
result = runner.run_task(task)
|
result = runner.run_task(task)
|
||||||
with open(output_file, 'w', encoding='utf-8') as f:
|
with open(output_file, 'w', encoding='utf-8') as f:
|
||||||
|
|||||||
Reference in New Issue
Block a user