AdithyaSK's picture
AdithyaSK HF Staff
Deploy standalone SmolDataEnv Native OpenCode
9ce3f34 verified
Raw History Blame Contribute Delete
2.69 kB
"""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")),
)