|
| 1 | +"""mini-SWE-agent adapter for lagent.""" |
| 2 | + |
| 3 | +from __future__ import annotations |
| 4 | + |
| 5 | +import asyncio |
| 6 | +import os |
| 7 | +from pathlib import Path |
| 8 | +from typing import Any, Optional |
| 9 | + |
| 10 | +from .base import AsyncExternalAgent |
| 11 | + |
| 12 | + |
| 13 | +class MiniSWEAgentAdapter(AsyncExternalAgent): |
| 14 | + """Wrap mini-SWE-agent as a lagent async external agent.""" |
| 15 | + |
| 16 | + def __init__( |
| 17 | + self, |
| 18 | + model: str, |
| 19 | + api_base: Optional[str] = None, |
| 20 | + api_key: Optional[str] = None, |
| 21 | + cwd: Optional[str] = None, |
| 22 | + step_limit: int = 30, |
| 23 | + command_timeout: int = 300, |
| 24 | + trajectory_path: str = "/tmp/mini.traj.json", |
| 25 | + model_kwargs: Optional[dict[str, Any]] = None, |
| 26 | + mini_config: Optional[dict[str, Any]] = None, |
| 27 | + **kwargs, |
| 28 | + ): |
| 29 | + kwargs.setdefault("name", "mini-swe-agent") |
| 30 | + kwargs.setdefault("description", "mini-SWE-agent") |
| 31 | + super().__init__(**kwargs) |
| 32 | + self.model = model |
| 33 | + self.api_base = api_base |
| 34 | + self.api_key = api_key |
| 35 | + self.cwd = cwd or self.working_dir or os.getcwd() |
| 36 | + self.step_limit = step_limit |
| 37 | + self.command_timeout = command_timeout |
| 38 | + self.trajectory_path = Path(trajectory_path) |
| 39 | + self.model_kwargs = model_kwargs or {} |
| 40 | + self.mini_config = mini_config or {} |
| 41 | + |
| 42 | + def setup(self) -> None: |
| 43 | + os.environ.setdefault("MSWEA_CONFIGURED", "true") |
| 44 | + os.environ.setdefault("MSWEA_SILENT_STARTUP", "true") |
| 45 | + os.environ.setdefault("MSWEA_GLOBAL_CONFIG_DIR", "/tmp/mswea-config") |
| 46 | + try: |
| 47 | + import minisweagent # noqa: F401 |
| 48 | + except ImportError as exc: |
| 49 | + raise RuntimeError("mini-swe-agent is required. Install with: pip install mini-swe-agent") from exc |
| 50 | + |
| 51 | + async def run_external_async(self, task: str, **kwargs) -> str: |
| 52 | + def run() -> str: |
| 53 | + from minisweagent.agents import get_agent |
| 54 | + from minisweagent.config import get_config_from_spec |
| 55 | + from minisweagent.environments import get_environment |
| 56 | + from minisweagent.models import get_model |
| 57 | + from minisweagent.utils.serialize import recursive_merge |
| 58 | + |
| 59 | + model_kwargs = { |
| 60 | + "drop_params": True, |
| 61 | + "custom_llm_provider": "openai", |
| 62 | + **self.model_kwargs, |
| 63 | + "api_base": self.proxy.url |
| 64 | + if self.proxy |
| 65 | + else self.api_base or os.environ.get("RL_LLM_BASE_URL") or os.environ.get("OPENAI_BASE_URL", ""), |
| 66 | + } |
| 67 | + api_key = ( |
| 68 | + f"sk-proxy-{self.session_id}" |
| 69 | + if self.proxy |
| 70 | + else self.api_key or os.environ.get("RL_LLM_API_KEY") or os.environ.get("OPENAI_API_KEY", "") |
| 71 | + ) |
| 72 | + if api_key: |
| 73 | + model_kwargs["api_key"] = api_key |
| 74 | + |
| 75 | + config = recursive_merge( |
| 76 | + get_config_from_spec("mini"), |
| 77 | + { |
| 78 | + "agent": { |
| 79 | + "agent_class": "default", |
| 80 | + "step_limit": self.step_limit, |
| 81 | + "cost_limit": 0.0, |
| 82 | + "output_path": self.trajectory_path, |
| 83 | + }, |
| 84 | + "environment": { |
| 85 | + "environment_class": "local", |
| 86 | + "cwd": self.cwd, |
| 87 | + "timeout": self.command_timeout, |
| 88 | + }, |
| 89 | + "model": { |
| 90 | + "model_class": "litellm", |
| 91 | + "model_name": self.model, |
| 92 | + "model_kwargs": model_kwargs, |
| 93 | + "cost_tracking": "ignore_errors", |
| 94 | + }, |
| 95 | + }, |
| 96 | + self.mini_config, |
| 97 | + ) |
| 98 | + agent = get_agent( |
| 99 | + get_model(config=config["model"]), |
| 100 | + get_environment(config["environment"], default_type="local"), |
| 101 | + config["agent"], |
| 102 | + default_type="default", |
| 103 | + ) |
| 104 | + result = agent.run(task, **kwargs) |
| 105 | + if result.get("submission"): |
| 106 | + return result["submission"] |
| 107 | + for message in reversed(agent.messages): |
| 108 | + if message.get("role") == "assistant" and message.get("content"): |
| 109 | + return str(message["content"]) |
| 110 | + return "" |
| 111 | + |
| 112 | + return await asyncio.to_thread(run) |
0 commit comments