File size: 3,505 Bytes
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
"""Native OpenCode sessions on the fixed SmolDataEnvs tasks."""

import json
import os

from opencode_env.config import OpenCodeConfig
from opencode_env.harness import OpenCodeSessionFactory
from opencode_env.task import OpenCodeTask
from openenv.core.harness import VerifyResult

from .daytona import DaytonaSandboxBackend
from .tasks import grade_answer, stage_task


class TaskSession:
    def __init__(self, session):
        self.session = session

    def __getattr__(self, name):
        return getattr(self.session, name)

    def fetch_training_trace(self):
        trace = self.session.fetch_training_trace()
        try:
            events = [
                json.loads(line)
                for line in self.session.fetch_trace().splitlines()
                if line.strip()
            ]
            parts = [
                event["part"] for event in events if event.get("type") == "tool_use"
            ]
            ids = [part["callID"] for part in parts]
            if not any(event.get("type") == "step_finish" for event in events):
                raise ValueError("Incomplete OpenCode event stream")
            if any(not value for value in ids) or len(set(ids)) != len(ids):
                raise ValueError("Missing or repeated tool action IDs")
            if any(
                part.get("state", {}).get("status") not in {"completed", "error"}
                for part in parts
            ):
                raise ValueError("Unfinished tool action")
            calls = len(ids)
        except (OSError, ValueError, KeyError, TypeError):
            calls = None
        for turn in trace.turns:
            turn.metadata["native_tool_calls"] = calls
        return trace


def verify(sandbox, task):
    return VerifyResult(
        env_reward=grade_answer(sandbox, task.metadata["folder"]), done=True
    )


class TaskFactory:
    def __init__(self, args, tasks, *, sampling):
        self.tasks = {task["instruction"]: task for task in tasks}
        config = OpenCodeConfig(
            base_url=os.environ["SANDBOX_VLLM_URL"].rstrip("/") + "/v1",
            api_key=os.environ["SANDBOX_VLLM_KEY"],
            model=args.model,
            opencode_version="1.18.31",
            sandbox_home="/root",
            proxy_disable_thinking=True,
            proxy_max_tokens_cap=4096,
            agent_timeout_s=600,
            run_format="json",
            extra_setup_shell="mkdir -p /workdir && rmdir /root/workdir && ln -s /workdir /root/workdir",
            disabled_tools=["webfetch", "question", "task"],
            extra_opencode_json={"permission": {"*": "allow"}},
        )
        self.factory = OpenCodeSessionFactory(
            config=config,
            sampling=sampling,
            verifier=verify,
            mode="transparent_proxy",
            sandbox_backend=DaytonaSandboxBackend(
                image="docker.io/savatar101/env-data-agent-train:base"
            ),
        )

    def create(self, task, seed=None, episode_id=None):
        row = self.tasks[task[-1]["content"]]
        task = OpenCodeTask(
            instruction=row["instruction"], metadata={"folder": row["folder"]}
        )
        session = self.factory.create(
            task, seed=seed, episode_id=episode_id, start_agent=False
        )
        try:
            stage_task(session.sandbox, row["folder"])
            session.start_agent()
            return TaskSession(session)
        except BaseException:
            session.close()
            raise