File size: 2,685 Bytes
81663e8 9ce3f34 81663e8 9ce3f34 81663e8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 | """Run native OpenCode on SmolDataEnvs through the OpenEnv protocol."""
import os
from types import SimpleNamespace
from fastmcp import FastMCP
from openenv.core.env_server import create_app
from openenv.core.env_server.mcp_environment import MCPEnvironment
from openenv.core.env_server.mcp_types import CallToolAction, CallToolObservation
from openenv.core.env_server.types import Observation, State
from .catalog import TaskCatalog, task_by_name
from .environment import TaskFactory
from .ui import build_ui
class OpenCodeEnvironment(TaskCatalog, MCPEnvironment):
SUPPORTS_CONCURRENT_SESSIONS = True
def __init__(self):
self._state = State()
mcp = FastMCP("smoldataenv-opencode")
@mcp.tool
def run_rollout(split: str, task_name: str, model: str, sampling: dict) -> dict:
"""Run OpenCode, grade its answer, and export its typed training trace."""
row = task_by_name(split, task_name)
factory = TaskFactory(
SimpleNamespace(model=model), [row], sampling=sampling
)
session = factory.create([{"role": "user", "content": row["instruction"]}])
try:
session.wait_for_completion(timeout_s=900)
try:
trace = session.fetch_training_trace()
except (ValueError, TypeError, KeyError) as exc:
return {"capture_error": str(exc)}
grade = session.verify([]).env_reward
return {
"training_trace": trace.model_dump(mode="json"),
"correctness": grade,
}
finally:
session.close()
super().__init__(mcp)
def reset(self, seed=None, episode_id=None, **kwargs):
self._state = State(episode_id=episode_id)
return Observation(metadata={"status": "Call run_rollout with a task name"})
@property
def state(self):
return self._state
def _step_impl(self, action, **kwargs):
raise ValueError("Use an MCP tool action")
def step(self, action, timeout_s=None, **kwargs):
return super().step(action, timeout_s=timeout_s or 1800, **kwargs)
async def step_async(self, action, timeout_s=None, **kwargs):
return await super().step_async(action, timeout_s=timeout_s or 1800, **kwargs)
os.environ.setdefault("ENABLE_WEB_INTERFACE", "true")
app = create_app(
OpenCodeEnvironment,
CallToolAction,
CallToolObservation,
gradio_builder=build_ui,
show_default_tab=False,
env_name="smoldataenv_opencode",
max_concurrent_envs=int(os.environ.get("MAX_CONCURRENT_ENVS", "40")),
)
|