diff --git a/.github/workflows/README.md b/.github/workflows/README.md index 4bb0ef059b..40febc96d6 100644 --- a/.github/workflows/README.md +++ b/.github/workflows/README.md @@ -13,7 +13,7 @@ This directory contains automated workflows for the verifiers project. **What it does**: - Runs ruff for linting and formatting checks -- Runs ty type checks with `uv run ty check verifiers packages/tasksets/tasksets packages/harnesses/harnesses` +- Runs ty type checks with `uv run ty check verifiers/v1 packages/tasksets/tasksets packages/harnesses/harnesses` - Runs Semgrep policy checks from the isolated `policy` dependency group. - Uses configuration from `pyproject.toml`, `.pre-commit-config.yaml`, and `.semgrep/verifiers.yml` @@ -55,7 +55,7 @@ To run checks locally the same way they run in CI: ```bash # Ty parity with CI (Python 3.13 target configured in `pyproject.toml`) -uv run ty check verifiers packages/tasksets/tasksets packages/harnesses/harnesses +uv run ty check verifiers/v1 packages/tasksets/tasksets packages/harnesses/harnesses # Verifiers-specific policy lint env PYTHONWARNINGS=ignore::SyntaxWarning uv run --no-dev --group policy semgrep --metrics=off --disable-version-check --config .semgrep/verifiers.yml --error --quiet diff --git a/.github/workflows/style.yml b/.github/workflows/style.yml index 25e5920a0f..913bbcc13b 100644 --- a/.github/workflows/style.yml +++ b/.github/workflows/style.yml @@ -41,7 +41,7 @@ jobs: - name: Install dependencies run: uv sync - name: Run ty - run: uv run ty check verifiers packages/tasksets/tasksets packages/harnesses/harnesses + run: uv run ty check verifiers/v1 packages/tasksets/tasksets packages/harnesses/harnesses semgrep: name: Semgrep runs-on: ubuntu-latest diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 17565f1fbf..713f3d1e1f 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -70,7 +70,7 @@ jobs: run: | test -n "$PRIME_API_KEY" test -n "$PRIME_TEAM_ID" - uv run pytest tests/test_v1_runtime_lifecycle.py -v -m prime_sandbox --cov=verifiers --cov-append --cov-report=xml --cov-report=term + uv run pytest tests/test_v1_core.py -v -m prime_sandbox --cov=verifiers --cov-append --cov-report=xml --cov-report=term test-envs: name: Environments if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 37f1e4265d..5ec80013a9 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -28,7 +28,7 @@ repos: pass_filenames: false - id: ty name: ty (ci parity) - entry: uv run --python 3.13 ty check verifiers packages/tasksets/tasksets packages/harnesses/harnesses + entry: uv run --python 3.13 ty check verifiers/v1 packages/tasksets/tasksets packages/harnesses/harnesses language: system pass_filenames: false stages: [pre-push] diff --git a/.semgrep/verifiers.yml b/.semgrep/verifiers.yml index c20a000616..bca88d7431 100644 --- a/.semgrep/verifiers.yml +++ b/.semgrep/verifiers.yml @@ -1,10 +1,4 @@ rules: - - id: verifiers-no-future-annotations - languages: [python] - severity: ERROR - message: Do not use `from __future__ import annotations`; quote only the specific forward references that need it. - pattern: from __future__ import annotations - - id: verifiers-no-skip-validation languages: [python] severity: ERROR @@ -40,89 +34,16 @@ rules: metavariable: $ANNOT regex: "(Any|Mapping\\[str, object\\]|dict\\[str, object\\]|.*\\|.*\\|.*)" - - id: verifiers-v1-load-environment-config-required + - id: verifiers-v1-no-env-config-subclasses languages: [python] severity: ERROR - message: v1 load_environment must take a strict config object; do not accept None or synthesize config with config-or-Config fallback. + message: EnvConfig is final. Put defaults and public settings on TasksetConfig and HarnessConfig; the library loader owns vf.Env construction. paths: include: - /environments/**/*.py - pattern-either: - - pattern: | - def load_environment(..., config: $CONFIG = None, ...): - ... - - pattern: | - def load_environment(..., config: $CONFIG | None, ...): - ... - - pattern: | - def load_environment(..., config: $CONFIG | None = $DEFAULT, ...): - ... - - pattern: | - def load_environment(..., config: Optional[$CONFIG], ...): - ... - - pattern: | - def load_environment(..., config: Optional[$CONFIG] = $DEFAULT, ...): - ... - - - id: verifiers-v1-no-load-environment-config-fallback - languages: [python] - severity: ERROR - message: v1 load_environment receives a concrete config object from the framework; do not use config or Config() fallback. - paths: - include: - - /environments/**/*.py - patterns: - - pattern-inside: | - def load_environment(...): - ... - - pattern: $CONFIG = $CONFIG or $DEFAULT() - - metavariable-regex: - metavariable: $DEFAULT - regex: ".*Config" - - - id: verifiers-v1-load-environment-canonical-shim - languages: [python] - severity: ERROR - message: v1 load_environment is only the Taskset/Harness loader shim. Keep load_environment as-is; implement the config surface through TasksetConfig, HarnessConfig, load_taskset, and load_harness instead of patching root-loader behavior. - paths: - include: - - /environments/**/*.py - - /verifiers/scripts/init.py - exclude: - - /environments/alphabet_sort/alphabet_sort_v1.py - - /environments/bfcl_v3/bfcl_v3.py - - /environments/dspy_flights/dspy_flights.py - - /environments/dspy_rlm/dspy_rlm.py - - /environments/hello_group_reward_v1/hello_group_reward_v1.py - - /environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1.py - - /environments/hello_self_judge_v1/hello_self_judge_v1.py - - /environments/hello_subagent_v1/hello_subagent_v1.py - - /environments/langchain_deep_agents_wikispeedia/langchain_deep_agents_wikispeedia.py - - /environments/math_python/math_python_v1.py - - /environments/mcp_search_env/mcp_search_env.py - - /environments/nemo_gym_env/nemo_gym_env/env.py - - /environments/nested_harness_v1/nested_harness_v1.py - - /environments/openai_agents_env/openai_agents_env.py - - /environments/rlm_swe_v1/rlm_swe_v1.py - - /environments/tau2_bench_v1/tau2_bench_v1.py - - /environments/wiki_search/wiki_search_v1.py - patterns: - - pattern: | - def load_environment(config: $CONFIG) -> vf.Env: - ... - - pattern-not: | - def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) - - pattern-not: | - def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) + pattern: | + class $CONFIG(vf.EnvConfig): + ... - id: verifiers-v1-no-config-object-defaults languages: [python] @@ -159,23 +80,21 @@ rules: - /environments/**/*.py exclude: - /verifiers/v1/types.py - - /environments/openenv_*/proj/**/*.py + - /environments/**/proj/**/*.py pattern-regex: "\\bAny\\b" - id: verifiers-no-raw-object-containers-v1 languages: [python] severity: ERROR - message: Do not spell broad object containers in v1 or environment code; use a concrete model or a named boundary type that describes the real payload. ConfigData is only for actual config-shaped data; do not use it as a patch to satisfy this rule for protocol, state, task, or endpoint payloads. + message: Do not spell broad object containers in v1 environment or package authoring code; use a concrete model or a named boundary type that describes the real payload. paths: include: - - /verifiers/v1/**/*.py - /environments/**/*.py + - /packages/tasksets/tasksets/**/*.py + - /packages/harnesses/harnesses/**/*.py exclude: - - /verifiers/v1/types.py - - /verifiers/v1/utils/object_utils.py - - /verifiers/v1/utils/task_freeze_utils.py - /environments/openenv_*/proj/**/*.py - pattern-regex: "(?x)(\\b(?:Mapping|MutableMapping|dict|list|Sequence|Iterable|Callable|Awaitable|tuple)\\[[^\\n\\]]*\\bobject\\b|\\bcast\\([^\\n)]*\\bobject\\b)" + pattern-regex: "(?x)(def\\s+\\w+\\([^\\n)]*:\\s*(?:Mapping|MutableMapping|dict|list|Sequence|Iterable|Callable|Awaitable|tuple)\\[[^\\n\\]]*\\bobject\\b|->\\s*(?:Mapping|MutableMapping|dict|list|Sequence|Iterable|Callable|Awaitable|tuple)\\[[^\\n\\]]*\\bobject\\b)" - id: verifiers-no-raw-mapping-annotations-v1 languages: [python] @@ -183,25 +102,12 @@ rules: message: Do not use raw Mapping annotations in v1 or environment code; normalize at the boundary to a dict, a typed Config model, or a named v1 boundary type. Do not patch around unclear contracts with loose Mapping types. paths: include: - - /verifiers/v1/**/*.py - /environments/**/*.py - exclude: - - /verifiers/v1/types.py - - /environments/openenv_*/proj/**/*.py - pattern-regex: "\\b(?:Mapping|MutableMapping)\\[" - - - id: verifiers-v1-package-no-mapping - languages: [python] - severity: ERROR - message: Package taskset/harness implementations should not use Mapping as a loose shape escape hatch; use concrete dict payloads, typed Config models, or package utils that encode a real protocol boundary. - paths: - include: - /packages/tasksets/tasksets/**/*.py - /packages/harnesses/harnesses/**/*.py - - /verifiers/v1/**/*.py exclude: - - /verifiers/v1/types.py - pattern-regex: "\\b(?:Mapping|MutableMapping)\\b" + - /environments/openenv_*/proj/**/*.py + pattern-regex: "(?x)(def\\s+\\w+\\([^\\n)]*:\\s*(?:Mapping|MutableMapping)\\[|->\\s*(?:Mapping|MutableMapping)\\[)" - id: verifiers-get-messages-typed languages: [python] @@ -350,7 +256,7 @@ rules: - id: verifiers-v1-no-env-var-reads-in-signals languages: [python] severity: ERROR - message: v1 reward/update handlers must not read env vars directly; use state.get_client(api="chat") and optionally read state.get_endpoint_config(api="chat").model for model selection. + message: v1 reward/update handlers must not read env vars directly; put model/backend settings on typed configs and resolve clients through the harness runtime. paths: include: - /environments/**/*.py @@ -503,24 +409,6 @@ rules: os.environ[$KEY] ... - - id: verifiers-v1-no-judge-endpoint-config-fields - languages: [python] - severity: ERROR - message: Judge rewards must use state.get_endpoint_config(api="chat"); expose only judge_model, not judge_base_url or judge_api_key_var config fields. - paths: - include: - - /environments/**/*.py - exclude: - - /environments/**/proj/**/*.py - - /environments/**/tasks/**/*.py - patterns: - - pattern-inside: | - class $CONFIG(vf.TasksetConfig): - ... - - pattern-either: - - pattern: "judge_base_url: $TYPE = $VALUE" - - pattern: "judge_api_key_var: $TYPE = $VALUE" - - id: verifiers-v1-no-raw-kwargs-in-subclass-init languages: [python] severity: ERROR @@ -572,27 +460,3 @@ rules: - pattern: | def __init__(...): ... - - - id: verifiers-v1-package-no-static-classmethod-helpers - languages: [python] - severity: ERROR - message: Package taskset/harness implementations should use direct instance methods and standard lifecycle hooks; do not patch around weak ownership with staticmethod/classmethod helpers. - paths: - include: - - /packages/tasksets/tasksets/*.py - - /packages/harnesses/harnesses/*.py - pattern-either: - - pattern: "@staticmethod" - - pattern: "@classmethod" - - - id: verifiers-v1-no-transcript-user-callback-param - languages: [python] - severity: ERROR - message: Do not use transcript as a user callback parameter name in v1 environments; use messages for the rendered prompt/completion messages. - paths: - include: - - /environments/**/*.py - exclude: - - /environments/**/proj/**/*.py - - /environments/**/tasks/**/*.py - pattern-regex: "def\\s+[A-Za-z_][A-Za-z0-9_]*\\s*\\([^)]*\\btranscript\\b" diff --git a/AGENTS.md b/AGENTS.md index 019fce0b65..d35380e925 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -9,7 +9,7 @@ These points are direct restatements of Verifiers docs so agents can follow the - Environments are expected to expose `load_environment(...) -> vf.Environment` and be installable with `prime env install `. (See `docs/overview.md` and `docs/environments.md`.) - Validate environment behavior with `prime eval run ...` before sharing/publishing changes. Treat `prime eval run` as the canonical eval path: it saves results automatically, and agents should not add opt-out flags such as `--skip-upload` unless the user explicitly requests that deviation so runs stay visible in the private Evaluations tab and in `prime eval view`. (See `docs/overview.md` and `docs/development.md`.) - Agents should assume they are allowed to make live model calls through the user's authenticated Prime CLI when a live smoke test is useful. For Prime Inference models, use `prime eval run ` with the base eval configuration from the environment's `pyproject.toml`; do not edit that `pyproject.toml`, and do not add model/config flags unless the task truly requires them. Agents do not need to manage API keys. If sandboxing blocks outbound requests, request elevated permissions for `prime eval run`, preferably as an ongoing approval instead of per run. -- For new taskset/harness environments, start with `prime env init --v1` or `prime env init --v1 --with-harness`. Edit the generated `TasksetConfig` for task settings, `Taskset.load_tasks()` for train/eval task records, `Taskset.load_toolsets()` for task-owned tools, `User` subclasses for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Keep the generated `load_taskset(config: MyTasksetConfig)`, optional `load_harness(config: MyHarnessConfig)`, and `load_environment(config: vf.EnvConfig)` shapes as the component entrypoints. Environment READMEs must use the generated `prime env init` section structure; freeform environment READMEs are not allowed. Treat [BYO Harness](docs/byo-harness.md) as the canonical authoring guide for reusable tasksets, reusable harnesses, framework programs, endpoint interception, and sandboxed Python/command programs. +- For new taskset/harness environments, start with `prime env init --v1` or `prime env init --v1 --with-harness`. Edit the generated `TasksetConfig` for task settings and `ToolsetConfig` entries, `Taskset.load_tasks()` for train/eval task records, `servers/toolset.py` for task-owned tools, `servers/user.py` for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Keep `taskset.py` with `load_taskset(config: MyTasksetConfig)` and optional `harness.py` with `load_harness(config: MyHarnessConfig)` as the component entrypoints. Environment READMEs must use the generated `prime env init` section structure; freeform environment READMEs are not allowed. Treat [BYO Harness](docs/byo-harness.md) as the canonical authoring guide for reusable tasksets, reusable harnesses, framework programs, endpoint interception, and sandboxed Python/command programs. - Use `ToolEnv`/`MCPEnv` for stateless tools and `StatefulToolEnv` when per-rollout state must persist (sandbox/session/db handles). (See `docs/environments.md`.) - If external API keys are required, validate them in `load_environment()` with `vf.ensure_keys(...)` so failures are explicit and early. (See `docs/environments.md`.) @@ -21,7 +21,7 @@ Use these rules when shaping user-facing Verifiers APIs, configs, and environmen - Keep user-facing APIs incredibly minimal and elegant. The best surface is usually golfy but intuitive: one obvious field, one obvious constructor, and no redundant knobs unless there is a concrete long-term reason. - Use Pydantic config models wherever structured configuration is needed. Pydantic is always acceptable and preferred over loose dictionaries when it clarifies the contract. - Prefer strict, narrow types. Use `object`, broad unions, or untyped mappings only at explicit framework boundaries where arbitrary user values are genuinely part of the contract. -- Basic v1 environments should fit in a few dozen self-contained, idiomatic lines: import `verifiers as vf`, define typed taskset/harness config classes when needed, keep policy values in config subclasses, and put implementation logic on the owning taskset or harness class. Static prompts should usually be `system_prompt` config fields; override `load_system_prompt` only for computed prompt loading. Use bindings for shared resources owned by tasksets, toolsets, users, programs, or harnesses; object entries should be loader specs, not pre-initialized resources. +- Basic v1 environments should fit in a few dozen self-contained, idiomatic lines: import `verifiers.v1 as vf`, define typed taskset/harness/toolset config classes when needed, keep policy values in config subclasses, and put implementation logic on the owning taskset, harness, toolset, or user class. Static prompts should usually be `system_prompt` config fields; override `load_system_prompt` only for computed prompt loading. Configure tools and users with `ToolsetConfig.loader` / `UserConfig.loader`; pass serializable config, not pre-initialized resources. - Avoid module globals. Acceptable globals are imports, immutable literals, factory functions, and carefully managed process-level resource constraints such as locks or semaphores. Put all other behavior and state in well-named utility modules, taskset/harness classes, toolsets, users, programs, or user code. - Additional code should have a clear home. Do not hide utilities at the bottom of files or scatter one-off helpers through environment entrypoints. @@ -39,7 +39,7 @@ Use this guidance when contributing to the `verifiers` repository itself. - Before v0.2.0, breaking backward compatibility inside v1 Taskset/Harness APIs is acceptable and encouraged when it improves the core design. Preserve v0 multi-turn environment compatibility unless the user explicitly asks for a v0 migration. - Treat public configuration and docs as part of the API. Keep TOML shapes consistent across eval, GEPA, RL, and Hosted Training; normalize legacy inputs at the ingestion boundary instead of spreading compatibility branches through examples. - For v1 Taskset/Harness work, make the taskset own task data, task controls, task tools, user behavior, metrics, rewards, and task-specific configuration. Make the harness own reusable execution mechanisms such as programs, command agents, primary sandboxes, endpoint interception, framework adapters, and execution artifacts. Use the base `vf.Harness` unless the harness really owns such a mechanism. -- Keep v1 construction explicit: `vf.Env` receives concrete taskset/harness objects, while `load_taskset(config: MyTasksetConfig)` and `load_harness(config: MyHarnessConfig)` define child config types. `EnvConfig` only carries the two child configs; ordinary environment packages use the generated `load_environment(config: vf.EnvConfig)` shim. +- Keep v1 construction explicit: `vf.Env` receives concrete taskset/harness objects, while `taskset.py` / `harness.py` expose `load_taskset(config: MyTasksetConfig)` and `load_harness(config: MyHarnessConfig)` to define child config types. `EnvConfig` only carries the two child configs; the package loader assembles `vf.Env` from those components. - Put class-owned behavior on the taskset or harness class through config fields, `load_*` methods, `User` subclasses, `Toolset`, and `@vf.*` lifecycle methods. `load_taskset` and `load_harness` provide the typed entrypoints to those classes. - Do not override `Taskset.__init__`, `Harness.__init__`, or `User.__init__` in v1 implementations. Put initialization policy in config fields, public load methods, lifecycle handlers, task rows, `Toolset`, `User.get_response`, or utility modules when genuinely shared. - Do not add one-off private helper methods or bottom-of-file helper functions to make taskset/harness classes look shorter. Core lifecycle logic should live on the class with standard public method names or `@vf.*` decorators; reusable multi-line plumbing belongs in a named utility module. diff --git a/README.md b/README.md index d1d9256b39..00149416d7 100644 --- a/README.md +++ b/README.md @@ -138,17 +138,25 @@ def load_environment(dataset_name: str = 'gsm8k') -> vf.Environment: ``` For new environments with reusable tasksets, toolsets, custom programs, or -custom harnesses, use the v1 Taskset/Harness path: +custom harnesses, there is also a separate v1 Taskset/Harness path. v1 is under +active development and may change before release; v0 remains the top-level +`verifiers` surface. ```python -# my_env.py -import verifiers as vf +# my_env/taskset.py +import verifiers.v1 as vf class MyTasksetConfig(vf.TasksetConfig): system_prompt: vf.SystemPrompt = "Reverse text exactly." +class MyTask(vf.Task): + answer: str + + class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: rows = [ { @@ -161,23 +169,16 @@ class MyTaskset(vf.Taskset[MyTasksetConfig]): return [row for row in rows if row["split"] == split] @vf.reward(weight=1.0) - async def contains_answer(self, task, state) -> float: - return float(task["answer"] in str(state.get("completion") or "")) + async def contains_answer(self, task: MyTask, state: vf.State) -> float: + content = state.completion[-1].content if state.completion else "" + return float(task.answer in str(content)) def load_taskset(config: MyTasksetConfig) -> MyTaskset: return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` -The child loader annotation defines the taskset config shape; root -`load_environment` stays typed as `vf.EnvConfig`. See +The child loader annotation defines the taskset config shape. The package +loader assembles `vf.Env` from `taskset.py` and optional `harness.py`. See **[BYO Harness](docs/byo-harness.md)** for the advanced v1 taskset/harness API. Reusable taskset and harness packages live in `tasksets` and `harnesses`. Install them with `uv add "verifiers[packages]"`, or with the narrower diff --git a/V1_TEST_COMMANDS.md b/V1_TEST_COMMANDS.md new file mode 100644 index 0000000000..7265fb5bda --- /dev/null +++ b/V1_TEST_COMMANDS.md @@ -0,0 +1,50 @@ +# V1 Test Commands + +Run these from the repository root. + +## Default V1 Evals + +Each command uses the environment default settings. For the current example +packages, that means `n=5` and `r=3`. `--disable-tui` and +`--abbreviated-summary` only change display. + +Run these in separate terminals when you want broad coverage: + +```bash +prime eval run reverse-text-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run alphabet-sort-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run mcp-search-env-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run math-python-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run hello-group-reward-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run sft-replay-v1 --disable-tui --abbreviated-summary +``` + +Stateful user/tool environments are slower and noisier. Run these separately +from the quick eval batch: + +```bash +prime eval run openenv-echo-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run openenv-textarena-v1 --disable-tui --abbreviated-summary +``` + +```bash +prime eval run tau2-bench-v1 --disable-tui --abbreviated-summary +``` diff --git a/assets/agents/common_best_practices.md b/assets/agents/common_best_practices.md index f0bf3c6e5b..20c21e5ba3 100644 --- a/assets/agents/common_best_practices.md +++ b/assets/agents/common_best_practices.md @@ -5,7 +5,7 @@ These points are direct restatements of Verifiers docs so agents can follow the - Environments are expected to expose `load_environment(...) -> vf.Environment` and be installable with `prime env install `. (See `docs/overview.md` and `docs/environments.md`.) - Validate environment behavior with `prime eval run ...` before sharing/publishing changes. Treat `prime eval run` as the canonical eval path: it saves results automatically, and agents should not add opt-out flags such as `--skip-upload` unless the user explicitly requests that deviation so runs stay visible in the private Evaluations tab and in `prime eval view`. (See `docs/overview.md` and `docs/development.md`.) - Agents should assume they are allowed to make live model calls through the user's authenticated Prime CLI when a live smoke test is useful. For Prime Inference models, use `prime eval run ` with the base eval configuration from the environment's `pyproject.toml`; do not edit that `pyproject.toml`, and do not add model/config flags unless the task truly requires them. Agents do not need to manage API keys. If sandboxing blocks outbound requests, request elevated permissions for `prime eval run`, preferably as an ongoing approval instead of per run. -- For new taskset/harness environments, start with `prime env init --v1` or `prime env init --v1 --with-harness`. Edit the generated `TasksetConfig` for task settings, `Taskset.load_tasks()` for train/eval task records, `Taskset.load_toolsets()` for task-owned tools, `User` subclasses for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Keep the generated `load_taskset(config: MyTasksetConfig)`, optional `load_harness(config: MyHarnessConfig)`, and `load_environment(config: vf.EnvConfig)` shapes as the component entrypoints. Environment READMEs must use the generated `prime env init` section structure; freeform environment READMEs are not allowed. Treat [BYO Harness](docs/byo-harness.md) as the canonical authoring guide for reusable tasksets, reusable harnesses, framework programs, endpoint interception, and sandboxed Python/command programs. +- For new taskset/harness environments, start with `prime env init --v1` or `prime env init --v1 --with-harness`. Edit the generated `TasksetConfig` for task settings and `ToolsetConfig` entries, `Taskset.load_tasks()` for train/eval task records, `servers/toolset.py` for task-owned tools, `servers/user.py` for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Keep `taskset.py` with `load_taskset(config: MyTasksetConfig)` and optional `harness.py` with `load_harness(config: MyHarnessConfig)` as the component entrypoints. Environment READMEs must use the generated `prime env init` section structure; freeform environment READMEs are not allowed. Treat [BYO Harness](docs/byo-harness.md) as the canonical authoring guide for reusable tasksets, reusable harnesses, framework programs, endpoint interception, and sandboxed Python/command programs. - Use `ToolEnv`/`MCPEnv` for stateless tools and `StatefulToolEnv` when per-rollout state must persist (sandbox/session/db handles). (See `docs/environments.md`.) - If external API keys are required, validate them in `load_environment()` with `vf.ensure_keys(...)` so failures are explicit and early. (See `docs/environments.md`.) @@ -17,6 +17,6 @@ Use these rules when shaping user-facing Verifiers APIs, configs, and environmen - Keep user-facing APIs incredibly minimal and elegant. The best surface is usually golfy but intuitive: one obvious field, one obvious constructor, and no redundant knobs unless there is a concrete long-term reason. - Use Pydantic config models wherever structured configuration is needed. Pydantic is always acceptable and preferred over loose dictionaries when it clarifies the contract. - Prefer strict, narrow types. Use `object`, broad unions, or untyped mappings only at explicit framework boundaries where arbitrary user values are genuinely part of the contract. -- Basic v1 environments should fit in a few dozen self-contained, idiomatic lines: import `verifiers as vf`, define typed taskset/harness config classes when needed, keep policy values in config subclasses, and put implementation logic on the owning taskset or harness class. Static prompts should usually be `system_prompt` config fields; override `load_system_prompt` only for computed prompt loading. Use bindings for shared resources owned by tasksets, toolsets, users, programs, or harnesses; object entries should be loader specs, not pre-initialized resources. +- Basic v1 environments should fit in a few dozen self-contained, idiomatic lines: import `verifiers.v1 as vf`, define typed taskset/harness/toolset config classes when needed, keep policy values in config subclasses, and put implementation logic on the owning taskset, harness, toolset, or user class. Static prompts should usually be `system_prompt` config fields; override `load_system_prompt` only for computed prompt loading. Configure tools and users with `ToolsetConfig.loader` / `UserConfig.loader`; pass serializable config, not pre-initialized resources. - Avoid module globals. Acceptable globals are imports, immutable literals, factory functions, and carefully managed process-level resource constraints such as locks or semaphores. Put all other behavior and state in well-named utility modules, taskset/harness classes, toolsets, users, programs, or user code. - Additional code should have a clear home. Do not hide utilities at the bottom of files or scatter one-off helpers through environment entrypoints. diff --git a/assets/agents/repo_development_best_practices.md b/assets/agents/repo_development_best_practices.md index 925182caf3..04fd98a78c 100644 --- a/assets/agents/repo_development_best_practices.md +++ b/assets/agents/repo_development_best_practices.md @@ -12,7 +12,7 @@ Use this guidance when contributing to the `verifiers` repository itself. - Before v0.2.0, breaking backward compatibility inside v1 Taskset/Harness APIs is acceptable and encouraged when it improves the core design. Preserve v0 multi-turn environment compatibility unless the user explicitly asks for a v0 migration. - Treat public configuration and docs as part of the API. Keep TOML shapes consistent across eval, GEPA, RL, and Hosted Training; normalize legacy inputs at the ingestion boundary instead of spreading compatibility branches through examples. - For v1 Taskset/Harness work, make the taskset own task data, task controls, task tools, user behavior, metrics, rewards, and task-specific configuration. Make the harness own reusable execution mechanisms such as programs, command agents, primary sandboxes, endpoint interception, framework adapters, and execution artifacts. Use the base `vf.Harness` unless the harness really owns such a mechanism. -- Keep v1 construction explicit: `vf.Env` receives concrete taskset/harness objects, while `load_taskset(config: MyTasksetConfig)` and `load_harness(config: MyHarnessConfig)` define child config types. `EnvConfig` only carries the two child configs; ordinary environment packages use the generated `load_environment(config: vf.EnvConfig)` shim. +- Keep v1 construction explicit: `vf.Env` receives concrete taskset/harness objects, while `taskset.py` / `harness.py` expose `load_taskset(config: MyTasksetConfig)` and `load_harness(config: MyHarnessConfig)` to define child config types. `EnvConfig` only carries the two child configs; the package loader assembles `vf.Env` from those components. - Put class-owned behavior on the taskset or harness class through config fields, `load_*` methods, `User` subclasses, `Toolset`, and `@vf.*` lifecycle methods. `load_taskset` and `load_harness` provide the typed entrypoints to those classes. - Do not override `Taskset.__init__`, `Harness.__init__`, or `User.__init__` in v1 implementations. Put initialization policy in config fields, public load methods, lifecycle handlers, task rows, `Toolset`, `User.get_response`, or utility modules when genuinely shared. - Do not add one-off private helper methods or bottom-of-file helper functions to make taskset/harness classes look shorter. Core lifecycle logic should live on the class with standard public method names or `@vf.*` decorators; reusable multi-line plumbing belongs in a named utility module. diff --git a/assets/lab/AGENTS.md b/assets/lab/AGENTS.md index e1970a0e1f..bac281f82d 100644 --- a/assets/lab/AGENTS.md +++ b/assets/lab/AGENTS.md @@ -12,7 +12,7 @@ These points are direct restatements of Verifiers docs so agents can follow the - Environments are expected to expose `load_environment(...) -> vf.Environment` and be installable with `prime env install `. (See `docs/overview.md` and `docs/environments.md`.) - Validate environment behavior with `prime eval run ...` before sharing/publishing changes. Treat `prime eval run` as the canonical eval path: it saves results automatically, and agents should not add opt-out flags such as `--skip-upload` unless the user explicitly requests that deviation so runs stay visible in the private Evaluations tab and in `prime eval view`. (See `docs/overview.md` and `docs/development.md`.) - Agents should assume they are allowed to make live model calls through the user's authenticated Prime CLI when a live smoke test is useful. For Prime Inference models, use `prime eval run ` with the base eval configuration from the environment's `pyproject.toml`; do not edit that `pyproject.toml`, and do not add model/config flags unless the task truly requires them. Agents do not need to manage API keys. If sandboxing blocks outbound requests, request elevated permissions for `prime eval run`, preferably as an ongoing approval instead of per run. -- For new taskset/harness environments, start with `prime env init --v1` or `prime env init --v1 --with-harness`. Edit the generated `TasksetConfig` for task settings, `Taskset.load_tasks()` for train/eval task records, `Taskset.load_toolsets()` for task-owned tools, `User` subclasses for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Keep the generated `load_taskset(config: MyTasksetConfig)`, optional `load_harness(config: MyHarnessConfig)`, and `load_environment(config: vf.EnvConfig)` shapes as the component entrypoints. Environment READMEs must use the generated `prime env init` section structure; freeform environment READMEs are not allowed. Treat [BYO Harness](docs/byo-harness.md) as the canonical authoring guide for reusable tasksets, reusable harnesses, framework programs, endpoint interception, and sandboxed Python/command programs. +- For new taskset/harness environments, start with `prime env init --v1` or `prime env init --v1 --with-harness`. Edit the generated `TasksetConfig` for task settings and `ToolsetConfig` entries, `Taskset.load_tasks()` for train/eval task records, `servers/toolset.py` for task-owned tools, `servers/user.py` for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Keep `taskset.py` with `load_taskset(config: MyTasksetConfig)` and optional `harness.py` with `load_harness(config: MyHarnessConfig)` as the component entrypoints. Environment READMEs must use the generated `prime env init` section structure; freeform environment READMEs are not allowed. Treat [BYO Harness](docs/byo-harness.md) as the canonical authoring guide for reusable tasksets, reusable harnesses, framework programs, endpoint interception, and sandboxed Python/command programs. - Use `ToolEnv`/`MCPEnv` for stateless tools and `StatefulToolEnv` when per-rollout state must persist (sandbox/session/db handles). (See `docs/environments.md`.) - If external API keys are required, validate them in `load_environment()` with `vf.ensure_keys(...)` so failures are explicit and early. (See `docs/environments.md`.) @@ -24,7 +24,7 @@ Use these rules when shaping user-facing Verifiers APIs, configs, and environmen - Keep user-facing APIs incredibly minimal and elegant. The best surface is usually golfy but intuitive: one obvious field, one obvious constructor, and no redundant knobs unless there is a concrete long-term reason. - Use Pydantic config models wherever structured configuration is needed. Pydantic is always acceptable and preferred over loose dictionaries when it clarifies the contract. - Prefer strict, narrow types. Use `object`, broad unions, or untyped mappings only at explicit framework boundaries where arbitrary user values are genuinely part of the contract. -- Basic v1 environments should fit in a few dozen self-contained, idiomatic lines: import `verifiers as vf`, define typed taskset/harness config classes when needed, keep policy values in config subclasses, and put implementation logic on the owning taskset or harness class. Static prompts should usually be `system_prompt` config fields; override `load_system_prompt` only for computed prompt loading. Use bindings for shared resources owned by tasksets, toolsets, users, programs, or harnesses; object entries should be loader specs, not pre-initialized resources. +- Basic v1 environments should fit in a few dozen self-contained, idiomatic lines: import `verifiers.v1 as vf`, define typed taskset/harness/toolset config classes when needed, keep policy values in config subclasses, and put implementation logic on the owning taskset, harness, toolset, or user class. Static prompts should usually be `system_prompt` config fields; override `load_system_prompt` only for computed prompt loading. Configure tools and users with `ToolsetConfig.loader` / `UserConfig.loader`; pass serializable config, not pre-initialized resources. - Avoid module globals. Acceptable globals are imports, immutable literals, factory functions, and carefully managed process-level resource constraints such as locks or semaphores. Put all other behavior and state in well-named utility modules, taskset/harness classes, toolsets, users, programs, or user code. - Additional code should have a clear home. Do not hide utilities at the bottom of files or scatter one-off helpers through environment entrypoints. diff --git a/assets/lab/environments/AGENTS.md b/assets/lab/environments/AGENTS.md index 3ff3017a60..e934d064cc 100644 --- a/assets/lab/environments/AGENTS.md +++ b/assets/lab/environments/AGENTS.md @@ -294,13 +294,11 @@ async def my_reward_func(completion, my_helper) -> float: return await my_helper.score(completion) ``` -For taskset/harness environments, keep shared dependencies behind the taskset or -harness that owns them. Bindings are the canonical way to inject shared -resources into rewards, updates, tools, and programs. Configured binding -objects should use serializable loader paths when they cross a TOML or CLI -boundary; Python-only construction may use factory callables directly when a -resource cannot be serialized. Required Taskset and Toolset factory parameters -must be supplied through bindings. +For taskset/harness environments, keep shared dependencies behind the taskset, +harness, toolset, user, or runtime that owns them. v1 toolsets and users receive +shared rollout data through `@vf.tool(args=...)` and write serializable rollout +data through `sets` and `extends`. Configured objects use serializable loader +paths across TOML and CLI boundaries. Judges are used for tasks where deterministic evaluation is impractical, and an LLM is used to score responses. **JudgeRubric** stores an LLM client inside the @@ -703,23 +701,23 @@ environments/my_env/ ### v1 Env Shape -The v1 template teaches the standard object layout: one taskset class, one -typed `load_taskset(config: MyTasksetConfig)` child factory, and a tiny -`load_environment(config: vf.EnvConfig)` root loader that delegates through -`vf.load_taskset(config=config.taskset)` and -`vf.load_harness(config=config.harness)`. The child factory annotation defines -the taskset config type for TOML, CLI, eval, GEPA, RL, and Hosted Training. +v1 is component-first, under active development, and intentionally separate +from v0. v1 environment code imports `verifiers.v1 as vf`. Packages expose +`taskset.py` and, only when they own reusable execution behavior, `harness.py`. +They do not define a root `load_environment`; the library loader assembles +`vf.Env` from discovered components. Factory annotations define config types +for TOML, CLI, eval, GEPA, RL, and Hosted Training. After `prime env init my-env --v1`, edit the generated taskset class: 1. Add task settings to `TasksetConfig`. 2. Return task records from `load_tasks(split=...)`. -3. Return task-owned tools from `load_toolsets` when needed. +3. Add task-owned toolsets or a user when needed. 4. Add lifecycle, metric, reward, and advantage methods with `@vf.*`. Add a harness config, harness class, and `load_harness(config: MyHarnessConfig)` when the environment owns reusable rollout behavior. -Otherwise the generated root loader uses the base harness. +Otherwise the component loader uses the base harness. `EnvConfig` is the lightweight envelope for the two child configs. Put environment knobs on `TasksetConfig` or `HarnessConfig`. @@ -727,14 +725,20 @@ environment knobs on `TasksetConfig` or `HarnessConfig`. The taskset-only shape is: ```python -import verifiers as vf +import verifiers.v1 as vf class MyTasksetConfig(vf.TasksetConfig): system_prompt: vf.SystemPrompt = "Answer exactly." +class MyTask(vf.Task): + answer: str + + class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: """Return serializable task records as a list, generator, or Dataset.""" if split == "eval": @@ -748,24 +752,15 @@ class MyTaskset(vf.Taskset[MyTasksetConfig]): ] @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - if not messages: + async def correct_answer(self, task: MyTask, state: vf.State) -> float: + if not state.completion: return 0.0 - response = str(messages[-1].content or "").strip() - return float(response == task["answer"]) + response = str(state.completion[-1].content or "").strip() + return float(response == task.answer) def load_taskset(config: MyTasksetConfig) -> MyTaskset: return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` With a reusable harness, keep the same explicit object boundary: @@ -785,43 +780,17 @@ def load_taskset(config: MyTasksetConfig) -> MyTaskset: def load_harness(config: MyHarnessConfig) -> MyHarness: return MyHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -Keep v1 dependencies behind the owning taskset or harness. Do not pass -already-instantiated resource objects through environment loaders. Bindings are -allowed wherever the owning taskset, toolset, user, program, or harness wires -callables. `objects` entries should be loader specs: prefer serializable import -paths in config, and use factory callables directly only for Python-only -construction when the dependency cannot be serialized. Required Taskset and -Toolset factory parameters must be supplied through bindings. - -Judge-style rewards should read endpoint details from the rollout state: - -```python -@vf.reward(weight=1.0) -async def judge_reward(task, state) -> float: - endpoint = state.get_endpoint_config(api="chat") - client = state.get_client(api="chat") - model = str(task.get("judge_model") or endpoint.model) - ... ``` -Expose at most `judge_model: str | None = None` on the taskset config. Do not -add judge endpoint URL/API-key fields or read `os.environ` inside reward/update -handlers. +Keep v1 dependencies behind the owning taskset, harness, toolset, user, or runtime +provider. Config, task rows, state, tool specs, and user specs must stay +serializable. Live clients, runtimes, functions, and file handles do not +cross those boundaries. Custom mutable rollout data belongs in `state.extras`. For reusable tasksets and harnesses, [BYO Harness](byo-harness.md) is the canonical v1 implementation guide. It covers ownership, configs, task controls, -system prompts, users, toolsets, programs, sandboxes, artifacts, nested -harnesses, package adapters, and TOML/CLI overrides. +system prompts, users, toolsets, runtimes, artifacts, nested harnesses, package +adapters, and TOML/CLI overrides. ### pyproject.toml diff --git a/configs/endpoints.toml b/configs/endpoints.toml index 10f9f5d479..48a73f1e61 100644 --- a/configs/endpoints.toml +++ b/configs/endpoints.toml @@ -1,184 +1,44 @@ [[endpoint]] -endpoint_id = "olmo3-32b-t" -model = "allenai/olmo-3-32b-think" +endpoint_id = "haiku-4.5" +model = "anthropic/claude-haiku-4.5" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "olmo3-7b-i" -model = "allenai/olmo-3-7b-instruct" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "olmo3-7b-t" -model = "allenai/olmo-3-7b-think" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "trinity-mini" -model = "arcee/trinity-mini" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "haiku" -model = "claude-haiku-4-5" -url = "https://api.anthropic.com" -key = "ANTHROPIC_API_KEY" -type = "anthropic_messages" - -[[endpoint]] -endpoint_id = "sonnet" -model = "claude-sonnet-4-5" -url = "https://api.anthropic.com" -key = "ANTHROPIC_API_KEY" -type = "anthropic_messages" - -[[endpoint]] -endpoint_id = "opus" -model = "claude-opus-4-5" -url = "https://api.anthropic.com" -key = "ANTHROPIC_API_KEY" -type = "anthropic_messages" - -[[endpoint]] -endpoint_id = "deepseek-chat" -model = "deepseek-chat" -url = "https://api.deepseek.com/v1" -key = "DEEPSEEK_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "deepseek-reasoner" -model = "deepseek-reasoner" -url = "https://api.deepseek.com/v1" -key = "DEEPSEEK_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "deepseek-chat-anth" -model = "deepseek-chat" -url = "https://api.deepseek.com/anthropic" -key = "DEEPSEEK_API_KEY" -type = "anthropic_messages" - -[[endpoint]] -endpoint_id = "deepseek-reasoner-anth" -model = "deepseek-reasoner" -url = "https://api.deepseek.com/anthropic" -key = "DEEPSEEK_API_KEY" type = "anthropic_messages" [[endpoint]] -endpoint_id = "gemini-2.5-flash" -model = "google/gemini-2.5-flash" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "gemini-2.5-pro" -model = "google/gemini-2.5-pro" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "gemini-3-flash" -model = "google/gemini-3-flash" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "gemini-3-pro" -model = "google/gemini-3-pro-preview" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "gemini-3-pro-exp" -model = "google/gemini-3-pro-preview" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "qwen3-30b-i" -model = "qwen/qwen3-30b-a3b-instruct-2507" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "qwen3-30b-t" -model = "qwen/qwen3-30b-a3b-thinking-2507" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "qwen3-235b-i" -model = "qwen/qwen3-235b-a22b-instruct-2507" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "qwen3-235b-t" -model = "qwen/qwen3-235b-a22b-thinking-2507" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "qwen3-vl-30b-i" -model = "qwen/qwen3-vl-30b-a3b-instruct" -url = "https://api.pinference.ai/api/v1" -key = "PRIME_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "qwen3-vl-30b-t" -model = "qwen/qwen3-vl-30b-a3b-thinking" +endpoint_id = "sonnet-4.5" +model = "anthropic/claude-sonnet-4.5" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" -type = "openai_chat_completions" +type = "anthropic_messages" [[endpoint]] -endpoint_id = "qwen3-vl-235b-i" -model = "qwen/qwen3-vl-235b-a22b-instruct" +endpoint_id = "sonnet-4.6" +model = "anthropic/claude-sonnet-4.6" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" -type = "openai_chat_completions" +type = "anthropic_messages" [[endpoint]] -endpoint_id = "qwen3-vl-235b-t" -model = "qwen/qwen3-vl-235b-a22b-thinking" +endpoint_id = "opus-4.5" +model = "anthropic/claude-opus-4.5" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" -type = "openai_chat_completions" +type = "anthropic_messages" [[endpoint]] -endpoint_id = "kimi-k2" -model = "moonshotai/kimi-k2-0905" +endpoint_id = "opus-4.6" +model = "anthropic/claude-opus-4.6" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" -type = "openai_chat_completions" +type = "anthropic_messages" [[endpoint]] -endpoint_id = "kimi-k2-t" -model = "moonshotai/kimi-k2-thinking" +endpoint_id = "opus-4.7" +model = "anthropic/claude-opus-4.7" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" -type = "openai_chat_completions" +type = "anthropic_messages" [[endpoint]] endpoint_id = "gpt-oss-120b" @@ -224,84 +84,70 @@ type = "openai_chat_completions" [[endpoint]] endpoint_id = "gpt-4.1-nano" -model = "gpt-4.1-nano" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" +model = "openai/gpt-4.1-nano" +url = "https://api.pinference.ai/api/v1" +key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] endpoint_id = "gpt-4.1-mini" -model = "gpt-4.1-mini" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" +model = "openai/gpt-4.1-mini" +url = "https://api.pinference.ai/api/v1" +key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] endpoint_id = "gpt-4.1" -model = "gpt-4.1" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "gpt-5-nano" -model = "gpt-5-nano" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" -type = "openai_chat_completions" - -[[endpoint]] -endpoint_id = "gpt-5-mini" -model = "gpt-5-mini" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" +model = "openai/gpt-4.1" +url = "https://api.pinference.ai/api/v1" +key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "gpt-5" -model = "gpt-5" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" +endpoint_id = "gpt-5.2" +model = "openai/gpt-5.2" +url = "https://api.pinference.ai/api/v1" +key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "gpt-5.1" -model = "gpt-5.1" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" +endpoint_id = "glm-4.7" +model = "z-ai/glm-4.7" +url = "https://api.pinference.ai/api/v1" +key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "gpt-5.2" -model = "gpt-5.2" -url = "https://api.openai.com/v1" -key = "OPENAI_API_KEY" +endpoint_id = "glm-5.1" +model = "z-ai/glm-5.1" +url = "https://api.pinference.ai/api/v1" +key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "glm-4.5" -model = "z-ai/glm-4.5" +endpoint_id = "qwen3-vl-30b-i" +model = "qwen/qwen3-vl-30b-a3b-instruct" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "glm-4.5-air" -model = "z-ai/glm-4.5-air" +endpoint_id = "qwen3-vl-30b-t" +model = "qwen/qwen3-vl-30b-a3b-thinking" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "glm-4.6" -model = "z-ai/glm-4.6" +endpoint_id = "qwen3-vl-235b-i" +model = "qwen/qwen3-vl-235b-a22b-instruct" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" type = "openai_chat_completions" [[endpoint]] -endpoint_id = "glm-4.7" -model = "z-ai/glm-4.7" +endpoint_id = "qwen3-vl-235b-t" +model = "qwen/qwen3-vl-235b-a22b-thinking" url = "https://api.pinference.ai/api/v1" key = "PRIME_API_KEY" type = "openai_chat_completions" diff --git a/docs/assets/v1-composition-lifecycle.svg b/docs/assets/v1-composition-lifecycle.svg index 1cbb7c4849..f71959fd5b 100644 --- a/docs/assets/v1-composition-lifecycle.svg +++ b/docs/assets/v1-composition-lifecycle.svg @@ -52,7 +52,7 @@ updates - sync trajectory, enrich state + sync transcript, enrich state score + cleanup diff --git a/docs/assets/v1-task-harness-state.svg b/docs/assets/v1-task-harness-state.svg index 47a9c6cb71..c0ca0e3f44 100644 --- a/docs/assets/v1-task-harness-state.svg +++ b/docs/assets/v1-task-harness-state.svg @@ -37,7 +37,7 @@ model/client endpoint toolsets and MCP sandbox leases - trajectory capture + transcript capture Base loop, Python entrypoint, or CLI program @@ -46,7 +46,7 @@ State Mutable rollout record - trajectory + transcript completion metrics / reward artifacts / errors diff --git a/docs/byo-harness.md b/docs/byo-harness.md index 1f40a86403..ed75e4cf9a 100644 --- a/docs/byo-harness.md +++ b/docs/byo-harness.md @@ -1,638 +1,342 @@ # v1 Taskset/Harness Environments -Use the v1 Taskset/Harness path for reusable environments: dataset adapters, -tool environments, user simulators, sandboxed programs, command agents, -framework harnesses, packaged benchmark formats, and environments that need the -same config shape from Python, TOML, eval, GEPA, RL, and Hosted Training. +v1 is the active-development Taskset/Harness API and may change before release. +Use it for new reusable tasksets, reusable harnesses, MCP tools, user +simulators, framework adapters, endpoint interception, command agents, and +runtime-backed environments. -For short v0-style `SingleTurnEnv`, `ToolEnv`, or `MultiTurnEnv` examples, see -[Environments](environments.md). For API signatures, see -[Reference](reference.md). +v1 is intentionally separate from v0. Author v1 files with: -## Start From The Template +```python +import verifiers.v1 as vf +from verifiers.utils.response_utils import parse_response_message +``` -Initialize the v1 Taskset/Harness template first: +The top-level `verifiers` package remains the v0 authoring surface. The public +loader can still load v1 packages, but v1 environment code should import from +`verifiers.v1` so the boundary is explicit. -```bash -prime env init my-env --v1 -``` +## Package Shape -Use a custom harness template only when the environment owns reusable execution -behavior such as a command agent, framework adapter, sandboxed program, browser -loop, endpoint interceptor, or nested harness: +Start from the v1 template: ```bash +prime env init my-env --v1 prime env init my-env --v1 --with-harness ``` -Open the generated `my_env.py` and edit it in this order: - -1. Add user-facing task settings to `MyTasksetConfig`. -2. Fill `MyTaskset.load_tasks(split=...)` with train/eval task records. -3. Add task-owned tools with `MyTaskset.load_toolsets()` when the task defines - an action space. -4. Add task behavior with `@vf.setup`, `@vf.update`, `@vf.reward`, `@vf.metric`, - `@vf.cleanup`, and related lifecycle methods on `MyTaskset`. -5. Add a `User` subclass and `load_user()` when the task owns simulated user - behavior. -6. If `--with-harness` is used, put execution-level program, sandbox, endpoint, - model, or harness lifecycle behavior on `MyHarness`. -7. Keep the generated loaders as the typed entrypoints: `load_taskset`, - optional `load_harness`, and the root `load_environment`. +A v1 package is component-first: + +```text +my_env/ + pyproject.toml + README.md + my_env/ + __init__.py + taskset.py + harness.py # only when the package owns a custom harness + servers/ + __init__.py + toolset.py # optional toolset implementation + user.py # optional user implementation +``` -## Golden Shape +Do not define a package-level `load_environment`. The loader imports the package +and discovers `taskset.py` and `harness.py`. -Every v1 environment has one root loader and typed child loaders: +## Minimal Taskset ```python -import verifiers as vf +import verifiers.v1 as vf -class MyTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPrompt = "Answer exactly." +class ReverseTask(vf.Task): + answer: str -class MyTaskset(vf.Taskset[MyTasksetConfig]): +class ReverseTasksetConfig(vf.TasksetConfig): + system_prompt: vf.SystemPrompt = "Return only the reversed string." + + +class ReverseTaskset(vf.Taskset[ReverseTasksetConfig]): + task_type = ReverseTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - """Return serializable task records as a list, generator, or Dataset.""" if split == "eval": return [] return [ { + "row_id": 0, "prompt": [{"role": "user", "content": "Reverse abc."}], "answer": "cba", "max_turns": 1, } ] - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - return float(response.strip() == task["answer"]) + @vf.reward + async def exact(self, task: ReverseTask, state: vf.State) -> float: + return float(state.completion[-1].content == task.answer) -def load_taskset(config: MyTasksetConfig) -> MyTaskset: - return MyTaskset(config=config) +def load_taskset(config: ReverseTasksetConfig) -> ReverseTaskset: + return ReverseTaskset(config=config) +``` +`load_taskset(config: ReverseTasksetConfig)` is the taskset entrypoint and defines +the `[taskset]` config schema for eval, RL, GEPA, and hosted workers. -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` +## Custom Harness -Add a custom harness only when the environment owns reusable execution behavior: +Use the base `vf.Harness` until the package owns a reusable execution protocol. +Add `harness.py` when the environment owns a command/program runner, framework +adapter, browser loop, nested harness, interception protocol, runtime placement, +or execution-level lifecycle. ```python -async def run_agent(task: vf.Task, state: vf.State) -> vf.State: - client = state.get_client(api="chat") - response = await client.chat.completions.create( - model=state.get_model(), - messages=[*state.get("system_prompt", []), *task["prompt"]], - ) - message = response.choices[0].message - state["completion"] = [{"role": "assistant", "content": message.content or ""}] - return state +import verifiers.v1 as vf +from verifiers.utils.response_utils import parse_response_message class MyHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="my_env:run_agent") + max_turns: int = 3 class MyHarness(vf.Harness[MyHarnessConfig]): - pass + async def run_with_context(self, context: vf.Context) -> None: + task = context.task + state = context.state + toolsets = context.toolsets + prompt = self.initial_messages(task) + response = await context.model_client.get_response( + prompt=prompt, + model=context.model, + sampling_args=context.sampling_args, + tools=toolsets.tools() if toolsets is not None else None, + state=state, + ) + turn = vf.Turn( + prompt=prompt, + completion=await parse_response_message(response), + tool_calls=list(response.message.tool_calls or []), + response_id=response.id, + model=response.model, + created=response.created, + finish_reason=response.message.finish_reason, + usage=vf.TurnUsage.from_usage(response.usage), + tokens=vf.TurnTokens.from_response( + response.message.tokens, + is_truncated=bool(response.message.is_truncated), + ), + is_truncated=bool(response.message.is_truncated), + ) + state.transcript.append(turn) + state.is_truncated = state.is_truncated or turn.is_truncated + state.stop("assistant_completed") def load_harness(config: MyHarnessConfig) -> MyHarness: return MyHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` -The child loader annotations are load-bearing. `load_taskset(config: -MyTasksetConfig)` defines the `[env.taskset]` schema; `load_harness(config: -MyHarnessConfig)` defines the `[env.harness]` schema. Keep -`load_environment(config: vf.EnvConfig)` as-is: implement the config surface through taskset and harness configs, not root loader kwargs. +`load_harness(config: MyHarnessConfig)` defines the `[harness]` config schema. +Packages without `harness.py` use the base harness. -Start with a taskset and the base harness. Add a custom harness only when the -environment owns a reusable execution protocol such as a command agent, -third-party framework adapter, browser loop, endpoint interceptor, primary -sandbox placement, or program runner. +## State -## Implementation Map +`vf.State` is a strict Pydantic model. It is not a dict and has no convenience +helpers for arbitrary mutation. -- Import as `import verifiers as vf`. -- Use `XXXConfig` classes for structured settings. -- Put task behavior on the taskset config/class. -- Put execution behavior on the harness config/class. -- Use `vf.Env` and `vf.EnvConfig` for ordinary environment packages. -- Let the base `Taskset`, `Harness`, and `User` constructors handle - construction; customize with config fields and public methods. -- Put taskset/harness behavior on the owning class with standard public methods - or `@vf.*` decorators. -- Use `system_prompt` for system messages. -- Keep reusable multi-line internals in utility modules with clear names. - -Utility modules are appropriate only for reused, nontrivial internals or messy -upstream adapters that users should not think about. - -## Ownership - -| Object | Owns | -| --- | --- | -| `Taskset` | Task loading, task data, task prompts, task controls, task-owned tools, user behavior, task-specific lifecycle, metrics, rewards, advantages, and task-owned program/sandbox inputs. | -| `Harness` | Rollout execution, execution-level system prompts, model/client defaults, programs, command agents, framework adapters, endpoint interception, primary sandbox placement, harness-owned tools, and execution artifacts. | -| `Env` | The adapter that makes one taskset/harness pair usable by eval and training workers. | +Use: -If a tool or state transition defines the task action space, observations, or -success condition, it belongs to the taskset. If a class describes how a model -or external agent attempts arbitrary tasks, it belongs to the harness. +- `state.transcript: list[vf.Turn]` for model request/completion turns. +- `state.prompt`, `state.completion`, and `state.messages` for derived views of + the current transcript. +- `state.extras` for user-owned per-rollout mutable JSON. +- `state.metrics`, `state.reward`, and `Turn.tokens.*_advantages` for scoring. +- `state.artifacts` for serializable outputs worth saving. -Examples: - -- Wikispeedia link tools belong to the Wikispeedia taskset. -- TextArena game state and user responses belong to the TextArena taskset. -- Harbor task directories, uploads, and tests belong to `HarborTaskset`. -- OpenCode, Pi, mini-swe-agent, Terminus, and RLM execution belong to harness - classes. -- Endpoint routing and interception belong to the harness/runtime, not task - rows. - -## Config - -Config values must be serializable. Use import-ref strings such as -`"my_env.module:factory"` when config needs to name a callable across TOML, CLI, -or package boundaries. Python constructors may pass runtime objects only where -the constructor explicitly accepts them, such as `vf.Toolset(tools=[...])` or -standalone `vf.Harness(model=..., client=...)`. - -Common owner config fields: - -| Field | Meaning | -| --- | --- | -| `system_prompt` | String, system-message list, or `vf.SystemPromptConfig`. | -| `user` | `UserConfig` subclass that materializes a registered `User`. | -| `toolsets` | Configured toolset collection. | -| `objects` | Private dependency factories owned by this object. | -| `bindings` | Hidden argument bindings for handlers, tools, users, and programs. | -| `artifacts` | Text/JSON artifacts owned by this object. | -| lifecycle lists | Import-ref `setups`, `updates`, `metrics`, `rewards`, `cleanups`, etc. | -| `scoring` | Per-handler tuning or skipping by handler name. | - -Put taskset fields on `TasksetConfig`; put harness fields on `HarnessConfig`. -Avoid broad unions and untyped mappings unless arbitrary JSON is the actual -task payload. +Do not store live clients, functions, runtimes, file handles, or other +non-serializable objects on `State`, `Task`, tool specs, user specs, or config. ## Tasks -Tasksets load train and eval data through `load_tasks(split=...)`: +Task rows become typed `vf.Task` instances. Define a task subclass when a field +matters to the environment: ```python -class GSM8KTasksetConfig(vf.TasksetConfig): - dataset_name: str = "gsm8k" - num_examples: int | None = None - - -class GSM8KTaskset(vf.Taskset[GSM8KTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - """Return serializable task records as a list, generator, or Dataset.""" - dataset_split = "test" if split == "eval" else "train" - dataset = load_dataset(self.config.dataset_name, "main", split=dataset_split) - if self.config.num_examples is not None: - dataset = dataset.select(range(self.config.num_examples)) - return dataset +class SearchTask(vf.Task): + query: str + answer: str ``` -`vf.Tasks` may be a `datasets.Dataset`, an iterable of serializable records, or -an iterable of `vf.Task` objects. During rollout, records become immutable -`vf.Task` objects. - -Return a `datasets.Dataset` directly when the source already has standard -columns such as `question` and `answer`; the framework derives `prompt` from -`question`. Hardcode fixed upstream split names inside `load_tasks(split=...)`. -Only expose split-name config when the upstream split choice is -genuine user-space configuration, not the way v1 decides whether eval exists. -Return `[]` for `split == "eval"` when the taskset has no explicit eval source; -`vf.Env` treats the empty split as an absent eval dataset, so -`Environment.get_eval_dataset()` falls back to train data with the standard -warning. - -Common task fields: +Task records must be serializable. Common fields: | Field | Meaning | | --- | --- | -| `prompt` | User/developer/tool messages. No system messages. | -| `system_prompt` | Per-task taskset-side system prompt override. | +| `row_id` | Explicit upstream row/example identifier. | +| `prompt` | Initial non-system prompt messages. May be empty when a user server starts the rollout. | +| `system_prompt` | Per-task taskset system prompt override. | | `answer` | Reference answer or target data. | -| `info` | Serializable metadata. | +| `image` | Per-task runtime image selection. | +| `resources` | Per-task runtime CPU, memory, GPU, and disk requests. Runtime config wins when explicitly set. | | `max_turns` | Per-task base-loop turn limit. | -| `toolsets` | Toolset visibility: `{"show": [...]}` or `{"hide": [...]}`. | -| `tools` | Per-toolset tool visibility: `{"search": {"show": [...]}}`. | -| `sandbox` | Per-task sandbox override. | -| `program` | Task-owned files, dirs, setup, env, artifacts, bindings, and args. | -| `artifacts` | Task-owned artifacts collected after program execution. | - -Users should not manage task/example IDs. Preserve upstream IDs only as ordinary -metadata when they matter. - -Do not copy config defaults into every row. Use `max_turns`, `sandbox`, -`program`, and tool visibility fields in task records only when they genuinely -vary by example. - -## System Prompts - -System prompt resolution happens per task during rollout setup. - -There are two sides: +| `toolsets` | Per-task toolset visibility. | +| `tools` | Per-tool visibility. | +| `extras` fields | Put custom mutable rollout data in `state.extras`, not task rows. | -- `T`: the resolved taskset side. `task["system_prompt"]` wins for that task; - otherwise the taskset uses `TasksetConfig.system_prompt`. -- `H`: the harness side from `HarnessConfig.system_prompt`. +## Tools -`HarnessConfig.system_prompt_strategy` decides how those two sides resolve: - -| Strategy | Meaning | -| --- | --- | -| `HT` | Harness side followed by resolved taskset side. Default. | -| `TH` | Resolved taskset side followed by harness side. | -| `H_OR_T` | Harness side when present, otherwise resolved taskset side. | -| `T_OR_H` | Resolved taskset side when present, otherwise harness side. | -| `H` | Harness side only. | -| `T` | Resolved taskset side only. | -| `REJECT` | Error if both sides are present. | - -Static prompts belong in config: +Tasksets declare toolsets through `ToolsetConfig`. A toolset implementation is +a `vf.Toolset` subclass with `@vf.tool` methods. ```python -class WordleTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPrompt = ( - "Play Wordle. Submit guesses inside ... tags." - ) -``` +class SearchToolsetConfig(vf.ToolsetConfig): + scope: vf.Scope = "rollout" -For GEPA or other file-backed prompt optimization, use config too: -```python -class WordleTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPromptConfig = vf.SystemPromptConfig( - path="system_prompt.txt" +class SearchToolset(vf.Toolset): + @vf.tool( + args={"case": "state.metadata.case"}, + extends={"events": "state.extras.search_events"}, + sets={"last_result": "state.extras.last_search_result"}, ) -``` - -Override `load_system_prompt(config)` only when prompt loading is computed from -other config fields or package resources. + def search(self, query: str, case: str) -> dict: + ... -## Toolsets -Toolsets package model-visible schemas, hidden bindings, private objects, -artifacts, lifecycle hooks, and optional runtime scope. - -```python class SearchTasksetConfig(vf.TasksetConfig): - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"index": "my_env.search:load_index"} - ) - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - {"search.query.index": "objects.index"} - ) - - -async def query(index, q: str) -> str: - return index.search(q) - - -class SearchTaskset(vf.Taskset[SearchTasksetConfig]): - def load_toolsets(self, config: SearchTasksetConfig) -> vf.Toolsets: - return {"search": vf.Toolset(tools=[query])} + toolsets: vf.ToolsetConfigs = {"search": SearchToolsetConfig()} ``` -Bindings inject hidden arguments that the model does not see. Common sources -include `task.*`, `state.*`, `objects.*`, and `tools.*`. +The mapping key is the tool prefix exposed to the model. A config file can +override taskset-defined keys directly, and can add a new key by setting +`source` to a `ToolsetConfig` class path. -Tasks show all toolsets and tools by default. Restrict visibility in task data: +Use `scope="rollout"` for per-rollout servers and `scope="env"` for servers +that live for one `EnvRun`. Eval creates one `EnvRun` per evaluation, so +env-scope servers are shared across all rollouts in that evaluation. +Group-shared resources should use `state.group_id` plus env-scope server state. -```python -yield { - "prompt": [{"role": "user", "content": "Use the calculator only."}], - "toolsets": {"show": ["math"]}, - "tools": {"math": {"show": ["calculate"]}}, -} -``` - -Use rollout-scoped toolsets for resources that exist only during a rollout, -such as OpenReward sessions or sandbox-backed servers. Keep live backend -handles on `state`; keep task rows serializable. Dynamic schemas use -`state.add_tool("toolset_name", vf.Tool(...))` during rollout setup against a -named rollout toolset. - -MCP servers are normal tool entries: - -```python -class FetchTasksetConfig(vf.TasksetConfig): - toolsets: dict[str, vf.ToolsetConfig] = { - "fetch": vf.ToolsetConfig( - tools=[vf.MCPToolConfig(command="uvx", args=["mcp-server-fetch"])], - scope="rollout", - ) - } -``` +`@vf.tool(args=..., sets=..., extends=...)` is the data-flow contract. `args` +are removed from the model-visible tool schema and injected from serialized +task/state paths at call time. `sets` consume return keys and replace state +paths. `extends` consume return keys and append lists to list paths. Multiple +same-path extends in one tool-call batch are allowed; their relative order is +not a contract. -Custom harness programs should consume resolved tools from state: - -```python -async def run_agent(task: vf.Task, state: vf.State) -> vf.State: - tools = state.get_tools() - result = await framework_agent(task["prompt"], tools=list(tools.values())) - state["completion"] = [{"role": "assistant", "content": result}] - return state -``` +Unbound result keys remain visible to the model. A `content` key is used as the +tool message content when it is the only unbound key. ## Users -A `User` simulates environment/user responses between model turns. It is not a -callable; subclass `vf.User` and implement `get_response`. +Users follow the same pattern. A user implementation exposes a hidden +`respond` tool. The harness calls it after each assistant turn, and also before +the first model request when `task.prompt` is empty. ```python -class GameUserConfig(vf.UserConfig): +class DialogueUserConfig(vf.UserConfig): pass -class GameUser(vf.User[GameUserConfig]): - async def get_response( - self, - task: vf.Task, - state: vf.State, - messages: list[vf.Message], - ) -> list[vf.UserMessage]: - observation = state["game"].observe(messages) - return [{"role": "user", "content": observation}] +class DialogueUser(vf.User): + @vf.user( + args={"transcript": "state.transcript"}, + sets={"turn_count": "state.extras.turn_count"}, + ) + def respond(self, transcript: list[dict]) -> dict: + return { + "messages": [{"role": "user", "content": "continue"}], + "turn_count": len(transcript), + } -class GameTasksetConfig(vf.TasksetConfig): - user: GameUserConfig = GameUserConfig() +class DialogueTasksetConfig(vf.TasksetConfig): + user: vf.UserConfig | None = DialogueUserConfig() ``` -Use a user when the environment naturally replies after model turns. Use tools -when the model chooses an explicit schema action. Use setup/update handlers when -state should change without adding conversation messages. +## Runtime -## Programs, Harnesses, And Sandboxes +The harness owns live runtimes. Runtime config is resolved from taskset, +harness, and environment config into one provider/runtime pair for each rollout. +State stores only serializable runtime metadata and artifacts. -`HarnessConfig.program` is a `vf.ProgramConfig`: +Available runtime providers: -| Form | Meaning | -| --- | --- | -| `vf.ProgramConfig()` | Base endpoint-backed tool loop. | -| `vf.ProgramConfig(base=True)` | Explicit base loop, usually with sandbox options. | -| `vf.ProgramConfig(fn="my_env:run")` | Importable Python program. | -| `vf.ProgramConfig(command=["agent", "run"])` | Local or sandboxed command. | +- `vf.SubprocessRuntimeConfig` +- `vf.DockerRuntimeConfig` +- `vf.PrimeRuntimeConfig` -The preferred Python program signature is: +Reserved runtime provider configs: -```python -async def program(task: vf.Task, state: vf.State) -> vf.State: - state["answer"] = task["answer"] - return state -``` +- `vf.ModalRuntimeConfig` +- `vf.DaytonaRuntimeConfig` -Programs may call models, call tools, run solvers, replay cached solutions, or -adapt third-party frameworks. They should read immutable task data, mutate -serializable state, and let lifecycle handlers collect artifacts and score. +Live runtimes expose: -Tasksets can contribute task-local program data through `task["program"]`. -Harnesses still own the program kind, channel wiring, and primary sandbox -placement. Duplicate files, env vars, artifacts, or bindings fail fast. +- `start` +- `stop` +- `expose` +- `run` +- `read` +- `write` -Put sandbox config on the harness when it is part of the execution mechanism: +Custom runtimes implement `vf.RuntimeProvider` and `vf.Runtime`. -```python -class PythonHarnessConfig(vf.HarnessConfig): - sandbox: vf.SandboxConfig = vf.SandboxConfig( - image="python:3.11-slim", - scope="rollout", - ) - program: vf.ProgramConfig = vf.ProgramConfig( - fn="my_env.solver:solve", - sandbox=True, - ) -``` - -Put sandbox overrides on tasks only when the taskset owns per-task images, -files, resource sizing, or setup. +## Scoring -## Lifecycle And Scoring - -Lifecycle decorators attach behavior to the owning class: +Tasksets define task lifecycle, metrics, and rewards. Harnesses may define +lifecycle and metrics for reusable execution telemetry: ```python -class QAATaskset(vf.Taskset[QAATasksetConfig]): - @vf.update - async def extract_answer(self, task: vf.Task, state: vf.State) -> None: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - state["answer"] = str(messages[-1].content or "") if messages else "" +class MyTaskset(vf.Taskset[MyTasksetConfig]): + @vf.setup + async def setup(self, task: MyTask, state: vf.State) -> None: + state.extras["started"] = True @vf.metric - async def response_length(self, task: vf.Task, state: vf.State) -> float: - return float(len(str(state.get("answer") or ""))) - - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State) -> float: - return float(state.get("answer") == task["answer"]) -``` - -Rollout handlers can request `task`, `state`, `completion`, `prompt`, and bound -hidden args. Group handlers use `tasks` and `states` and must return one value -per state when scoring. - -## Objects, Bindings, And Artifacts + async def length(self, state: vf.State) -> float: + return float(len(str(state.completion[-1].content))) -`objects` are private dependency factories. `bindings` connect those objects, -task fields, state fields, or runtime values to hidden callable arguments. - -```python -class ExtractTasksetConfig(vf.TasksetConfig): - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"extractor": "my_env.extractors:load_answer_extractor"} - ) - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - {"exact.extractor": "objects.extractor"} - ) - - -class ExtractTaskset(vf.Taskset[ExtractTasksetConfig]): - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State, extractor) -> float: - return float(extractor(state.get("completion") or []) == task["answer"]) + @vf.reward + async def exact(self, task: MyTask, state: vf.State) -> float: + return float(state.completion[-1].content == task.answer) ``` -Artifacts are text/JSON files copied into serialized state: +Group rewards run through `env.score_group(tasks, states)`. Env-level advantage +functions mutate token-level advantages in place. The default v1 advantage is +`"rl"`; pass `advantage=None` only when the caller wants trainer-owned +advantages. ```python -class AgentHarnessConfig(vf.HarnessConfig): - artifacts: vf.ArtifactsConfig = vf.ArtifactsConfig.model_validate( - {"agent_log": {"path": "/app/agent.log", "format": "text", "optional": True}} - ) +env = vf.Env(taskset=MyTaskset(), advantage="grpo") ``` -Artifacts can live on tasksets, harnesses, users, toolsets, programs, or tasks. -The owner determines which sandbox/filesystem is searched first. - -## Nested Harnesses - -Nested harnesses are ordinary harness runs. Create a child task, create a child -state, and run the child harness. - -```python -async def ask_child(name: str, state: vf.State) -> str: - harness = vf.Harness( - config=vf.HarnessConfig(program=vf.ProgramConfig(fn="my_env.children:greet")) - ) - task = vf.Task( - {"prompt": [{"role": "user", "content": f"Say hello to {name}."}]} - ).freeze() - child_state = await harness.run(task, state.for_task(task)) - messages = vf.get_messages(child_state.get("completion") or [], role="assistant") - return str(messages[-1].content or "") if messages else "" -``` - -Borrow runtime handles only when the child intentionally reuses live parent -resources: - -```python -child_state = state.for_task(child_task, borrow="model", tools=["search"]) -``` +## Evaluation -Borrowed resources are process-local and stripped before state serialization. - -## Packaged Tasksets And Harnesses - -Reusable implementations live in standalone packages under `packages/`: +Run v1 environments the same way as v0 packages: ```bash -uv add "verifiers[packages]" -uv add "verifiers[tasksets]" -uv add "verifiers[harnesses]" -uv add "verifiers[openenv]" -uv add "verifiers[openreward]" -uv add "verifiers[ta]" -uv add "verifiers[nemogym]" +prime eval run reverse-text-v1 -n 5 -r 1 ``` -Package-backed environments use the same loader shape: - -```python -import verifiers as vf -from harnesses import OpenCode, OpenCodeConfig -from tasksets import HarborTaskset, HarborTasksetConfig - - -def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: - return HarborTaskset(config=config) - - -def load_harness(config: OpenCodeConfig) -> OpenCode: - return OpenCode(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -Tasksets include Harbor, OpenEnv, OpenReward, ReplayTaskset, TextArena, and -NeMoGym. Harnesses include OpenCode, Pi, mini-swe-agent, Terminus, RLM, -ReplayHarness, and NeMoGymHarness. - -## TOML And CLI - -Eval and training config owns the run: model, endpoint, sampling, examples, and -rollout count. v1 child config owns environment behavior: - -```toml -model = "openai/gpt-5.4-mini" -num_examples = 5 -rollouts_per_example = 3 - -[[eval]] -env_id = "my-v1-env" - -[eval.sampling] -max_tokens = 4096 - -[eval.taskset] -system_prompt = "Answer exactly." - -[eval.harness] -max_turns = 4 -``` - -CLI overrides target typed child fields: +Typed overrides address component config directly: ```bash -prime eval run my-v1-env --taskset.system-prompt "Answer exactly." --harness.max-turns 4 -``` - -For package-only composition, TOML can name the taskset and harness packages -directly: - -```toml -[[eval]] - -[eval.taskset] -id = "tasksets.harbor" -tasks_dir = "tasks" - -[eval.harness] -id = "harnesses.opencode" -max_turns = 8 -``` - -Callable config uses import refs: - -```toml -[[env.taskset.rewards]] -fn = "my_env.rewards:exact" -weight = 1.0 -priority = 0 -``` - -Use `[...scoring.function_name]` to tune or skip an existing class-defined -metric/reward without creating a new signal: - -```toml -[env.taskset.scoring.exact] -weight = 0.5 +prime eval run my-env-v1 --taskset.system-prompt "Be terse." --harness.max-turns 4 ``` -## Checklist - -Before publishing or asking for review: - -1. `load_environment(config: vf.EnvConfig)` is the only root loader shape. -2. Custom tasksets have `load_taskset(config: MyTasksetConfig)`. -3. Custom harnesses have `load_harness(config: MyHarnessConfig)`. -4. No `Taskset`, `Harness`, or `User` subclass overrides `__init__`. -5. No ordinary environment subclass of `vf.Env` or `vf.EnvConfig` exists. -6. Config fields are serializable and named `XXXConfig`. -7. Taskset-owned behavior is not hidden in the harness. -8. Harness-owned execution is not hidden in task rows. -9. Static prompts live in config; computed prompts use `load_system_prompt`. -10. Tools are exposed through `vf.Toolset`; task rows only show/hide them. -11. Runtime-only resources live on state or runtime-managed owners. -12. Metrics/rewards/setup/update/cleanup are decorated with `@vf.*`. -13. Generated component loaders remain the typed taskset/harness entrypoints. -14. One-off helper methods and bottom-of-file helper functions are absent. -15. Install/load/eval has been validated with `prime eval run` or the relevant - package-install test. +The package name selects the default taskset/harness. `--taskset.id` and +`--harness.id` may point at other installed v1 component packages when an eval +needs to compose them. + +## Design Rules + +- Import `verifiers.v1 as vf` in v1 code. +- Prefer one package-level taskset and, only when necessary, one harness. +- Keep tools and users in `servers/`. +- Keep config serializable and typed. +- Keep live Python functions only as methods on loaded owner objects. +- Mutate framework state only through typed `State` fields. +- Put user-owned mutable rollout data in `state.extras`. +- Do not pass runtimes, clients, functions, or other live objects across + task/state/tool/user/config boundaries. diff --git a/docs/development.md b/docs/development.md index 9cf6b32d0d..eaedb1184e 100644 --- a/docs/development.md +++ b/docs/development.md @@ -181,7 +181,7 @@ def test_with_mock(mock_client): ### Code Style - Strict `ruff` enforcement via pre-commit hooks -- `ty` runs in the pre-push hook via `uv run --python 3.13 ty check verifiers` +- `ty` runs in the pre-push hook via `uv run --python 3.13 ty check verifiers/v1 packages/tasksets/tasksets packages/harnesses/harnesses` - Use type hints for function parameters and returns - Write docstrings for public functions/classes - Keep functions focused and modular @@ -314,7 +314,7 @@ uv run pytest tests/test_envs.py -k math_python # Specific environment # Linting uv run ruff check --fix . # Fix lint errors uv run ruff format --check verifiers tests # Verify Python formatting -uv run ty check verifiers # Type check (matches CI Ty target) +uv run ty check verifiers/v1 packages/tasksets/tasksets packages/harnesses/harnesses # Type check (matches CI Ty target) # Environment tools prime env init new-env # Create v0 environment stub diff --git a/docs/environments.md b/docs/environments.md index fc97465274..676a3b4c31 100644 --- a/docs/environments.md +++ b/docs/environments.md @@ -287,13 +287,11 @@ async def my_reward_func(completion, my_helper) -> float: return await my_helper.score(completion) ``` -For taskset/harness environments, keep shared dependencies behind the taskset or -harness that owns them. Bindings are the canonical way to inject shared -resources into rewards, updates, tools, and programs. Configured binding -objects should use serializable loader paths when they cross a TOML or CLI -boundary; Python-only construction may use factory callables directly when a -resource cannot be serialized. Required Taskset and Toolset factory parameters -must be supplied through bindings. +For taskset/harness environments, keep shared dependencies behind the taskset, +harness, toolset, user, or runtime that owns them. v1 toolsets and users receive +shared rollout data through `@vf.tool(args=...)` and write serializable rollout +data through `sets` and `extends`. Configured objects use serializable loader +paths across TOML and CLI boundaries. Judges are used for tasks where deterministic evaluation is impractical, and an LLM is used to score responses. **JudgeRubric** stores an LLM client inside the @@ -696,23 +694,23 @@ environments/my_env/ ### v1 Env Shape -The v1 template teaches the standard object layout: one taskset class, one -typed `load_taskset(config: MyTasksetConfig)` child factory, and a tiny -`load_environment(config: vf.EnvConfig)` root loader that delegates through -`vf.load_taskset(config=config.taskset)` and -`vf.load_harness(config=config.harness)`. The child factory annotation defines -the taskset config type for TOML, CLI, eval, GEPA, RL, and Hosted Training. +v1 is component-first, under active development, and intentionally separate +from v0. v1 environment code imports `verifiers.v1 as vf`. Packages expose +`taskset.py` and, only when they own reusable execution behavior, `harness.py`. +They do not define a root `load_environment`; the library loader assembles +`vf.Env` from discovered components. Factory annotations define config types +for TOML, CLI, eval, GEPA, RL, and Hosted Training. After `prime env init my-env --v1`, edit the generated taskset class: 1. Add task settings to `TasksetConfig`. 2. Return task records from `load_tasks(split=...)`. -3. Return task-owned tools from `load_toolsets` when needed. +3. Add task-owned toolsets or a user when needed. 4. Add lifecycle, metric, reward, and advantage methods with `@vf.*`. Add a harness config, harness class, and `load_harness(config: MyHarnessConfig)` when the environment owns reusable rollout behavior. -Otherwise the generated root loader uses the base harness. +Otherwise the component loader uses the base harness. `EnvConfig` is the lightweight envelope for the two child configs. Put environment knobs on `TasksetConfig` or `HarnessConfig`. @@ -720,14 +718,20 @@ environment knobs on `TasksetConfig` or `HarnessConfig`. The taskset-only shape is: ```python -import verifiers as vf +import verifiers.v1 as vf class MyTasksetConfig(vf.TasksetConfig): system_prompt: vf.SystemPrompt = "Answer exactly." +class MyTask(vf.Task): + answer: str + + class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: """Return serializable task records as a list, generator, or Dataset.""" if split == "eval": @@ -741,24 +745,15 @@ class MyTaskset(vf.Taskset[MyTasksetConfig]): ] @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - if not messages: + async def correct_answer(self, task: MyTask, state: vf.State) -> float: + if not state.completion: return 0.0 - response = str(messages[-1].content or "").strip() - return float(response == task["answer"]) + response = str(state.completion[-1].content or "").strip() + return float(response == task.answer) def load_taskset(config: MyTasksetConfig) -> MyTaskset: return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` With a reusable harness, keep the same explicit object boundary: @@ -778,43 +773,17 @@ def load_taskset(config: MyTasksetConfig) -> MyTaskset: def load_harness(config: MyHarnessConfig) -> MyHarness: return MyHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -Keep v1 dependencies behind the owning taskset or harness. Do not pass -already-instantiated resource objects through environment loaders. Bindings are -allowed wherever the owning taskset, toolset, user, program, or harness wires -callables. `objects` entries should be loader specs: prefer serializable import -paths in config, and use factory callables directly only for Python-only -construction when the dependency cannot be serialized. Required Taskset and -Toolset factory parameters must be supplied through bindings. - -Judge-style rewards should read endpoint details from the rollout state: - -```python -@vf.reward(weight=1.0) -async def judge_reward(task, state) -> float: - endpoint = state.get_endpoint_config(api="chat") - client = state.get_client(api="chat") - model = str(task.get("judge_model") or endpoint.model) - ... ``` -Expose at most `judge_model: str | None = None` on the taskset config. Do not -add judge endpoint URL/API-key fields or read `os.environ` inside reward/update -handlers. +Keep v1 dependencies behind the owning taskset, harness, toolset, user, or runtime +provider. Config, task rows, state, tool specs, and user specs must stay +serializable. Live clients, runtimes, functions, and file handles do not +cross those boundaries. Custom mutable rollout data belongs in `state.extras`. For reusable tasksets and harnesses, [BYO Harness](byo-harness.md) is the canonical v1 implementation guide. It covers ownership, configs, task controls, -system prompts, users, toolsets, programs, sandboxes, artifacts, nested -harnesses, package adapters, and TOML/CLI overrides. +system prompts, users, toolsets, runtimes, artifacts, nested harnesses, package +adapters, and TOML/CLI overrides. ### pyproject.toml diff --git a/docs/evaluation.md b/docs/evaluation.md index bb11ebbb72..546a1a2080 100644 --- a/docs/evaluation.md +++ b/docs/evaluation.md @@ -29,7 +29,7 @@ Run evaluations directly against a local or Hub environment: prime eval run my-env -m openai/gpt-4.1-mini -n 10 ``` -`prime eval` resolves and installs the environment when needed, imports the environment module using Python's import system, calls its `load_environment()` function, runs 5 examples with 3 rollouts each (the default), scores them using the environment's rubric, and prints aggregate metrics. +`prime eval` resolves and installs the environment when needed, imports the environment module using Python's import system, loads the environment, runs 5 examples with 3 rollouts each (the default), scores them, and prints aggregate metrics. v0 packages load through `load_environment(...)`; v1 packages load through discovered `taskset.py` and `harness.py` components. ## Hosted Evaluations @@ -68,9 +68,9 @@ The positional argument accepts two formats: Environment IDs are converted to Python module names (`my-env` → `my_env`) and imported after `prime eval run` resolves the environment package. -For v1 `load_environment(config: vf.EnvConfig)` loaders, prefer typed -taskset/harness overrides. These flags are parsed against the concrete child -config types from `load_taskset(config: ...)` and `load_harness(config: ...)`: +For v1 packages, prefer typed taskset/harness overrides. These flags are parsed +against the concrete child config types from `load_taskset(config: ...)` and +`load_harness(config: ...)`: ```bash prime eval run my-v1-env --taskset.id my-taskset --harness.id my-harness --harness.max-turns 4 @@ -81,8 +81,8 @@ harness flags are fields on the typed child configs. If the loaded package does not provide a local child loader, `--taskset.id` and `--harness.id` select the taskset and harness loader packages. -For legacy or direct-constructor environments, the `--env-args` flag passes -arguments to your `load_environment()` function: +For v0 or direct-constructor environments, the `--env-args` flag passes +arguments to the package `load_environment()` function: ```bash prime eval run my-env -a '{"difficulty": "hard", "num_examples": 100}' @@ -136,7 +136,7 @@ env.set_concurrency(256) | `--api-client-type` | — | `openai_chat_completions` | Client type: `openai_completions`, `openai_chat_completions`, `openai_chat_completions_token`, `openai_responses`, `renderer`, `anthropic_messages`, or `nemorl_chat_completions` | | `--endpoints-path` | `-e` | `./configs/endpoints.toml` | Path to TOML endpoints registry | | `--header` | — | — | Extra HTTP header (`Name: Value`), repeatable | -| `--header-from-state` | — | framework session id | Per-request header whose value is read from rollout state (`Name: state_key`), repeatable | +| `--header-from-state` | — | framework session id | Per-request header whose value is read from the serialized client state (`Name: state_key`), repeatable | The `renderer` client type requires the optional renderer package. Install it with `uv add "verifiers[renderers]"` before running evals with `--api-client-type renderer`. @@ -178,7 +178,7 @@ headers = { "X-Custom-Header" = "value" } In `[[eval]]` TOML configs you can set extra headers as `headers = { ... }` and/or as a list `header = ["Name: Value", ...]` (same form as repeated `--header`). Merge order is: registry row, then the `headers` table, then each `header` / `--header` line, with later entries overriding the same name. -For per-request headers that need to vary per rollout, use `headers_from_state = { "X-Name" = "state_key" }` and/or `header_from_state = ["X-Name: state_key", ...]` (same form as repeated `--header-from-state`). The value for each request is resolved at send time as `state[state_key]`. If unset, Verifiers supplies a framework-managed `X-Session-ID`. +For per-request headers that need to vary per rollout, use `headers_from_state = { "X-Name" = "state_key" }` and/or `header_from_state = ["X-Name: state_key", ...]` (same form as repeated `--header-from-state`). The value for each request is resolved at send time from the serialized client state. If unset, Verifiers supplies a framework-managed `X-Session-ID`. To define equivalent replicas, add multiple `[[endpoint]]` entries with the same `endpoint_id`. diff --git a/docs/overview.md b/docs/overview.md index 4e9a3d7d4a..cefff1bab1 100644 --- a/docs/overview.md +++ b/docs/overview.md @@ -57,6 +57,9 @@ For the v1 Taskset/Harness authoring path: prime env init my-env --v1 prime env init my-env --v1 --with-harness ``` +v1 is under active development and may change before release. The top-level +`verifiers` package remains the v0 authoring surface; v1 environment code +imports `verifiers.v1 as vf`. This will create a new module called `my_env` with a runnable environment template. For v1 templates, start by editing the generated `TasksetConfig`, @@ -87,17 +90,23 @@ def load_environment(dataset_name: str = 'gsm8k') -> vf.Environment: ``` For new environments with reusable tasksets, toolsets, custom programs, or -custom harnesses, use the v1 Taskset/Harness path: +custom harnesses, use the separate v1 Taskset/Harness path: ```python -# my_env.py -import verifiers as vf +# my_env/taskset.py +import verifiers.v1 as vf class MyTasksetConfig(vf.TasksetConfig): system_prompt: vf.SystemPrompt = "Answer exactly." +class MyTask(vf.Task): + answer: str + + class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: """Return serializable task records as a list, generator, or Dataset.""" if split == "eval": @@ -111,30 +120,21 @@ class MyTaskset(vf.Taskset[MyTasksetConfig]): ] @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - if not messages: + async def correct_answer(self, task: MyTask, state: vf.State) -> float: + if not state.completion: return 0.0 - response = str(messages[-1].content or "").strip() - return float(response == task["answer"]) + response = str(state.completion[-1].content or "").strip() + return float(response == task.answer) def load_taskset(config: MyTasksetConfig) -> MyTaskset: return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` See [BYO Harness](byo-harness.md) for the advanced v1 taskset/harness API. -The child loader annotation is the config contract: keep root -`load_environment` typed as `vf.EnvConfig`, put task settings on -`TasksetConfig`, and add `load_harness(config: MyHarnessConfig)` only when the -environment owns a reusable harness. +The child loader annotation is the config contract: put task settings on +`TasksetConfig`, and add `harness.py` with +`load_harness(config: MyHarnessConfig)` only when the environment owns a +reusable harness. The package loader assembles `vf.Env` from those components. Reusable v1 taskset and harness packages live in `tasksets` and `harnesses`. Install them with `uv add "verifiers[packages]"`, or with the narrower `verifiers[tasksets]`, `verifiers[harnesses]`, and backend-specific extras such diff --git a/docs/reference.md b/docs/reference.md index a50811f4aa..fcd6401082 100644 --- a/docs/reference.md +++ b/docs/reference.md @@ -123,6 +123,9 @@ Selects which `Client` implementation to use. Set via `ClientConfig.client_type` ## Data Types +This section documents the top-level v0 data types. v1 data types live under +`verifiers.v1` and are documented in [v1 Taskset/Harness Classes](#v1-tasksetharness-classes). + ### State ```python @@ -365,6 +368,9 @@ class RolloutScores(TypedDict): ### Environment Classes +Top-level environment classes are v0. v1 code should import `verifiers.v1 as vf` +and use the taskset/harness classes below. + #### Environment ```python @@ -605,202 +611,138 @@ state, or output field. ### v1 Taskset/Harness Classes -The v1 API is exposed from the top-level `verifiers` namespace and documented -in [BYO Harness](byo-harness.md). Its core unit is: +v1 is under active development and is exposed from `verifiers.v1`, not the +top-level `verifiers` namespace: ```python -state = await harness.run(task, state=None) +import verifiers.v1 as vf ``` -`Taskset` and `Env` package that runner for datasets, evals, and training. - -#### Task +The core direct runner is: ```python -class Task(dict): - def freeze(self) -> Task: ... -``` - -Immutable, JSON-serializable input data. A task is usually created by a -`Taskset`, but can be run directly through a standalone `Harness`. - -Common top-level fields: - -| Field | Description | -|-------|-------------| -| `prompt` | User/developer/tool messages for the rollout. Must not contain system messages. | -| `system_prompt` | Per-task system messages or string. | -| `answer` | Reference answer or target data. Stays on task, not state. | -| `info` | Serializable metadata. | -| `max_turns` | Per-task base-loop turn limit. | -| `tools` | Toolset-keyed tool visibility: `{"wiki": {"show": [...]}}` or `{"wiki": {"hide": [...]}}`. | -| `toolsets` | Toolset visibility: `{"show": [...]}` or `{"hide": [...]}`. | -| `sandbox` | Per-task sandbox overrides for sandboxed programs. | -| `artifacts` | Per-task text/JSON files collected after program execution. | -| `program` | Task-owned files, dirs, env, setup, artifacts, bindings, and command args. | - -`task.runtime` is not public schema. Runtime metadata belongs on `State`. - -#### State - -```python -class State(dict): - @classmethod - def for_task(task: Task, ...) -> State: ... - def stop(self, condition: str = "state_done") -> None: ... - def get_model(self) -> str: ... - def get_client(api: str = "chat_completions", *, sync: bool = False) -> object: ... - def get_endpoint_config(api: str = "chat_completions") -> EndpointConfig: ... - def get_tools() -> dict[str, Callable[..., Awaitable[object]]]: ... - def get_max_turns(default: int) -> int: ... - def finalize() -> State: ... +state = await harness.run( + task="hello world", + model="openai/gpt-5", + score=False, +) ``` -Mutable rollout output. State starts from a task and accumulates trajectory, -completion, metrics, reward, timing, artifacts, errors, and user-defined -serializable fields. +`Harness.run(...)` is generation-only by default. `Env.run_rollout(...)` calls +it with `score=True`; nested judge/self-check calls should pass a parent +`context` and keep `score=False`. -Framework-managed fields such as `is_completed`, `stop_condition`, -`is_truncated`, and `error` cannot be written directly. Use `state.stop(...)` or -raise `vf.Error` subclasses. - -`State.for_task(...)` can borrow selected active runtime handles from another -state: +#### Task ```python -child_state = state.for_task(child_task, borrow=["model", "sandbox"], tools="bash") -child_state = await child_harness.run(child_task, child_state) +class Task(BaseModel, extra="forbid", frozen=True): + task_id: str + row_id: int + prompt: Messages + name: str | None + description: str | None + image: str | None + max_turns: int | None ``` -Borrowed handles are process-local and stripped before state crosses the -serialization boundary. +Tasks are immutable, serializable Pydantic specs. Subclass `vf.Task` for +benchmark-specific fields; do not use task as an arbitrary bag. -#### Taskset +#### State ```python -class Taskset: - def __init__(config: TasksetConfig | None = None): ... - - def to_task(task: Task | JsonData) -> Task: ... - def load_tasks(split: TaskSplit = "train") -> Tasks: ... - async def init_group(task: Task, num_rollouts: int) -> tuple[list[Task], list[State]]: ... - def get_dataset() -> Dataset: ... - def get_eval_dataset() -> Dataset: ... +class State(BaseModel, extra="forbid"): + task_id: str | None + transcript: list[Turn] + extras: JsonData + metadata: JsonData + metrics: dict[str, float] + reward: float + artifacts: JsonData ``` -Packages tasks and task-owned behavior. Tasksets define -`load_tasks(split="train" | "eval")`. During rollout, -records are always materialized as `vf.Task`. -`Taskset.__init__` is final; subclasses customize behavior through config, -task-loading methods, lifecycle handlers, `load_toolsets`, `load_user`, and -other public load methods. +`State` is the rollout record. `state.transcript` is canonical; there is no live +`trajectory` alias. User-owned mutable rollout data belongs in `state.extras`. +Advantages are token-level fields on `TurnTokens`, not scalar state fields. +Live model clients, runtimes, MCP sessions, functions, and file handles never go +on `State`. #### Harness ```python class Harness: - def __init__( - config: HarnessConfig | None = None, + async def run( + task: Task | str, + state: State | None = None, *, - model: str | ModelConfig | None = None, - client: Client | ClientConfig | str | None = None, - sampling_args: SamplingArgs | None = None, - ): ... + model: ModelConfig | str | None = None, + teacher: ModelConfig | str | None = None, + context: Context | None = None, + score: bool = False, + ) -> State: ... - async def run(task: Task, state: State | None = None) -> State: ... - async def score_group(tasks: list[Task], states: list[State]) -> list[State]: ... - async def cleanup_group(tasks: list[Task], states: list[State]) -> None: ... - async def teardown() -> None: ... + async def run_with_context(context: Context) -> None: ... ``` -Runs one task. All model calls go through the v1 interception endpoint so -trajectory capture, sampling args, tool forwarding, and protocol translation use -one path across local Python, sandboxed Python, command programs, and the base -tool loop. -`Harness.__init__` is final; subclasses customize behavior through config, -`load_sandbox`, `load_toolsets`, `load_system_prompt`, lifecycle handlers, and -program config. +Subclass `run_with_context(...)` when the harness owns an execution mechanism +such as a command agent, framework adapter, endpoint interception loop, or nested +harness. `Context` is the live rollout object: task, state, model client, +teacher client, runtime, MCP tools/user registries, parent context, and scoring +flags. -`HarnessConfig.program` is a `ProgramConfig`. Dict/TOML inputs are accepted as -shorthand for the same config object: +`EnvRun` is the environment execution object. It starts env-scope toolsets/users +once, creates one runtime per rollout, and owns grouped rollout coordination. +`Group` owns the tasks/states for one grouped example and calls `score_group` +after its member rollouts finish. -| Form | Meaning | -|------|---------| -| `ProgramConfig()` | Default endpoint-backed tool loop. | -| `ProgramConfig(base=True, ...)` | Explicit default loop, usually with sandbox options. | -| `ProgramConfig(fn="pkg.module:run", ...)` | Importable Python program. | -| `ProgramConfig(command=["cmd", "arg"], ...)` | Local or sandboxed command. | +#### Runtime -Reusable command harnesses should subclass `ProgramConfig` and implement -`resolve()` with `self.resolve_command(command=..., ...)` so typed harness -settings resolve to a canonical command program without a second loader layer. +`RuntimeConfig` is the serializable backend spec. `Runtime` is the live handle +with `start`, `stop`, `expose`, `run`, `read`, and `write`. Built-in configs are +`SubprocessRuntimeConfig`, `DockerRuntimeConfig`, and `PrimeRuntimeConfig`. +`ModalRuntimeConfig` and `DaytonaRuntimeConfig` are reserved stubs. -Sandboxed `program.fn` refs resolve their owning local package from the resolved -module root: single-file modules use `pyproject.toml` in the same directory as -the module file, and package modules use `pyproject.toml` inside the package -directory. v1 uploads and installs that package in the program sandbox. Package -dependencies come from normal `[project.dependencies]`. - -#### Env +#### Toolset And User ```python -class Env(vf.Environment): - def __init__(config, *, taskset=Taskset, harness=Harness): ... -``` +class SearchToolsetConfig(vf.ToolsetConfig): + scope: vf.Scope = "rollout" -Adapter that makes a v1 taskset/harness pair usable by eval and training -workers. `config` is an `EnvConfig` object or mapping with nested `taskset` and -`harness` sections. -#### Toolset And MCPTool +class SearchToolset(vf.Toolset): + @vf.tool( + args={"case": "state.metadata.case"}, + extends={"events": "state.extras.search_events"}, + ) + def search(self, query: str, case: str) -> dict: + ... -```python -class Toolset: - def __init__( - tools=None, - show=None, - hide=None, - bindings=None, - objects=None, - write: bool | None = None, - scope: Literal["rollout", "group", "global"] | None = None, - sandbox: SandboxConfig | Literal["program"] | None = None, - stops=None, - setups=None, - updates=None, - cleanups=None, - teardowns=None, - config: ToolsetConfig | None = None, - ): ... -class MCPTool: - def __init__(command: str, args=None, env=None, cwd: str | None = None): ... +class MyTasksetConfig(vf.TasksetConfig): + toolsets: vf.ToolsetConfigs = {"search": SearchToolsetConfig()} ``` -Toolsets package callable tools, MCP servers, private dependency factories, -hidden bindings, and tool-owned lifecycle handlers. `objects.*` bindings are -private to the owning toolset/user and are not directly accessible from state. -String binding sources are framework paths; literal strings should be bound via -callable sources. -Tasks show all toolsets/tools by default and can restrict them with `show` or -`hide` visibility at `task["toolsets"]` and `task["tools"]`. +The `toolsets` mapping key is the tool prefix. Config can add a new key with +`source = "my_env.servers.search.config:SearchToolsetConfig"`, which points to +the config class; the tool implementation is derived from that class. +`UserConfig` is a sibling config for `User` servers over the same server base. +`@vf.tool` hides bound `args` from model-visible schemas and binds selected +return keys back into state through `sets` and `extends`. Binding inputs may +read `task.*`, `state.*`, `extras.*`, and server-local `resources.*`; binding +outputs may write `state.*` or `extras.*`. Toolsets and users return messages +through the shared `vf.ServerResponse` shape. -#### v1 Config Models +#### Env ```python -TasksetConfig(...) -HarnessConfig(...) -ToolsetConfig(...) -SandboxConfig(...) -UserConfig(...) -MCPToolConfig(...) +class Env: + def __init__(*, taskset: Taskset, harness: Harness | None = None): ... + def run() -> EnvRun: ... + async def run_rollout(input, *, model, teacher=None, state=None) -> State: ... + async def score_group(tasks: list[Task], states: list[State], ...) -> list[State]: ... ``` -v1 config models are strict Pydantic models. Python code builds them directly, -and TOML config validates into the same models at the loader boundary. TOML -uses `"module:object"` refs for Python callables and loaders. Users are typed -`UserConfig` objects materialized through registered `User` subclasses, not -string refs. Unknown fields fail validation. +`Env` is the thin eval/training adapter for a concrete taskset/harness pair. --- @@ -1026,50 +968,48 @@ Provider-agnostic tool definition. Environments define tools using this type; ea ### v1 Config ```python -class Config(BaseModel): +class Config(BaseConfig): ... +@final class EnvConfig(Config): - taskset: TasksetConfig - harness: HarnessConfig + taskset: dict[str, object] = Field(default_factory=dict) + harness: dict[str, object] = Field(default_factory=dict) + runtime: RuntimeConfig | None = None + advantage: AdvantageConfig = "rl" class TasksetConfig(Config): - taskset_id: str | None = None # `id` shorthand accepted + id: str | None = None system_prompt: SystemPrompt = None + runtime: RuntimeConfig | None = None user: UserConfig | None = None - bindings: BindingsConfig = BindingsConfig() - objects: ObjectsConfig = ObjectsConfig() - artifacts: ArtifactsConfig = ArtifactsConfig() + toolsets: ToolsetConfigs = Field(default_factory=dict) + extras: Extras | None = None class HarnessConfig(Config): - harness_id: str | None = None # `id` shorthand accepted - program: ProgramConfig = ProgramConfig() - model: ModelConfig = ModelConfig() + id: str | None = None system_prompt: SystemPrompt = None system_prompt_strategy: SystemPromptStrategy = "HT" - sandbox: SandboxConfig | None = None - user: UserConfig | None = None - bindings: BindingsConfig = BindingsConfig() - objects: ObjectsConfig = ObjectsConfig() - artifacts: ArtifactsConfig = ArtifactsConfig() - max_turns: int = -1 # <= 0 means unbounded (run until a stop condition) + max_turns: int = -1 + runtime: RuntimeConfig | None = None + extras: Extras | None = None class ModelConfig(Config): - name: str | None = None - client: ClientConfig | str | None = None - sampling_args: SamplingArgs = {} + client: ClientConfig = ClientConfig() + model: str + sampling_args: JsonData = {} ``` -`EnvConfig` is the typed v1 loader envelope. TOML `[env.taskset]` and -`[env.harness]` sections populate `EnvConfig.taskset` and `EnvConfig.harness`. -The normal environment package loader stays typed as `load_environment(config: -vf.EnvConfig)` and delegates child coercion to `vf.load_taskset` / -`vf.load_harness`. Environment-specific fields belong on the taskset or harness -config that owns them; do not subclass `EnvConfig` just to narrow child config -types in ordinary environment packages. +`EnvConfig` is the typed v1 component-loader envelope. TOML `[env.taskset]` +and `[env.harness]` sections populate `EnvConfig.taskset` and +`EnvConfig.harness`. The package loader discovers `taskset.py` and optional +`harness.py`, then assembles `vf.Env` from the typed child loaders. +Environment-specific fields belong on the taskset or harness config that owns +them; `EnvConfig` is final. -`Config` subclasses are strict Pydantic config models. Validate raw mappings -with `MyConfig.model_validate(...)` or use the typed object directly. +`Config` is the v1 alias for `pydantic_config.BaseConfig`. Config subclasses +validate raw mappings with `MyConfig.model_validate(...)` or accept typed +objects directly. ### ClientConfig diff --git a/docs/training.md b/docs/training.md index 655d057c99..fba19dd629 100644 --- a/docs/training.md +++ b/docs/training.md @@ -100,15 +100,16 @@ max_turns = 8 system_prompt = "Answer exactly." [env.taskset.toolsets.search] -tools = ["my_env.tools:search"] -objects = { index = "my_env.tools:load_index" } -bindings = { "search.index" = "objects.index" } +scope = "rollout" [[env.taskset.rewards]] fn = "my_env.signals:exact_answer" weight = 1.0 ``` +This overrides a taskset-defined `search` toolset. To add a new toolset key +from config, set `source` to a `ToolsetConfig` class path under that key. + See [BYO Harness](byo-harness.md#toml-config) for the matching eval config shape and v1 callable/toolset patterns. @@ -177,7 +178,7 @@ In TOML configs, set GEPA parameters such as `max_calls`, `num_train`, `num_val` After optimization, you'll find: - `system_prompt.txt` - The optimized system prompt. For v1 environments, expose the owner prompt that GEPA should optimize as a `system_prompt` config field and default it to `vf.SystemPromptConfig(path="system_prompt.txt")` when the prompt should be file-backed. Override `load_system_prompt(config)` only when prompt loading is computed from config or package resources. -- `results.jsonl` - Candidate prompt rows for evaluation upload; GEPA-specific fields live under `info`. +- `results.jsonl` - Prompt rows for evaluation upload; GEPA-specific fields live under `info`. - `pareto_frontier.jsonl` - Best candidate references per validation example - `metadata.json` - Run configuration and summary diff --git a/environments/AGENTS.md b/environments/AGENTS.md index 681088a1d4..3498ea2b5e 100644 --- a/environments/AGENTS.md +++ b/environments/AGENTS.md @@ -293,13 +293,11 @@ async def my_reward_func(completion, my_helper) -> float: return await my_helper.score(completion) ``` -For taskset/harness environments, keep shared dependencies behind the taskset or -harness that owns them. Bindings are the canonical way to inject shared -resources into rewards, updates, tools, and programs. Configured binding -objects should use serializable loader paths when they cross a TOML or CLI -boundary; Python-only construction may use factory callables directly when a -resource cannot be serialized. Required Taskset and Toolset factory parameters -must be supplied through bindings. +For taskset/harness environments, keep shared dependencies behind the taskset, +harness, toolset, user, or runtime that owns them. v1 toolsets and users receive +shared rollout data through `@vf.tool(args=...)` and write serializable rollout +data through `sets` and `extends`. Configured objects use serializable loader +paths across TOML and CLI boundaries. Judges are used for tasks where deterministic evaluation is impractical, and an LLM is used to score responses. **JudgeRubric** stores an LLM client inside the @@ -702,23 +700,23 @@ environments/my_env/ ### v1 Env Shape -The v1 template teaches the standard object layout: one taskset class, one -typed `load_taskset(config: MyTasksetConfig)` child factory, and a tiny -`load_environment(config: vf.EnvConfig)` root loader that delegates through -`vf.load_taskset(config=config.taskset)` and -`vf.load_harness(config=config.harness)`. The child factory annotation defines -the taskset config type for TOML, CLI, eval, GEPA, RL, and Hosted Training. +v1 is component-first, under active development, and intentionally separate +from v0. v1 environment code imports `verifiers.v1 as vf`. Packages expose +`taskset.py` and, only when they own reusable execution behavior, `harness.py`. +They do not define a root `load_environment`; the library loader assembles +`vf.Env` from discovered components. Factory annotations define config types +for TOML, CLI, eval, GEPA, RL, and Hosted Training. After `prime env init my-env --v1`, edit the generated taskset class: 1. Add task settings to `TasksetConfig`. 2. Return task records from `load_tasks(split=...)`. -3. Return task-owned tools from `load_toolsets` when needed. +3. Add task-owned toolsets or a user when needed. 4. Add lifecycle, metric, reward, and advantage methods with `@vf.*`. Add a harness config, harness class, and `load_harness(config: MyHarnessConfig)` when the environment owns reusable rollout behavior. -Otherwise the generated root loader uses the base harness. +Otherwise the component loader uses the base harness. `EnvConfig` is the lightweight envelope for the two child configs. Put environment knobs on `TasksetConfig` or `HarnessConfig`. @@ -726,14 +724,20 @@ environment knobs on `TasksetConfig` or `HarnessConfig`. The taskset-only shape is: ```python -import verifiers as vf +import verifiers.v1 as vf class MyTasksetConfig(vf.TasksetConfig): system_prompt: vf.SystemPrompt = "Answer exactly." +class MyTask(vf.Task): + answer: str + + class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: """Return serializable task records as a list, generator, or Dataset.""" if split == "eval": @@ -747,24 +751,15 @@ class MyTaskset(vf.Taskset[MyTasksetConfig]): ] @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - if not messages: + async def correct_answer(self, task: MyTask, state: vf.State) -> float: + if not state.completion: return 0.0 - response = str(messages[-1].content or "").strip() - return float(response == task["answer"]) + response = str(state.completion[-1].content or "").strip() + return float(response == task.answer) def load_taskset(config: MyTasksetConfig) -> MyTaskset: return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` With a reusable harness, keep the same explicit object boundary: @@ -784,43 +779,17 @@ def load_taskset(config: MyTasksetConfig) -> MyTaskset: def load_harness(config: MyHarnessConfig) -> MyHarness: return MyHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -Keep v1 dependencies behind the owning taskset or harness. Do not pass -already-instantiated resource objects through environment loaders. Bindings are -allowed wherever the owning taskset, toolset, user, program, or harness wires -callables. `objects` entries should be loader specs: prefer serializable import -paths in config, and use factory callables directly only for Python-only -construction when the dependency cannot be serialized. Required Taskset and -Toolset factory parameters must be supplied through bindings. - -Judge-style rewards should read endpoint details from the rollout state: - -```python -@vf.reward(weight=1.0) -async def judge_reward(task, state) -> float: - endpoint = state.get_endpoint_config(api="chat") - client = state.get_client(api="chat") - model = str(task.get("judge_model") or endpoint.model) - ... ``` -Expose at most `judge_model: str | None = None` on the taskset config. Do not -add judge endpoint URL/API-key fields or read `os.environ` inside reward/update -handlers. +Keep v1 dependencies behind the owning taskset, harness, toolset, user, or runtime +provider. Config, task rows, state, tool specs, and user specs must stay +serializable. Live clients, runtimes, functions, and file handles do not +cross those boundaries. Custom mutable rollout data belongs in `state.extras`. For reusable tasksets and harnesses, [BYO Harness](byo-harness.md) is the canonical v1 implementation guide. It covers ownership, configs, task controls, -system prompts, users, toolsets, programs, sandboxes, artifacts, nested -harnesses, package adapters, and TOML/CLI overrides. +system prompts, users, toolsets, runtimes, artifacts, nested harnesses, package +adapters, and TOML/CLI overrides. ### pyproject.toml diff --git a/environments/README.md b/environments/README.md index 5737f725d7..9ed32b588f 100644 --- a/environments/README.md +++ b/environments/README.md @@ -1,6 +1,6 @@ # Environments -This folder contains installable example environments that showcase common usage patterns in Verifiers. Each module exposes a `load_environment(...)` function that returns a ready-to-use `vf.Environment` object. +This folder contains installable example environments that showcase common usage patterns in Verifiers. v0 packages expose `load_environment(...)`; v1 packages expose `taskset.py` and optional `harness.py` components and are assembled by the library loader. ## Quick start @@ -23,8 +23,8 @@ This folder contains installable example environments that showcase common usage - **sentence_repeater**: Multi-turn Q/A over a paragraph; rewards compare assistant messages to expected answers. - **wordle**: Game-style interaction via `TextArenaEnv`; multiple rewards (correctness, partial credit, few-turn bonus) and XML formatting. - **wordle_v1**: Wordle on the reusable v1 `TextArenaTaskset`, with Wordle-specific prompt, feedback, and rewards kept in the environment package. -- **openenv_echo**: OpenEnv MCP integration example using upstream `echo_env`. -- **openenv_textarena**: OpenEnv gym integration example using upstream `textarena_env` (default `Wordle-v0`). +- **openenv_echo_v1**: OpenEnv MCP integration example using upstream `echo_env`. +- **openenv_textarena_v1**: OpenEnv gym integration example using upstream `textarena_env` (default `Wordle-v0`). ### Tool use - **ToolEnv (native function-calling)** @@ -40,27 +40,27 @@ This folder contains installable example environments that showcase common usage ### Experimental environments - **MCPEnv (MCP server integration)** - - **mcp_search_env**: Example environment demonstrating `vf.MCPEnv` for Model Context Protocol server integration. + - **mcp_search_env_v1**: Example environment demonstrating `vf.MCPEnv` for Model Context Protocol server integration. - **RLM (Recursive Language Model)** - **hello_rlm_v1**: v1 packaged `RLM` harness example with endpoint interception and metrics collection. - **V1 Taskset/Harness** - - **dspy_rlm**: DSPy RLM harness on GSM8K through `vf.Env`; DSPy uses the V1 interception endpoint from rollout state. - - **openai_agents_env**: OpenAI Agents SDK harness with a calculator tool on GSM8K through `vf.Env`. - - **langchain_deep_agents_wikispeedia**: LangChain Deep Agents harness on Wikispeedia navigation, where tool use is load-bearing. + - **dspy_rlm_v1**: DSPy RLM harness on GSM8K through `vf.Env`; DSPy uses the V1 interception endpoint from rollout state. + - **openai_agents_env_v1**: OpenAI Agents SDK harness with a calculator tool on GSM8K through `vf.Env`. + - **langchain_deep_agents_wikispeedia_v1**: LangChain Deep Agents harness on Wikispeedia navigation, where tool use is load-bearing. - **HarborEnv / CliAgentEnv (agent sandboxes)** - - **opencode_harbor**: Runs the OpenCode CLI agent on Harbor tasks with API interception via Prime Tunnel. + - **opencode_harbor_v1**: Runs the OpenCode CLI agent on Harbor tasks with API interception via Prime Tunnel. - **terminus_harbor**: Runs the Terminus agent on Harbor tasks with API interception via Prime Tunnel. - **hello_mcp_harbor**: Smallest runnable `HarborEnv` exercising framework-managed MCP server lifecycle (FastMCP `get_secret` server + OpenCode agent). - **Taskset/Harness v1** - - **bfcl_v3**: BFCL v3 function-calling eval using task-local dynamic tool schemas and v1 rewards. - - **dspy_flights**: Sandboxed DSPy flight-support `program.fn` entrypoint installed from its package `pyproject.toml` and configured against the v1 interception endpoint. + - **bfcl_v3_v1**: BFCL v3 function-calling eval using task-local dynamic tool schemas and v1 rewards. + - **dspy_flights_v1**: Sandboxed DSPy flight-support `program.fn` entrypoint installed from its package `pyproject.toml` and configured against the v1 interception endpoint. - **hello_group_reward_v1**: Deterministic v1 reference for group updates, metrics, rewards, advantages, and cleanup. - - **nemo_gym_env**: Minimal v1 example that wraps a packaged NeMo Gym task with `NeMoGymTaskset` and `NeMoGymHarness`. - - **sft-replay**: Thin v1 replay environment using `ReplayTaskset` and `ReplayHarness` to turn stored transcripts into trajectory steps without model calls. + - **nemo_gym_env_v1**: Minimal v1 example that wraps a packaged NeMo Gym task with `NeMoGymTaskset` and `NeMoGymHarness`. + - **sft_replay_v1**: Thin v1 replay environment using `ReplayTaskset` and `ReplayHarness` to turn stored transcripts into rollout `Turn`s without model calls. - **tau2_bench_v1**: `tau2-bench-v1` τ²-bench taskset/user/tool pattern on the v1 harness runtime. - **wordle_v1**: TextArena Wordle through the packaged v1 `TextArenaTaskset` boundary. @@ -71,7 +71,7 @@ This folder contains installable example environments that showcase common usage - **Nested harnesses** - **hello_subagent_v1**: Minimal parent/child harness hand-off through a tool. - **nested_harness_v1**: v1 example showing a tool that calls a child `Harness` as its own rollout scope. - - **hello_self_judge_v1**: v1 example where a judge harness shares model, endpoint, trajectory, and sandbox evidence from the answer rollout. + - **hello_self_judge_v1**: v1 example where a judge harness shares model, endpoint, transcript, and sandbox evidence from the answer rollout. - **hello_parallel_sandbox_v1**: v1 example where parallel child harnesses share a sandbox-backed tool across update and reward stages. - **RubricGroup** @@ -87,22 +87,22 @@ This folder contains installable example environments that showcase common usage - **mmmu**: Demonstrates passing images via chat `content` items with `{type: "image_url", image_url: {url: ...}}` and standard answer parsing. ## What to look at for each pattern -- **Minimal SingleTurnEnv**: `reverse_text`, `gsm8k` +- **Minimal SingleTurnEnv**: `reverse_text_v1`, `gsm8k` - **JudgeRubric end-to-end**: `continuation_quality`, `toxicity_explanation`, `self_reward` -- **ToolEnv with real tools**: `wiki_search`, `math_python` -- **Custom MultiTurnEnv**: `alphabet_sort`, `doublecheck`, `sentence_repeater`, `wordle` +- **ToolEnv with real tools**: `wiki_search_v1`, `math_python_v1` +- **Custom MultiTurnEnv**: `alphabet_sort_v1`, `doublecheck`, `sentence_repeater`, `wordle` - **GymEnv integration**: `gem_wordle` -- **OpenEnv integration (gym + MCP)**: `openenv_textarena`, `openenv_echo` -- **CLI agent sandboxes**: `opencode_harbor`, `terminus_harbor`, `hello_mcp_harbor` -- **MCP integration**: `mcp_search_env`, `hello_mcp_harbor` -- **Taskset/Harness v1**: use this pattern for new environments that need reusable tasksets, reusable harnesses, framework programs, endpoint interception, or sandboxed Python/command programs. Examples include `dspy_rlm`, `openai_agents_env`, `langchain_deep_agents_wikispeedia`, `reverse_text`, `alphabet_sort`, `wiki_search`, `math_python`, `mcp_search_env`, `opencode_harbor`, `bfcl_v3`, `hello_subagent_v1`, `nested_harness_v1`, `hello_self_judge_v1`, `hello_parallel_sandbox_v1`, `hello_group_reward_v1`, `hello_rlm_v1`, `rlm_swe_v1`, `dspy_flights`, `tau2-bench-v1`, and `wordle-v1`. - - `opencode_harbor` uses the packaged `HarborTaskset` + `OpenCode` boundary from `tasksets` and `harnesses`. -- **Environment and rubric composition**: `math_group`, `math_python`, `wiki_search` +- **OpenEnv integration (gym + MCP)**: `openenv_textarena_v1`, `openenv_echo_v1` +- **CLI agent sandboxes**: `opencode_harbor_v1`, `terminus_harbor`, `hello_mcp_harbor` +- **MCP integration**: `mcp_search_env_v1`, `hello_mcp_harbor` +- **Taskset/Harness v1**: use this pattern for new environments that need reusable tasksets, reusable harnesses, framework programs, endpoint interception, or sandboxed Python/command programs. Examples include `dspy_rlm_v1`, `openai_agents_env_v1`, `langchain_deep_agents_wikispeedia_v1`, `reverse_text_v1`, `alphabet_sort_v1`, `wiki_search_v1`, `math_python_v1`, `mcp_search_env_v1`, `opencode_harbor_v1`, `bfcl_v3_v1`, `hello_subagent_v1`, `nested_harness_v1`, `hello_self_judge_v1`, `hello_parallel_sandbox_v1`, `hello_group_reward_v1`, `hello_rlm_v1`, `rlm_swe_v1`, `dspy_flights_v1`, `tau2-bench-v1`, and `wordle-v1`. + - `opencode_harbor_v1` uses the packaged `HarborTaskset` + `OpenCode` boundary from `tasksets` and `harnesses`. +- **Environment and rubric composition**: `math_group`, `math_python_v1`, `wiki_search_v1` - **Procedural datasets**: `reasoning_gym_env` - **Multimodal**: `mmmu` ## Running examples -All environments export `load_environment(...)`. +All examples are loadable through `vf.load_environment(...)`. v0 packages do this with a package `load_environment(...)`; v1 packages do it through discovered taskset/harness components. In-line usage: ```python diff --git a/environments/alphabet_sort/alphabet_sort.py b/environments/alphabet_sort/alphabet_sort.py index fea91479f3..4f258a7e1c 100644 --- a/environments/alphabet_sort/alphabet_sort.py +++ b/environments/alphabet_sort/alphabet_sort.py @@ -200,35 +200,11 @@ def load_environment( dataset_name: str = "kalomaze/alphabetic-arxiv-authors-it1", dataset_split: str = "train", seed: int = 1337420, - v1: bool = False, **kwargs, ) -> vf.Environment: - if v1: - if kwargs: - unexpected = ", ".join(sorted(kwargs)) - raise TypeError(f"Unsupported v1 load_environment kwargs: {unexpected}") - - from alphabet_sort_v1 import ( - AlphabetSortEnvConfig, - AlphabetSortTasksetConfig, - load_environment as load_v1, - ) - - return load_v1( - config=AlphabetSortEnvConfig( - taskset=AlphabetSortTasksetConfig( - max_turns=max_turns, - min_turns=min_turns, - min_names_per_turn=min_names_per_turn, - max_names_per_turn=max_names_per_turn, - similarity_power=similarity_power, - power_per_turn=power_per_turn, - dataset_name=dataset_name, - dataset_split=dataset_split, - seed=seed, - ) - ) - ) + if kwargs: + unexpected = ", ".join(sorted(kwargs)) + raise TypeError(f"Unsupported load_environment kwargs: {unexpected}") assert min_turns >= 1, "min_turns must be at least 1" assert min_turns <= max_turns, "min_turns must be less than or equal to max_turns" diff --git a/environments/alphabet_sort/pyproject.toml b/environments/alphabet_sort/pyproject.toml index e56b652573..67c645f1c1 100644 --- a/environments/alphabet_sort/pyproject.toml +++ b/environments/alphabet_sort/pyproject.toml @@ -13,7 +13,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["alphabet_sort.py", "alphabet_sort_v1.py"] +include = ["alphabet_sort.py"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/alphabet_sort_v1/README.md b/environments/alphabet_sort_v1/README.md new file mode 100644 index 0000000000..9d61019b12 --- /dev/null +++ b/environments/alphabet_sort_v1/README.md @@ -0,0 +1,55 @@ +# alphabet-sort-v1 + + +Source Code + + +### Overview +- **Environment ID**: `alphabet-sort-v1` +- **Short description**: This task requires the model to maintain and update an alphabetically sorted list of names across multiple conversation turns, with new names being tagged appropriately. The dataset uses real author names from arXiv papers, with 1-3 turns per conversation and 2-5 total names (the turn and name counts are randomized during the data creation process by default). +- **Tags**: sorting, names, multi-turn, xml, synthetic, tools + +### Datasets +- **Primary dataset(s)**: `kalomaze/alphabetic-arxiv-authors-it1` (HF) used to sample name lists +- **Source links**: Hugging Face Datasets +- **Split sizes**: Procedurally constructs multi-turn sessions from the `train` split + +### Task +- **Type**: multi-turn +- **Rubric overview**: The reward function uses difflib to calculate sequence similarity between predicted and expected outputs, with the final score raised to the nth power (similarity_power, defaults to 4) to emphasize precision. + +### Quickstart +Run an evaluation with default settings: + +```bash +prime eval run alphabet-sort-v1 +``` + +Configure model and sampling: + +```bash +prime eval run alphabet-sort-v1 \ + -m openai/gpt-4.1-mini \ + -n 20 -r 3 -t 1024 -T 0.7 \ + -a '{"config": {"taskset": {"max_turns": 3, "min_turns": 1, "min_names_per_turn": 1, "max_names_per_turn": 5, "similarity_power": 4}}}' +``` + +Notes: +- v1 task settings belong under `config.taskset` when passed through `-a` / `--env-args`. + +### Taskset Config +| Arg | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `max_turns` | int | `3` | Maximum number of assistant turns | +| `min_turns` | int | `1` | Minimum number of assistant turns | +| `min_names_per_turn` | int | `1` | Minimum names per turn | +| `max_names_per_turn` | int | `5` | Maximum names per turn | +| `similarity_power` | int | `4` | Exponent applied to sequence similarity | +| `power_per_turn` | bool | `True` | Apply power scaling per turn (True) or to final average (False) | +| `hf_dataset_path` | str | `"kalomaze/alphabetic-arxiv-authors-it1"` | HF dataset path for names | +| `seed` | int | `1337420` | Random seed for dataset construction | + +### Metrics +| Metric | Meaning | +| ------ | ------- | +| `reward` | Average per-turn sequence similarity raised to `similarity_power` | diff --git a/environments/alphabet_sort_v1/alphabet_sort_v1/__init__.py b/environments/alphabet_sort_v1/alphabet_sort_v1/__init__.py new file mode 100644 index 0000000000..1faa179949 --- /dev/null +++ b/environments/alphabet_sort_v1/alphabet_sort_v1/__init__.py @@ -0,0 +1 @@ +"""alphabet-sort-v1 environment package.""" diff --git a/environments/alphabet_sort_v1/alphabet_sort_v1/servers/__init__.py b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/__init__.py new file mode 100644 index 0000000000..d87ea669ab --- /dev/null +++ b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for alphabet-sort-v1.""" diff --git a/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/__init__.py b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/__init__.py new file mode 100644 index 0000000000..be0177256d --- /dev/null +++ b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/__init__.py @@ -0,0 +1,3 @@ +from .config import UserConfig + +__all__ = ["UserConfig"] diff --git a/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/config.py b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/config.py new file mode 100644 index 0000000000..57349dd33b --- /dev/null +++ b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/config.py @@ -0,0 +1,5 @@ +import verifiers.v1 as vf + + +class UserConfig(vf.UserConfig): + pass diff --git a/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/user.py b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/user.py new file mode 100644 index 0000000000..109810dfed --- /dev/null +++ b/environments/alphabet_sort_v1/alphabet_sort_v1/servers/user/user.py @@ -0,0 +1,27 @@ +import verifiers.v1 as vf + +from .config import UserConfig + + +class User(vf.User[UserConfig]): + @vf.user( + args={ + "info": "task.info", + "transcript": "state.transcript", + } + ) + def respond(self, info: dict, transcript: list[dict]) -> dict: + follow_ups = info.get("follow_ups") or [] + assistant_count = 0 + for turn in transcript: + completion = turn.get("completion") or [] + assistant_count += sum( + 1 for message in completion if message.get("role") == "assistant" + ) + if assistant_count <= 0 or assistant_count > len(follow_ups): + return {"messages": []} + return { + "messages": [ + {"role": "user", "content": str(follow_ups[assistant_count - 1])} + ], + } diff --git a/environments/alphabet_sort/alphabet_sort_v1.py b/environments/alphabet_sort_v1/alphabet_sort_v1/taskset.py similarity index 82% rename from environments/alphabet_sort/alphabet_sort_v1.py rename to environments/alphabet_sort_v1/alphabet_sort_v1/taskset.py index bc16fbd0b5..3dd6011169 100644 --- a/environments/alphabet_sort/alphabet_sort_v1.py +++ b/environments/alphabet_sort_v1/alphabet_sort_v1/taskset.py @@ -5,9 +5,11 @@ import re from datasets import Dataset, load_dataset -from pydantic import model_validator +from pydantic import BaseModel -import verifiers as vf +import verifiers.v1 as vf + +from .servers.user import UserConfig logger = logging.getLogger(__name__) @@ -239,19 +241,20 @@ def score_response( def eval_turn( - completion: list[vf.ConfigData], + completion: vf.Messages, turn_num: int, - state: dict, + info: "AlphabetSortInfo", similarity_power: int, apply_power: bool, ) -> float: - ground_truths = state.get("info", {}).get("ground_truths", []) + ground_truths = info.ground_truths if turn_num > len(ground_truths): return 0.0 expected = ground_truths[turn_num - 1] assistant_msgs = [ str(message.content or "") - for message in vf.get_messages(completion, role="assistant") + for message in completion + if message.role == "assistant" ] if len(assistant_msgs) < turn_num: return 0.0 @@ -279,44 +282,30 @@ def eval_turn( return attempt_scores[-1] -@vf.reward(weight=1.0) -async def weighted_reward(task, state) -> float: - completion = state.get("completion") or [] - actual_turns = state["info"]["num_turns"] - similarity_power = int(task.get("similarity_power", 4)) - power_per_turn = bool(task.get("power_per_turn", True)) - total = 0.0 - for turn_num in range(1, actual_turns + 1): - total += eval_turn( - completion, - turn_num, - state, - similarity_power, - apply_power=power_per_turn, - ) - if actual_turns <= 0: - return 0.0 - if power_per_turn: - return total / actual_turns - return (total / actual_turns) ** similarity_power +class AlphabetSortInfo(BaseModel, extra="forbid"): + follow_ups: list[str] + turn_names: list[list[str]] + ground_truths: list[list[str]] + num_turns: int + sort_by_first: bool -class AlphabetUserConfig(vf.UserConfig): - pass +class AlphabetSortTask(vf.Task): + answer: str + info: AlphabetSortInfo + similarity_power: int = 4 + power_per_turn: bool = True -class AlphabetUser(vf.User[AlphabetUserConfig]): - async def get_response(self, task, state, messages) -> list[dict[str, str]]: - assistant_count = len(vf.get_messages(messages, role="assistant")) - follow_ups = state["info"]["follow_ups"] - if assistant_count <= 0 or assistant_count > len(follow_ups): - return [] - return [{"role": "user", "content": follow_ups[assistant_count - 1]}] +def transcript_completion_messages(state: vf.State) -> vf.Messages: + messages: vf.Messages = [] + for turn in state.transcript: + messages.extend(turn.completion) + return messages class AlphabetSortTasksetConfig(vf.TasksetConfig): - user: AlphabetUserConfig | None = AlphabetUserConfig() - rewards: list[str] = ["weighted_reward"] + user: vf.UserConfig | None = UserConfig() min_turns: int = 1 max_turns: int = 3 min_names_per_turn: int = 1 @@ -327,23 +316,10 @@ class AlphabetSortTasksetConfig(vf.TasksetConfig): dataset_split: str = "train" seed: int = 1337420 - @model_validator(mode="after") - def validate_task_shape(self) -> "AlphabetSortTasksetConfig": - validate_parameters( - min_turns=self.min_turns, - max_turns=self.max_turns, - min_names_per_turn=self.min_names_per_turn, - max_names_per_turn=self.max_names_per_turn, - ) - return self - - -class AlphabetSortEnvConfig(vf.EnvConfig): - taskset: AlphabetSortTasksetConfig = AlphabetSortTasksetConfig() - harness: vf.HarnessConfig = vf.HarnessConfig() - class AlphabetSortTaskset(vf.Taskset[AlphabetSortTasksetConfig]): + task_type = AlphabetSortTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: return load_tasks( min_turns=self.config.min_turns, @@ -357,9 +333,25 @@ def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: seed=self.config.seed, ) + @vf.reward(weight=1.0) + async def weighted_reward(self, task: AlphabetSortTask, state: vf.State) -> float: + completion = transcript_completion_messages(state) + actual_turns = task.info.num_turns + total = 0.0 + for turn_num in range(1, actual_turns + 1): + total += eval_turn( + completion, + turn_num, + task.info, + task.similarity_power, + apply_power=task.power_per_turn, + ) + if actual_turns <= 0: + return 0.0 + if task.power_per_turn: + return total / actual_turns + return (total / actual_turns) ** task.similarity_power -def load_environment(config: AlphabetSortEnvConfig) -> vf.Env: - return vf.Env( - taskset=AlphabetSortTaskset(config=config.taskset), - harness=vf.Harness(config=config.harness), - ) + +def load_taskset(config: AlphabetSortTasksetConfig) -> AlphabetSortTaskset: + return AlphabetSortTaskset(config=config) diff --git a/environments/alphabet_sort_v1/pyproject.toml b/environments/alphabet_sort_v1/pyproject.toml new file mode 100644 index 0000000000..19b12ba7ca --- /dev/null +++ b/environments/alphabet_sort_v1/pyproject.toml @@ -0,0 +1,20 @@ +[project] +name = "alphabet-sort-v1" +version = "0.1.12" +tags = ["sorting", "names", "multi-turn", "xml", "synthetic", "tools"] +license = "Apache-2.0" +description = "This task requires the model to maintain and update an alphabetically sorted list of names across multiple conversation turns." +dependencies = [ + "verifiers>=0.1.9", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build] +include = ["alphabet_sort_v1/**/*", "README.md", "pyproject.toml"] + +[tool.verifiers.eval] +num_examples = 5 +rollouts_per_example = 3 diff --git a/environments/bfcl_v3/bfcl_v3.py b/environments/bfcl_v3/bfcl_v3.py deleted file mode 100644 index 84258edcfd..0000000000 --- a/environments/bfcl_v3/bfcl_v3.py +++ /dev/null @@ -1,622 +0,0 @@ -import json -import re -from collections.abc import Sequence -from typing import cast - -import verifiers as vf -from verifiers.types import ( - AssistantMessage, - MessageContent, - Messages, - Tool, - ToolCall, - ToolMessage, - UserMessage, -) -from verifiers.utils.message_utils import message_role, normalize_messages - -from verifiers.v1.utils.endpoint_utils import assistant_completion_from_messages -from verifiers.v1.utils.config_utils import explicit_config_data -from verifiers.v1.utils.json_utils import json_args - -_BFCL_PATCHED = False -BFCLRawMessage = str | vf.JsonData -BFCLRawTurn = str | vf.JsonData | Sequence[BFCLRawMessage] | None - - -class BFCLTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["bfcl_reward"] - test_category: str = "simple_python" - test_categories: list[str] | None = None - examples_per_category: int = -1 - - -class BFCLHarnessConfig(vf.HarnessConfig): - test_category: str = "simple_python" - - -class BFCLEnvConfig(vf.EnvConfig): - taskset: BFCLTasksetConfig = BFCLTasksetConfig() - harness: BFCLHarnessConfig = BFCLHarnessConfig() - - -def modded_convert_func_name(function_name: str, model_name: str) -> str: - _ = model_name - return re.sub(r"\.", "_", function_name) - - -def patch_bfcl_eval() -> None: - global _BFCL_PATCHED - if _BFCL_PATCHED: - return - import bfcl_eval.constants.category_mapping as category_mapping - import bfcl_eval.eval_checker.ast_eval.ast_checker as ast_checker_module - - agentic_categories = category_mapping.AGENTIC_CATEGORY.copy() - category_mapping.AGENTIC_CATEGORY.clear() - for category in agentic_categories: - if category in category_mapping.ALL_SCORING_CATEGORIES: - category_mapping.ALL_SCORING_CATEGORIES.remove(category) - if category in category_mapping.ALL_CATEGORIES: - category_mapping.ALL_CATEGORIES.remove(category) - category_mapping.TEST_COLLECTION_MAPPING.pop("memory", None) - category_mapping.TEST_COLLECTION_MAPPING.pop("web_search", None) - category_mapping.TEST_COLLECTION_MAPPING.pop("agentic", None) - - non_scoring_categories = category_mapping.NON_SCORING_CATEGORY.copy() - category_mapping.NON_SCORING_CATEGORY.clear() - for category in non_scoring_categories: - if category in category_mapping.ALL_CATEGORIES: - category_mapping.ALL_CATEGORIES.remove(category) - category_mapping.TEST_COLLECTION_MAPPING.pop("format_sensitivity", None) - - setattr(ast_checker_module, "convert_func_name", modded_convert_func_name) - _BFCL_PATCHED = True - - -def bfcl_tool_defs(functions: object) -> list[Tool]: - patch_bfcl_eval() - from bfcl_eval.constants.enums import ModelStyle - from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI - from bfcl_eval.model_handler.utils import convert_to_tool - - oai_tools = convert_to_tool( - functions, GORILLA_TO_OPENAPI, ModelStyle.OPENAI_COMPLETIONS - ) - tool_defs = [] - for tool in oai_tools: - function = tool["function"] - tool_defs.append( - Tool( - name=str(function["name"]), - description=str(function.get("description") or ""), - parameters=dict(cast(vf.JsonData, function["parameters"])), - strict=False, - ) - ) - return tool_defs - - -class BFCLSchemaTool: - def __init__(self, tool_def: Tool): - self.name = tool_def.name - self.__name__ = tool_def.name - self.__doc__ = tool_def.description - self.tool_def = tool_def - - async def __call__(self, state: vf.State, **arguments: object) -> str: - calls = cast( - list[vf.ConfigData], state.setdefault("bfcl_executed_tool_calls", []) - ) - calls.append({self.name: arguments}) - return "recorded" - - -def bfcl_functions(task: vf.JsonData) -> object: - return task.get("function_with_hints") or task["function"] - - -def bfcl_missed_function(task: vf.JsonData) -> vf.JsonData: - value = task.get("missed_function_with_hints") or task.get("missed_function") or {} - if not isinstance(value, dict): - raise TypeError("BFCL missed_function must be a mapping.") - return cast(vf.JsonData, value) - - -def build_task_loader(test_category: str, examples_per_category: int = -1): - def factory(): - patch_bfcl_eval() - from bfcl_eval.utils import ( - is_multi_turn, - is_relevance_or_irrelevance, - load_dataset_entry, - load_ground_truth_entry, - ) - - entries = load_dataset_entry( - test_category, include_language_specific_hint=False - ) - entries_with_hints = load_dataset_entry( - test_category, include_language_specific_hint=True - ) - if is_relevance_or_irrelevance(test_category): - ground_truth_entries = [None] * len(entries) - else: - ground_truth_entries = load_ground_truth_entry(test_category) - limit = len(entries) if examples_per_category < 0 else examples_per_category - rows = [] - for index, (entry, hinted_entry, ground_truth) in enumerate( - zip(entries, entries_with_hints, ground_truth_entries) - ): - if index >= limit: - break - row = bfcl_row( - test_category, - entry, - hinted_entry, - cast(vf.JsonData | None, ground_truth), - ) - if is_multi_turn(test_category): - max_steps = maximum_step_limit() - row["max_steps_per_turn"] = max_steps - row["max_turns"] = ( - len(cast(Sequence[BFCLRawTurn], row["question"])) * max_steps - ) - else: - row["max_turns"] = 1 - rows.append(row) - return rows - - return factory - - -def load_tasks(test_category: str = "simple_python", examples_per_category: int = -1): - return build_task_loader(test_category, examples_per_category)() - - -def bfcl_row( - test_category: str, - entry: vf.JsonData, - hinted_entry: vf.JsonData, - ground_truth: vf.JsonData | None, -) -> vf.ConfigData: - question = cast(list[BFCLRawTurn], entry["question"]) - first_turn_system_prompt, first_turn_prompt = split_system_prompt( - normalize_turn(question[0]) - ) - row: vf.ConfigData = { - "task_id": str(entry["id"]), - "id": str(entry["id"]), - "category": test_category, - "prompt": first_turn_prompt, - "question": [ - first_turn_prompt, - *[normalize_turn(turn) for turn in question[1:]], - ], - "function": entry["function"], - "function_with_hints": hinted_entry["function"], - } - if first_turn_system_prompt: - row["system_prompt"] = first_turn_system_prompt - for key in ( - "initial_config", - "involved_classes", - ): - if key in entry: - row[key] = entry[key] - if "missed_function" in entry: - row["missed_function"] = entry["missed_function"] - if "missed_function" in hinted_entry: - row["missed_function_with_hints"] = hinted_entry["missed_function"] - if ground_truth is not None: - for key, value in ground_truth.items(): - row[key] = value - return row - - -def normalize_turn(value: object) -> list[vf.ConfigData]: - if value is None: - return [] - if isinstance(value, str): - return [{"role": "user", "content": value}] - if isinstance(value, dict): - return [dict(cast(vf.JsonData, value))] - if isinstance(value, Sequence): - messages = [] - for item in value: - if isinstance(item, str): - messages.append({"role": "user", "content": item}) - elif isinstance(item, dict): - messages.append(dict(cast(vf.JsonData, item))) - else: - raise TypeError(f"Unsupported BFCL message item: {type(item).__name__}") - return messages - raise TypeError(f"Unsupported BFCL prompt turn: {type(value).__name__}") - - -def split_system_prompt( - messages: Sequence[vf.JsonData], -) -> tuple[list[vf.ConfigData], list[vf.ConfigData]]: - system_prompt = [] - prompt = [] - for message in messages: - target = system_prompt if message.get("role") == "system" else prompt - target.append(dict(message)) - return system_prompt, prompt - - -def maximum_step_limit() -> int: - patch_bfcl_eval() - from bfcl_eval.constants.default_prompts import MAXIMUM_STEP_LIMIT - - return cast(int, MAXIMUM_STEP_LIMIT) - - -def model_name(state: vf.JsonData) -> str: - runtime = state.get("runtime") or {} - if isinstance(runtime, dict): - runtime_map = cast(vf.JsonData, runtime) - model = runtime_map.get("model") - if isinstance(model, str) and model: - return model - return "unknown" - - -def assistant_tool_calls(state: vf.JsonData) -> list[ToolCall]: - completion = state.get("completion") or [] - if not isinstance(completion, Sequence): - return [] - messages = vf.get_messages(completion, role="assistant") - if not messages: - return [] - return parse_tool_calls(messages[-1]) - - -def parse_tool_calls(message: object) -> list[ToolCall]: - if isinstance(message, AssistantMessage): - return list(message.tool_calls or []) - raw_tool_calls: object - if isinstance(message, dict): - message_map = cast(vf.JsonData, message) - raw_tool_calls = message_map.get("tool_calls") or [] - else: - raw_tool_calls = getattr(message, "tool_calls", []) or [] - if not isinstance(raw_tool_calls, Sequence): - return [] - calls = [] - for raw_call in raw_tool_calls: - if isinstance(raw_call, ToolCall): - calls.append(raw_call) - continue - if not isinstance(raw_call, dict): - continue - raw_call = cast(vf.JsonData, raw_call) - function = raw_call.get("function") - if isinstance(function, dict): - function_map = cast(vf.JsonData, function) - name = str(function_map.get("name") or "") - arguments = function_map.get("arguments") or "{}" - else: - name = str(raw_call.get("name") or "") - arguments = raw_call.get("arguments") or "{}" - if not name: - continue - calls.append( - ToolCall( - id=str(raw_call.get("id") or name), - name=name, - arguments=arguments - if isinstance(arguments, str) - else json.dumps(arguments), - ) - ) - return calls - - -def convert_to_gorilla(tool_calls: list[ToolCall]) -> list[vf.ConfigData]: - decoded_output = [] - for tool_call in tool_calls: - decoded_output.append({tool_call.name: json_args(tool_call.arguments)}) - return decoded_output - - -def convert_to_func_calls(tool_calls: list[ToolCall]) -> list[str]: - func_calls = [] - for tool_call in tool_calls: - params = json_args(tool_call.arguments) - args = ",".join(f"{key}={value!r}" for key, value in params.items()) - func_calls.append(f"{tool_call.name}({args})") - return func_calls - - -def json_clone(value: object) -> object: - return json.loads(json.dumps(value)) - - -@vf.reward(weight=1.0) -async def bfcl_reward(task: vf.Task, state: vf.State) -> float: - patch_bfcl_eval() - from bfcl_eval.utils import is_multi_turn, is_relevance_or_irrelevance - - category = str(task["category"]) - if is_relevance_or_irrelevance(category): - return relevance_reward(task, state) - if is_multi_turn(category): - return multi_turn_reward(task, state) - return ast_reward(task, state) - - -def relevance_reward(task: vf.JsonData, state: vf.JsonData) -> float: - patch_bfcl_eval() - from bfcl_eval.utils import is_empty_output - - category = str(task["category"]) - try: - gorilla_tool_calls = convert_to_gorilla(assistant_tool_calls(state)) - contain_func_call = not is_empty_output(gorilla_tool_calls) - except Exception: - contain_func_call = False - if "irrelevance" in category: - return float(not contain_func_call) - return float(contain_func_call) - - -def ast_reward(task: vf.JsonData, state: vf.JsonData) -> float: - patch_bfcl_eval() - from bfcl_eval.constants.enums import Language - from bfcl_eval.eval_checker.ast_eval.ast_checker import ast_checker - from bfcl_eval.utils import ( - is_function_calling_format_output, - is_java, - is_js, - ) - - category = str(task["category"]) - try: - gorilla_tool_calls = convert_to_gorilla(assistant_tool_calls(state)) - if not is_function_calling_format_output(gorilla_tool_calls): - return 0.0 - except Exception: - return 0.0 - - if is_java(category): - language = Language.JAVA - elif is_js(category): - language = Language.JAVASCRIPT - else: - language = Language.PYTHON - - checker_result = ast_checker( - task["function"], - gorilla_tool_calls, - task["ground_truth"], - language, - category, - model_name(state), - ) - return float(bool(checker_result["valid"])) - - -def multi_turn_reward(task: vf.JsonData, state: vf.JsonData) -> float: - patch_bfcl_eval() - from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_checker import ( - multi_turn_checker, - ) - from bfcl_eval.model_handler.base_handler import is_empty_execute_response - - completion = state.get("completion") or [] - if not isinstance(completion, Sequence): - return 0.0 - raw_ground_truth = task["ground_truth"] - if not isinstance(raw_ground_truth, Sequence): - return 0.0 - all_ground_truth = cast(list[list[str]], raw_ground_truth) - all_func_calls: list[list[list[str]]] = [[]] - try: - for message in completion: - role = message_role(message) - if role == "user": - all_func_calls.append([]) - elif role == "tool": - continue - elif role == "assistant": - func_calls = convert_to_func_calls(parse_tool_calls(message)) - if is_empty_execute_response(func_calls): - continue - all_func_calls[-1].append(func_calls) - elif role == "system": - continue - else: - return 0.0 - except Exception: - return 0.0 - - if len(all_func_calls) != len(all_ground_truth): - return 0.0 - - result = multi_turn_checker( - all_func_calls, - all_ground_truth, - { - "initial_config": task.get("initial_config", {}), - "involved_classes": task["involved_classes"], - "id": task["id"], - }, - str(task["id"]).rsplit("_", 1)[0], - model_name(state), - ) - return float(bool(result["valid"])) - - -async def bfcl_multi_turn_program( - task: vf.Task, state: vf.State, harness: vf.Harness -) -> vf.State: - patch_bfcl_eval() - from bfcl_eval.constants.default_prompts import ( - DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC, - ) - from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import ( - execute_multi_turn_func_call, - ) - from bfcl_eval.model_handler.base_handler import is_empty_execute_response - - messages = [ - *normalize_messages( - state.get("system_prompt", []), field_name="state.system_prompt" - ), - *normalize_messages(state.get("prompt", []), field_name="state.prompt"), - ] - prompt_messages = [message.model_dump(exclude_none=True) for message in messages] - - def sync_completion() -> list[vf.ConfigData]: - rendered_messages = [ - message.model_dump(exclude_none=True) for message in messages - ] - state["completion"] = assistant_completion_from_messages( - prompt_messages, rendered_messages - ) - return rendered_messages - - category = str(task["category"]) - tool_defs = bfcl_tool_defs(bfcl_functions(task)) - next_prompts = list(cast(Sequence[list[vf.ConfigData]], task["question"]))[1:] - holdout_function = bfcl_missed_function(task) - initial_config = cast(vf.ConfigData, json_clone(task.get("initial_config") or {})) - involved_classes = cast(list[str], json_clone(task["involved_classes"])) - max_steps_per_turn = int(task.get("max_steps_per_turn") or maximum_step_limit()) - turn_idx = 0 - steps_per_turn = 0 - runtime = harness.runtime - - execute_multi_turn_func_call( - [], - initial_config, - involved_classes, - model_name(state).replace("/", "_").replace("-", "_").replace(".", "_"), - str(task["id"]), - long_context=("long_context" in category or "composite" in category), - ) - - while True: - if await runtime.is_completed(task, state): - return state - response = await runtime.submit_model_request( - cast(Messages, messages), - task, - state, - tool_defs=tool_defs, - ) - messages.append(response.message) - sync_completion() - tool_calls = list(response.message.tool_calls or []) - try: - func_calls = convert_to_func_calls(tool_calls) - if is_empty_execute_response(func_calls): - func_calls = None - except Exception: - func_calls = None - - if func_calls: - execution_results, _ = execute_multi_turn_func_call( - func_call_list=func_calls, - initial_config=initial_config, - involved_classes=involved_classes, - model_name=model_name(state) - .replace("/", "_") - .replace("-", "_") - .replace(".", "_"), - test_entry_id=str(task["id"]), - long_context=("long_context" in category or "composite" in category), - ) - for execution_result, tool_call in zip(execution_results, tool_calls): - messages.append( - ToolMessage( - tool_call_id=tool_call.id, - content=cast(MessageContent, execution_result), - ) - ) - sync_completion() - steps_per_turn += 1 - if steps_per_turn >= max_steps_per_turn: - state.stop("max_steps_per_turn_reached") - return state - continue - - steps_per_turn = 0 - turn_idx += 1 - if not next_prompts: - state.stop("no_next_prompt_and_no_tool_calls") - return state - next_prompt = normalize_turn(next_prompts.pop(0)) - if str(turn_idx) in holdout_function: - tool_defs.extend(bfcl_tool_defs(holdout_function[str(turn_idx)])) - if next_prompt: - raise ValueError("BFCL holdout turns must not include user messages.") - messages.append( - UserMessage(content=DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC) - ) - else: - messages.extend(normalize_messages(cast(Messages, next_prompt))) - sync_completion() - - -class BFCLTaskset(vf.Taskset[BFCLTasksetConfig]): - def load_toolsets(self, config: BFCLTasksetConfig) -> vf.Toolsets: - _ = config - return {"bfcl": vf.Toolset(scope="rollout")} - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks( - test_category=self.config.test_category, - examples_per_category=self.config.examples_per_category, - ) - - @vf.setup - async def setup_bfcl_tools(self, task: vf.Task, state: vf.State) -> None: - patch_bfcl_eval() - from bfcl_eval.utils import is_multi_turn - - if is_multi_turn(str(task["category"])): - return - for tool_def in bfcl_tool_defs(bfcl_functions(task)): - state.add_tool("bfcl", BFCLSchemaTool(tool_def)) - - -def load_harness(config: BFCLHarnessConfig) -> vf.Harness: - patch_bfcl_eval() - from bfcl_eval.utils import is_multi_turn - - if is_multi_turn(config.test_category): - config = config.model_copy( - update={"program": vf.ProgramConfig(fn="bfcl_multi_turn_program")} - ) - return vf.Harness(config=config) - - -def load_environment(config: BFCLEnvConfig) -> vf.Env | vf.EnvGroup: - taskset_template = config.taskset - harness_template = config.harness - categories = taskset_template.test_categories or [taskset_template.test_category] - envs: list[vf.Env] = [] - for category in categories: - taskset_config = BFCLTasksetConfig.model_validate( - { - **explicit_config_data(taskset_template), - "test_category": category, - } - ) - harness_config = BFCLHarnessConfig.model_validate( - { - **explicit_config_data(harness_template), - "test_category": category, - } - ) - envs.append( - vf.Env( - taskset=BFCLTaskset(config=taskset_config), - harness=load_harness(config=harness_config), - ) - ) - if taskset_template.test_categories is not None: - return vf.EnvGroup(envs=envs, env_names=categories) - return envs[0] diff --git a/environments/bfcl_v3/README.md b/environments/bfcl_v3_v1/README.md similarity index 59% rename from environments/bfcl_v3/README.md rename to environments/bfcl_v3_v1/README.md index e8ac913f9e..e766a79154 100644 --- a/environments/bfcl_v3/README.md +++ b/environments/bfcl_v3_v1/README.md @@ -1,11 +1,14 @@ -# bfcl-v3 +# bfcl-v3-v1 Berkeley Function Calling Leaderboard v3 on the v1 Taskset/Harness runtime. ```bash -prime eval run bfcl-v3 -a '{"test_category": "simple_python"}' +prime eval run bfcl-v3-v1 -a '{"config": {"taskset": {"test_category": "simple_python"}}}' ``` Single-turn categories provision schema-backed tools into a taskset-owned rollout toolset. Multi-turn categories use a taskset-owned custom harness program for BFCL's official tool execution loop. + +Configure one category per v1 eval. Use multiple eval entries to run multiple +BFCL categories. diff --git a/environments/bfcl_v3_v1/bfcl_v3_v1/__init__.py b/environments/bfcl_v3_v1/bfcl_v3_v1/__init__.py new file mode 100644 index 0000000000..bf7978f0cf --- /dev/null +++ b/environments/bfcl_v3_v1/bfcl_v3_v1/__init__.py @@ -0,0 +1 @@ +"""bfcl-v3-v1 environment package.""" diff --git a/environments/bfcl_v3_v1/bfcl_v3_v1/harness.py b/environments/bfcl_v3_v1/bfcl_v3_v1/harness.py new file mode 100644 index 0000000000..930c8921f3 --- /dev/null +++ b/environments/bfcl_v3_v1/bfcl_v3_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import BFCLHarness as BFCLHarness +from .taskset import BFCLHarnessConfig as BFCLHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/bfcl_v3_v1/bfcl_v3_v1/taskset.py b/environments/bfcl_v3_v1/bfcl_v3_v1/taskset.py new file mode 100644 index 0000000000..31533eefb5 --- /dev/null +++ b/environments/bfcl_v3_v1/bfcl_v3_v1/taskset.py @@ -0,0 +1,687 @@ +import json +import re +import time +from collections.abc import Sequence +from typing import cast + +from pydantic import Field, field_serializer, field_validator, model_validator + +import verifiers.v1 as vf +from verifiers.types import ( + AssistantMessage, + Message, + MessageContent, + Messages, + Tool, + ToolCall, + ToolMessage, + UserMessage, +) +from verifiers.utils.response_utils import parse_response_message + +_BFCL_PATCHED = False +BFCLRawMessage = str | vf.JsonData +BFCLRawTurn = str | vf.JsonData | Sequence[BFCLRawMessage] | None + + +class BFCLTask(vf.Task, frozen=True): + category: str + question: list[list[vf.JsonData]] + function: vf.JsonValue + function_with_hints: vf.JsonValue | None = None + ground_truth: vf.JsonValue | None = None + initial_config: vf.JsonData = Field(default_factory=dict) + involved_classes: vf.JsonValue | None = None + missed_function: vf.JsonData = Field(default_factory=dict) + missed_function_with_hints: vf.JsonData = Field(default_factory=dict) + max_steps_per_turn: int | None = None + + @field_validator("function", "function_with_hints", "ground_truth", mode="before") + @classmethod + def parse_json_blob(cls, value: object) -> object: + if isinstance(value, str): + return json.loads(value) + return value + + @field_serializer( + "function", "function_with_hints", "ground_truth", when_used="json" + ) + def serialize_json_blob(self, value: vf.JsonValue | None) -> str | None: + if value is None: + return None + return json.dumps(value) + + +class BFCLTasksetConfig(vf.TasksetConfig): + id: str = "bfcl-v3-v1" + test_category: str = "simple_python" + test_categories: list[str] | None = None + examples_per_category: int = -1 + + @model_validator(mode="after") + def validate_category_routing(self) -> "BFCLTasksetConfig": + if self.test_categories is not None: + raise ValueError( + "BFCL v3 accepts one test_category per taskset. Configure separate " + "evals for multiple categories." + ) + return self + + +class BFCLHarnessConfig(vf.HarnessConfig): + max_turns: int = 1 + + +def modded_convert_func_name(function_name: str, model_name: str) -> str: + _ = model_name + return re.sub(r"\.", "_", function_name) + + +def patch_bfcl_eval() -> None: + global _BFCL_PATCHED + if _BFCL_PATCHED: + return + import bfcl_eval.constants.category_mapping as category_mapping + import bfcl_eval.eval_checker.ast_eval.ast_checker as ast_checker_module + + agentic_categories = category_mapping.AGENTIC_CATEGORY.copy() + category_mapping.AGENTIC_CATEGORY.clear() + for category in agentic_categories: + if category in category_mapping.ALL_SCORING_CATEGORIES: + category_mapping.ALL_SCORING_CATEGORIES.remove(category) + if category in category_mapping.ALL_CATEGORIES: + category_mapping.ALL_CATEGORIES.remove(category) + category_mapping.TEST_COLLECTION_MAPPING.pop("memory", None) + category_mapping.TEST_COLLECTION_MAPPING.pop("web_search", None) + category_mapping.TEST_COLLECTION_MAPPING.pop("agentic", None) + + non_scoring_categories = category_mapping.NON_SCORING_CATEGORY.copy() + category_mapping.NON_SCORING_CATEGORY.clear() + for category in non_scoring_categories: + if category in category_mapping.ALL_CATEGORIES: + category_mapping.ALL_CATEGORIES.remove(category) + category_mapping.TEST_COLLECTION_MAPPING.pop("format_sensitivity", None) + + setattr(ast_checker_module, "convert_func_name", modded_convert_func_name) + _BFCL_PATCHED = True + + +def bfcl_tool_defs(functions: object) -> list[Tool]: + patch_bfcl_eval() + from bfcl_eval.constants.enums import ModelStyle + from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI + from bfcl_eval.model_handler.utils import convert_to_tool + + oai_tools = convert_to_tool( + functions, GORILLA_TO_OPENAPI, ModelStyle.OPENAI_COMPLETIONS + ) + tool_defs: list[Tool] = [] + for tool in oai_tools: + function = tool["function"] + tool_defs.append( + Tool( + name=str(function["name"]), + description=str(function.get("description") or ""), + parameters=dict(cast(vf.JsonData, function["parameters"])), + strict=False, + ) + ) + return tool_defs + + +def bfcl_functions(task: vf.Task) -> object: + task = cast(BFCLTask, task) + return ( + task.function_with_hints + if task.function_with_hints is not None + else task.function + ) + + +def bfcl_missed_function(task: vf.Task) -> vf.JsonData: + task = cast(BFCLTask, task) + return task.missed_function_with_hints or task.missed_function + + +def build_task_loader(test_category: str, examples_per_category: int = -1): + def factory() -> list[vf.JsonData]: + patch_bfcl_eval() + from bfcl_eval.utils import ( + is_multi_turn, + is_relevance_or_irrelevance, + load_dataset_entry, + load_ground_truth_entry, + ) + + entries = load_dataset_entry( + test_category, include_language_specific_hint=False + ) + entries_with_hints = load_dataset_entry( + test_category, include_language_specific_hint=True + ) + if is_relevance_or_irrelevance(test_category): + ground_truth_entries = [None] * len(entries) + else: + ground_truth_entries = load_ground_truth_entry(test_category) + limit = len(entries) if examples_per_category < 0 else examples_per_category + rows: list[vf.JsonData] = [] + for index, (entry, hinted_entry, ground_truth) in enumerate( + zip(entries, entries_with_hints, ground_truth_entries) + ): + if index >= limit: + break + row = bfcl_row( + test_category, + cast(vf.JsonData, entry), + cast(vf.JsonData, hinted_entry), + cast(vf.JsonData | None, ground_truth), + ) + if is_multi_turn(test_category): + max_steps = maximum_step_limit() + row["max_steps_per_turn"] = max_steps + row["max_turns"] = ( + len(cast(Sequence[BFCLRawTurn], row["question"])) * max_steps + ) + else: + row["max_turns"] = 1 + rows.append(row) + return rows + + return factory + + +def load_tasks(test_category: str = "simple_python", examples_per_category: int = -1): + return build_task_loader(test_category, examples_per_category)() + + +def bfcl_row( + test_category: str, + entry: vf.JsonData, + hinted_entry: vf.JsonData, + ground_truth: vf.JsonData | None, +) -> vf.JsonData: + question = cast(list[BFCLRawTurn], entry["question"]) + first_turn_system_prompt, first_turn_prompt = split_system_prompt( + normalize_turn(question[0]) + ) + row = cast( + vf.JsonData, + { + "task_id": str(entry["id"]), + "category": test_category, + "prompt": first_turn_prompt, + "question": [ + first_turn_prompt, + *[normalize_turn(turn) for turn in question[1:]], + ], + "function": json.dumps(entry["function"]), + "function_with_hints": json.dumps(hinted_entry["function"]), + }, + ) + if first_turn_system_prompt: + row["system_prompt"] = cast(vf.JsonValue, first_turn_system_prompt) + for key in ("initial_config", "involved_classes"): + if key in entry: + row[key] = entry[key] + if "missed_function" in entry: + row["missed_function"] = entry["missed_function"] + if "missed_function" in hinted_entry: + row["missed_function_with_hints"] = hinted_entry["missed_function"] + if ground_truth is not None: + row.update(ground_truth) + row.pop("id", None) + if "ground_truth" in row: + row["ground_truth"] = json.dumps(row["ground_truth"]) + return row + + +def normalize_turn(value: object) -> list[vf.JsonData]: + if value is None: + return [] + if isinstance(value, str): + return [{"role": "user", "content": value}] + if isinstance(value, dict): + return [dict(cast(vf.JsonData, value))] + if isinstance(value, Sequence): + messages: list[vf.JsonData] = [] + for item in value: + if isinstance(item, str): + messages.append({"role": "user", "content": item}) + elif isinstance(item, dict): + messages.append(dict(cast(vf.JsonData, item))) + else: + raise TypeError(f"Unsupported BFCL message item: {type(item).__name__}") + return messages + raise TypeError(f"Unsupported BFCL prompt turn: {type(value).__name__}") + + +def split_system_prompt( + messages: Sequence[vf.JsonData], +) -> tuple[list[vf.JsonData], list[vf.JsonData]]: + system_prompt: list[vf.JsonData] = [] + prompt: list[vf.JsonData] = [] + for message in messages: + target = system_prompt if message.get("role") == "system" else prompt + target.append(dict(message)) + return system_prompt, prompt + + +def maximum_step_limit() -> int: + patch_bfcl_eval() + from bfcl_eval.constants.default_prompts import MAXIMUM_STEP_LIMIT + + return cast(int, MAXIMUM_STEP_LIMIT) + + +def model_name(state: vf.State) -> str: + value = state.metadata.get("model") + return value if isinstance(value, str) and value else "unknown" + + +def assistant_tool_calls(state: vf.State) -> list[ToolCall]: + messages = [message for message in state.completion if message.role == "assistant"] + if not messages: + return [] + return parse_tool_calls(messages[-1]) + + +def transcript_completion_messages(state: vf.State) -> Messages: + if not state.transcript: + return [] + seen = list(state.transcript[0].prompt) + messages: Messages = [] + for index, turn in enumerate(state.transcript): + if index: + prompt_delta = list(turn.prompt[len(seen) :]) + messages.extend(prompt_delta) + seen.extend(prompt_delta) + messages.extend(turn.completion) + seen.extend(turn.completion) + messages.extend(turn.tool_results) + seen.extend(turn.tool_results) + return messages + + +def parse_tool_calls(message: Message | vf.JsonData) -> list[ToolCall]: + if isinstance(message, AssistantMessage): + return list(message.tool_calls or []) + raw_tool_calls: object + if isinstance(message, dict): + raw_tool_calls = message.get("tool_calls") or [] + else: + raw_tool_calls = getattr(message, "tool_calls", []) or [] + if not isinstance(raw_tool_calls, Sequence): + return [] + calls: list[ToolCall] = [] + for raw_call in raw_tool_calls: + if isinstance(raw_call, ToolCall): + calls.append(raw_call) + continue + if not isinstance(raw_call, dict): + continue + raw_call = cast(vf.JsonData, raw_call) + function = raw_call.get("function") + if isinstance(function, dict): + function_map = cast(vf.JsonData, function) + name = str(function_map.get("name") or "") + arguments = function_map.get("arguments") or "{}" + else: + name = str(raw_call.get("name") or "") + arguments = raw_call.get("arguments") or "{}" + if not name: + continue + calls.append( + ToolCall( + id=str(raw_call.get("id") or name), + name=name, + arguments=arguments + if isinstance(arguments, str) + else json.dumps(arguments), + ) + ) + return calls + + +def convert_to_gorilla(tool_calls: list[ToolCall]) -> list[vf.JsonData]: + return [ + {tool_call.name: tool_args(tool_call.arguments)} for tool_call in tool_calls + ] + + +def convert_to_func_calls(tool_calls: list[ToolCall]) -> list[str]: + func_calls: list[str] = [] + for tool_call in tool_calls: + params = tool_args(tool_call.arguments) + args = ",".join(f"{key}={value!r}" for key, value in params.items()) + func_calls.append(f"{tool_call.name}({args})") + return func_calls + + +def json_clone(value: object) -> object: + return json.loads(json.dumps(value)) + + +def bfcl_involved_classes(task: BFCLTask) -> list[str]: + value = json_clone(task.involved_classes) + if not isinstance(value, list) or not all(isinstance(item, str) for item in value): + raise TypeError("BFCL multi-turn tasks require involved_classes.") + return cast(list[str], value) + + +def tool_args(value: str) -> vf.JsonData: + parsed = json.loads(value or "{}") + if not isinstance(parsed, dict): + raise TypeError("BFCL tool arguments must decode to an object.") + return cast(vf.JsonData, parsed) + + +def relevance_reward(task: vf.Task, state: vf.State) -> float: + patch_bfcl_eval() + from bfcl_eval.utils import is_empty_output + + task = cast(BFCLTask, task) + category = task.category + try: + gorilla_tool_calls = convert_to_gorilla(assistant_tool_calls(state)) + contain_func_call = not is_empty_output(gorilla_tool_calls) + except Exception: + contain_func_call = False + if "irrelevance" in category: + return float(not contain_func_call) + return float(contain_func_call) + + +def ast_reward(task: vf.Task, state: vf.State) -> float: + patch_bfcl_eval() + from bfcl_eval.constants.enums import Language + from bfcl_eval.eval_checker.ast_eval.ast_checker import ast_checker + from bfcl_eval.utils import ( + is_function_calling_format_output, + is_java, + is_js, + ) + + task = cast(BFCLTask, task) + category = task.category + try: + gorilla_tool_calls = convert_to_gorilla(assistant_tool_calls(state)) + if not is_function_calling_format_output(gorilla_tool_calls): + return 0.0 + except Exception: + return 0.0 + + if is_java(category): + language = Language.JAVA + elif is_js(category): + language = Language.JAVASCRIPT + else: + language = Language.PYTHON + + checker_result = ast_checker( + task.function, + gorilla_tool_calls, + task.ground_truth, + language, + category, + model_name(state), + ) + return float(bool(checker_result["valid"])) + + +def multi_turn_reward(task: vf.Task, state: vf.State) -> float: + patch_bfcl_eval() + from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_checker import ( + multi_turn_checker, + ) + from bfcl_eval.model_handler.base_handler import is_empty_execute_response + + task = cast(BFCLTask, task) + completion = transcript_completion_messages(state) + raw_ground_truth = task.ground_truth + if not isinstance(raw_ground_truth, Sequence): + return 0.0 + all_ground_truth = cast(list[list[str]], raw_ground_truth) + all_func_calls: list[list[list[str]]] = [[]] + try: + for message in completion: + role = message.role + if role == "user": + all_func_calls.append([]) + elif role == "tool": + continue + elif role == "assistant": + func_calls = convert_to_func_calls(parse_tool_calls(message)) + if is_empty_execute_response(func_calls): + continue + all_func_calls[-1].append(func_calls) + elif role == "system": + continue + else: + return 0.0 + except Exception: + return 0.0 + + if len(all_func_calls) != len(all_ground_truth): + return 0.0 + + result = multi_turn_checker( + all_func_calls, + all_ground_truth, + { + "initial_config": task.initial_config, + "involved_classes": task.involved_classes, + "id": task.task_id, + }, + task.task_id.rsplit("_", 1)[0], + model_name(state), + ) + return float(bool(result["valid"])) + + +class BFCLTaskset(vf.Taskset[BFCLTasksetConfig]): + task_type = BFCLTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + _ = split + return load_tasks(self.config.test_category, self.config.examples_per_category) + + @vf.reward(weight=1.0) + async def bfcl_reward(self, task: vf.Task, state: vf.State) -> float: + patch_bfcl_eval() + from bfcl_eval.utils import is_multi_turn, is_relevance_or_irrelevance + + task = cast(BFCLTask, task) + category = task.category + if is_relevance_or_irrelevance(category): + return relevance_reward(task, state) + if is_multi_turn(category): + return multi_turn_reward(task, state) + return ast_reward(task, state) + + +class BFCLHarness(vf.Harness[BFCLHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = BFCLTask.model_validate(context.task.model_dump()) + state = context.state + patch_bfcl_eval() + from bfcl_eval.utils import is_multi_turn + + state.metadata["model"] = context.model + if is_multi_turn(task.category): + await self.run_multi_turn(context, task, state) + return + prompt = self.initial_messages(task) + start = time.time() + response = await context.model_client.get_response( + prompt=prompt, + model=context.model, + sampling_args=self.sampling_args(task, context.sampling_args), + tools=bfcl_tool_defs(bfcl_functions(task)), + state=state, + ) + end = time.time() + turn = vf.Turn( + prompt=prompt, + completion=await parse_response_message(response), + tool_calls=list(response.message.tool_calls or []), + response_id=response.id, + model=response.model, + created=response.created, + finish_reason=response.message.finish_reason, + usage=vf.TurnUsage.from_usage(response.usage), + tokens=vf.TurnTokens.from_response( + response.message.tokens, + is_truncated=bool(response.message.is_truncated), + ), + is_truncated=bool(response.message.is_truncated), + timing=vf.TimeSpan(start=start, end=end), + ) + state.transcript.append(turn) + if turn.is_truncated: + state.is_truncated = True + state.stop("assistant_completed") + + async def run_multi_turn( + self, + context: vf.Context, + task: BFCLTask, + state: vf.State, + ) -> None: + from bfcl_eval.constants.default_prompts import ( + DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC, + ) + from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import ( + execute_multi_turn_func_call, + ) + from bfcl_eval.model_handler.base_handler import is_empty_execute_response + + messages = list(self.initial_messages(task)) + next_prompts = list(task.question)[1:] + holdout_function = bfcl_missed_function(task) + tool_defs = bfcl_tool_defs(bfcl_functions(task)) + initial_config = cast(vf.JsonData, json_clone(task.initial_config or {})) + involved_classes = bfcl_involved_classes(task) + simulator_model = ( + model_name(state).replace("/", "_").replace("-", "_").replace(".", "_") + ) + long_context = "long_context" in task.category or "composite" in task.category + max_steps_per_turn = int(task.max_steps_per_turn or maximum_step_limit()) + max_turns = self.max_turns(task) + model_turns = 0 + turn_idx = 0 + steps_per_turn = 0 + execute_multi_turn_func_call( + [], + initial_config, + involved_classes, + simulator_model, + task.task_id, + long_context=long_context, + ) + while max_turns <= 0 or model_turns < max_turns: + if await self.is_completed(context): + return + start = time.time() + response = await context.model_client.get_response( + prompt=messages, + model=context.model, + sampling_args=self.sampling_args(task, context.sampling_args), + tools=tool_defs, + state=state, + ) + end = time.time() + turn = vf.Turn( + prompt=list(messages), + completion=await parse_response_message(response), + tool_calls=list(response.message.tool_calls or []), + response_id=response.id, + model=response.model, + created=response.created, + finish_reason=response.message.finish_reason, + usage=vf.TurnUsage.from_usage(response.usage), + tokens=vf.TurnTokens.from_response( + response.message.tokens, + is_truncated=bool(response.message.is_truncated), + ), + is_truncated=bool(response.message.is_truncated), + timing=vf.TimeSpan(start=start, end=end), + ) + state.transcript.append(turn) + model_turns += 1 + if turn.is_truncated: + state.is_truncated = True + messages.extend(turn.completion) + tool_calls = list(turn.tool_calls) + try: + func_calls = convert_to_func_calls(tool_calls) + if is_empty_execute_response(func_calls): + func_calls = [] + except Exception: + func_calls = [] + if func_calls: + execution_results, _ = execute_multi_turn_func_call( + func_call_list=func_calls, + initial_config=initial_config, + involved_classes=involved_classes, + model_name=simulator_model, + test_entry_id=task.task_id, + long_context=long_context, + ) + tool_messages = [ + ToolMessage( + tool_call_id=tool_call.id, + content=cast(MessageContent, execution_result), + ) + for execution_result, tool_call in zip( + execution_results, tool_calls + ) + ] + turn.tool_results = tool_messages + messages.extend(tool_messages) + steps_per_turn += 1 + if steps_per_turn >= max_steps_per_turn: + state.stop("max_steps_per_turn_reached") + return + continue + + steps_per_turn = 0 + turn_idx += 1 + if not next_prompts: + state.stop("no_next_prompt_and_no_tool_calls") + return + next_prompt = normalize_turn(next_prompts.pop(0)) + if str(turn_idx) in holdout_function: + tool_defs.extend(bfcl_tool_defs(holdout_function[str(turn_idx)])) + if next_prompt: + raise ValueError( + "BFCL holdout turns must not include user messages." + ) + messages.append( + UserMessage(content=DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC) + ) + else: + for message in next_prompt: + role = message.get("role") + if role == "assistant": + messages.append(AssistantMessage.model_validate(message)) + elif role == "user": + messages.append(UserMessage.model_validate(message)) + elif role == "tool": + messages.append(ToolMessage.model_validate(message)) + elif role == "system": + raise ValueError( + "BFCL turn prompts must not include system messages." + ) + else: + raise ValueError( + f"Unsupported BFCL prompt message role: {role!r}." + ) + state.stop("max_turns") + + +def load_taskset(config: BFCLTasksetConfig) -> BFCLTaskset: + return BFCLTaskset(config=config) + + +def load_harness(config: BFCLHarnessConfig) -> BFCLHarness: + return BFCLHarness(config=config) diff --git a/environments/bfcl_v3/pyproject.toml b/environments/bfcl_v3_v1/pyproject.toml similarity index 61% rename from environments/bfcl_v3/pyproject.toml rename to environments/bfcl_v3_v1/pyproject.toml index a3c851caeb..a80c922fe2 100644 --- a/environments/bfcl_v3/pyproject.toml +++ b/environments/bfcl_v3_v1/pyproject.toml @@ -1,12 +1,12 @@ [project] -name = "bfcl-v3" +name = "bfcl-v3-v1" description = "BFCL v3 evaluation environment on the v1 Taskset/Harness runtime" tags = ["tool-use", "eval", "v1"] version = "0.1.0" requires-python = ">=3.11" dependencies = [ "verifiers>=0.1.13.dev0", - "bfcl-eval @ git+https://github.com/mikasenghaas/gorilla.git@898763a#subdirectory=berkeley-function-call-leaderboard", + "bfcl-eval>=2026.3.23", "soundfile>=0.13.0", ] @@ -14,11 +14,8 @@ dependencies = [ requires = ["hatchling"] build-backend = "hatchling.build" -[tool.hatch.metadata] -allow-direct-references = true - [tool.hatch.build] -include = ["bfcl_v3.py", "README.md", "pyproject.toml"] +include = ["bfcl_v3_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/dspy_flights/dspy_flights.py b/environments/dspy_flights/dspy_flights.py deleted file mode 100644 index 3ec38deb36..0000000000 --- a/environments/dspy_flights/dspy_flights.py +++ /dev/null @@ -1,455 +0,0 @@ -import asyncio -import functools -import random -import string -from collections.abc import Mapping -from typing import cast - -from pydantic import BaseModel - -import verifiers as vf - -PROGRAM_SANDBOX = { - "image": "python:3.11-slim", - "network_access": True, - "timeout_minutes": 60, - "command_timeout": 900, - "install_timeout": 900, -} - - -class DSPyFlightsHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig( - fn="run_dspy_flight_program", - sandbox=True, - ) - sandbox: vf.SandboxConfig = vf.SandboxConfig(**PROGRAM_SANDBOX) - - -class Date(BaseModel): - # Somehow LLM is bad at specifying `datetime.datetime`, so - # we define a custom class to represent the date. - year: int - month: int - day: int - hour: int - - -class UserProfile(BaseModel): - user_id: str - name: str - email: str - - -class Flight(BaseModel): - flight_id: str - date_time: Date - origin: str - destination: str - duration: float - price: float - - -class Itinerary(BaseModel): - confirmation_number: str - user_profile: UserProfile - flight: Flight - - -class Ticket(BaseModel): - user_request: str - user_profile: UserProfile - - -def user_database() -> dict[str, UserProfile]: - return { - "Adam": UserProfile(user_id="1", name="Adam", email="adam@gmail.com"), - "Bob": UserProfile(user_id="2", name="Bob", email="bob@gmail.com"), - "Chelsie": UserProfile(user_id="3", name="Chelsie", email="chelsie@gmail.com"), - "David": UserProfile(user_id="4", name="David", email="david@gmail.com"), - } - - -def flight_database() -> dict[str, Flight]: - return { - "DA123": Flight( - flight_id="DA123", - origin="SFO", - destination="JFK", - date_time=Date(year=2025, month=9, day=1, hour=1), - duration=3, - price=200, - ), - "DA125": Flight( - flight_id="DA125", - origin="SFO", - destination="JFK", - date_time=Date(year=2025, month=9, day=1, hour=7), - duration=9, - price=500, - ), - "DA456": Flight( - flight_id="DA456", - origin="SFO", - destination="SNA", - date_time=Date(year=2025, month=10, day=1, hour=1), - duration=2, - price=100, - ), - "DA460": Flight( - flight_id="DA460", - origin="SFO", - destination="SNA", - date_time=Date(year=2025, month=10, day=1, hour=9), - duration=2, - price=120, - ), - } - - -@vf.reward(weight=1.0) -async def expected_database_change(task, state) -> float: - expected = task["expected"] - if expected["kind"] == "book": - itineraries = state.get("itinerary_database", {}) - return float( - len(itineraries) == 1 - and any( - item["user_profile"]["name"] == expected["user"] - and item["flight"]["flight_id"] == expected["flight_id"] - for item in itineraries.values() - ) - ) - if expected["kind"] == "cancel": - return float( - expected["confirmation_number"] not in state.get("itinerary_database", {}) - ) - if expected["kind"] == "ticket": - tickets = state.get("ticket_database", {}) - return float( - len(tickets) == 1 - and any( - item["user_profile"]["name"] == expected["user"] - and expected["contains"].lower() in item["user_request"].lower() - for item in tickets.values() - ) - ) - raise ValueError(f"Unknown expected kind: {expected['kind']}") - - -@vf.metric -async def dspy_calls(task, state) -> float: - return float(len(state.get("trajectory", []))) - - -def load_tasks(split: vf.TaskSplit = "train"): - _ = split - - def record( - example_id: int, - user_request: str, - expected: vf.ConfigData, - initial_itineraries: dict[str, vf.ConfigData] | None = None, - ) -> vf.ConfigData: - task: vf.ConfigData = { - "example_id": example_id, - "user_request": user_request, - "prompt": [{"role": "user", "content": user_request}], - "expected": expected, - } - if initial_itineraries is not None: - task["initial_itineraries"] = initial_itineraries - return task - - return [ - record( - 0, - ( - "please help me book a flight from SFO to JFK on 09/01/2025, " - "my name is Adam" - ), - {"kind": "book", "user": "Adam", "flight_id": "DA123"}, - ), - record( - 1, - ( - "please help me book a flight from SFO to SNA on 10/01/2025, " - "my name is Bob" - ), - {"kind": "book", "user": "Bob", "flight_id": "DA456"}, - ), - record( - 2, - ( - "please cancel itinerary CH123 for Chelsie; she no longer wants " - "to travel" - ), - {"kind": "cancel", "confirmation_number": "CH123"}, - {"CH123": itinerary("CH123", "Chelsie", "DA125").model_dump()}, - ), - record( - 3, - ( - "my name is David and I need wheelchair assistance added to my " - "reservation" - ), - { - "kind": "ticket", - "user": "David", - "contains": "wheelchair assistance", - }, - ), - record( - 4, - ("my name is Adam and I need a vegetarian meal noted for my upcoming trip"), - { - "kind": "ticket", - "user": "Adam", - "contains": "vegetarian meal", - }, - ), - record( - 5, - "please cancel itinerary BO456 for Bob because his plans changed", - {"kind": "cancel", "confirmation_number": "BO456"}, - {"BO456": itinerary("BO456", "Bob", "DA456").model_dump()}, - ), - record( - 6, - ( - "please help me book a flight from SFO to JFK on 09/01/2025, " - "my name is Chelsie" - ), - {"kind": "book", "user": "Chelsie", "flight_id": "DA123"}, - ), - record( - 7, - ( - "please help me book a flight from SFO to SNA on 10/01/2025, " - "my name is David" - ), - {"kind": "book", "user": "David", "flight_id": "DA456"}, - ), - record( - 8, - "cancel confirmation AD460 for Adam; he will rebook later", - {"kind": "cancel", "confirmation_number": "AD460"}, - {"AD460": itinerary("AD460", "Adam", "DA460").model_dump()}, - ), - record( - 9, - "my name is Chelsie and I need to travel with a service animal", - { - "kind": "ticket", - "user": "Chelsie", - "contains": "service animal", - }, - ), - ] - - -def itinerary(confirmation_number: str, user_name: str, flight_id: str) -> Itinerary: - users = user_database() - flights = flight_database() - return Itinerary( - confirmation_number=confirmation_number, - user_profile=users[user_name], - flight=flights[flight_id], - ) - - -def build_airline_tools( - task, -) -> tuple[list[vf.Handler], dict[str, dict[str, BaseModel]]]: - users = user_database() - flights = flight_database() - itineraries = { - key: Itinerary.model_validate(value) - for key, value in (task.get("initial_itineraries") or {}).items() - } - tickets: dict[str, Ticket] = {} - - def fetch_flight_info(date: Date, origin: str, destination: str): - """Fetch flight information from origin to destination on the given date""" - date = Date.model_validate(date) - matching_flights = [] - - for flight in flights.values(): - if ( - flight.date_time.year == date.year - and flight.date_time.month == date.month - and flight.date_time.day == date.day - and flight.origin == origin - and flight.destination == destination - ): - matching_flights.append(flight) - if len(matching_flights) == 0: - raise ValueError("No matching flight found!") - return matching_flights - - def fetch_itinerary(confirmation_number: str): - """Fetch a booked itinerary information from database""" - return itineraries.get(confirmation_number) - - def pick_flight(flights: list[Flight]): - """Pick up the best flight that matches users' request. we pick the shortest, and cheaper one on ties.""" - sorted_flights = sorted( - flights, - key=lambda x: ( - x.get("duration") if isinstance(x, dict) else x.duration, - x.get("price") if isinstance(x, dict) else x.price, - ), - ) - return sorted_flights[0] - - def _generate_id(length=8): - chars = string.ascii_lowercase + string.digits - return "".join(random.choices(chars, k=length)) - - def book_flight(flight: Flight, user_profile: UserProfile): - """Book a flight on behalf of the user.""" - flight = Flight.model_validate(flight) - user_profile = UserProfile.model_validate(user_profile) - confirmation_number = _generate_id() - while confirmation_number in itineraries: - confirmation_number = _generate_id() - itineraries[confirmation_number] = Itinerary( - confirmation_number=confirmation_number, - user_profile=user_profile, - flight=flight, - ) - return confirmation_number, itineraries[confirmation_number] - - def cancel_itinerary(confirmation_number: str, user_profile: UserProfile): - """Cancel an itinerary on behalf of the user.""" - _ = UserProfile.model_validate(user_profile) - if confirmation_number in itineraries: - del itineraries[confirmation_number] - return - raise ValueError( - "Cannot find the itinerary, please check your confirmation number." - ) - - def get_user_info(name: str): - """Fetch the user profile from database with given name.""" - return users.get(name) - - def file_ticket(user_request: str, user_profile: UserProfile): - """File a customer support ticket if this is something the agent cannot handle.""" - user_profile = UserProfile.model_validate(user_profile) - ticket_id = _generate_id(length=6) - tickets[ticket_id] = Ticket( - user_request=user_request, - user_profile=user_profile, - ) - return ticket_id - - tools: list[vf.Handler] = [ - async_tool(fetch_flight_info), - async_tool(fetch_itinerary), - async_tool(pick_flight), - async_tool(book_flight), - async_tool(cancel_itinerary), - async_tool(get_user_info), - async_tool(file_ticket), - ] - databases: dict[str, dict[str, BaseModel]] = { - "itinerary_database": itineraries, - "ticket_database": tickets, - } - return tools, databases - - -def async_tool(fn: vf.Handler) -> vf.Handler: - @functools.wraps(fn) - async def wrapped(*args: object, **kwargs: object) -> object: - return await asyncio.to_thread(fn, *args, **kwargs) - - return wrapped - - -def dump_database(database: dict[str, BaseModel]) -> dict[str, vf.ConfigData]: - return {key: value.model_dump() for key, value in database.items()} - - -async def run_dspy_flight_program(task, state): - import dspy - from openai import OpenAI - - class DSPyAirlineCustomerService(dspy.Signature): - """You are an airline customer service agent that helps user book and manage flights. - - You are given a list of tools to handle user request, and you should decide the right tool to use in order to - fulfill users' request. - """ - - user_request: str = dspy.InputField() - process_result: str = dspy.OutputField( - desc=( - "Message that summarizes the process result, and the information users need, e.g., the " - "confirmation_number if a new flight is booked." - ) - ) - - tools, databases = build_airline_tools(task) - endpoint_config = state.get_endpoint_config(api="chat") - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - endpoint_api_key = endpoint_client.api_key - endpoint_client.close() - lm = dspy.LM( - f"openai/{endpoint_config.model}", - api_base=endpoint_config.base_url, - api_key=endpoint_api_key, - cache=False, - ) - agent = dspy.ReAct(DSPyAirlineCustomerService, tools=tools, max_iters=8) - with dspy.context(lm=lm): - result = await agent.acall(user_request=task["user_request"]) - - state["process_result"] = str(result.process_result) - state["reasoning"] = str(getattr(result, "reasoning", "")) - state["dspy_trajectory"] = stringify_nested(getattr(result, "trajectory", {})) - state["itinerary_database"] = dump_database(databases["itinerary_database"]) - state["ticket_database"] = dump_database(databases["ticket_database"]) - state["completion"] = [ - {"role": "assistant", "content": state["process_result"]}, - ] - return state - - -def stringify_nested(value: object) -> object: - if isinstance(value, BaseModel): - return stringify_nested(value.model_dump()) - if isinstance(value, Mapping): - return {str(key): stringify_nested(item) for key, item in value.items()} - if isinstance(value, list | tuple): - return [stringify_nested(item) for item in value] - if isinstance(value, str | int | float | bool) or value is None: - return value - return str(value) - - -class DSPyFlightsTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["expected_database_change"] - metrics: list[str] = ["dspy_calls"] - - -class DSPyFlightsTaskset(vf.Taskset[DSPyFlightsTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -class DSPyFlightsHarness(vf.Harness[DSPyFlightsHarnessConfig]): - pass - - -class DSPyFlightsEnvConfig(vf.EnvConfig): - taskset: DSPyFlightsTasksetConfig = DSPyFlightsTasksetConfig() - harness: DSPyFlightsHarnessConfig = DSPyFlightsHarnessConfig() - - -def load_environment(config: DSPyFlightsEnvConfig) -> vf.Env: - return vf.Env( - taskset=DSPyFlightsTaskset(config=config.taskset), - harness=DSPyFlightsHarness(config=config.harness), - ) diff --git a/environments/dspy_flights/README.md b/environments/dspy_flights_v1/README.md similarity index 94% rename from environments/dspy_flights/README.md rename to environments/dspy_flights_v1/README.md index 15da9832da..b4a97247fb 100644 --- a/environments/dspy_flights/README.md +++ b/environments/dspy_flights_v1/README.md @@ -1,4 +1,4 @@ -# dspy-flights +# dspy-flights-v1 Minimal v1 environment for a third-party DSPy flight-support program. diff --git a/environments/dspy_flights_v1/dspy_flights_v1/__init__.py b/environments/dspy_flights_v1/dspy_flights_v1/__init__.py new file mode 100644 index 0000000000..f11d8edce9 --- /dev/null +++ b/environments/dspy_flights_v1/dspy_flights_v1/__init__.py @@ -0,0 +1 @@ +"""dspy-flights-v1 environment package.""" diff --git a/environments/dspy_flights_v1/dspy_flights_v1/harness.py b/environments/dspy_flights_v1/dspy_flights_v1/harness.py new file mode 100644 index 0000000000..4124afe3c8 --- /dev/null +++ b/environments/dspy_flights_v1/dspy_flights_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import DSPyFlightsHarness as DSPyFlightsHarness +from .taskset import DSPyFlightsHarnessConfig as DSPyFlightsHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/dspy_flights_v1/dspy_flights_v1/taskset.py b/environments/dspy_flights_v1/dspy_flights_v1/taskset.py new file mode 100644 index 0000000000..ff7dd5d3d1 --- /dev/null +++ b/environments/dspy_flights_v1/dspy_flights_v1/taskset.py @@ -0,0 +1,499 @@ +import asyncio +import functools +import random +import re +import string +from collections.abc import Callable, Mapping +from typing import Literal, TypeAlias, cast + +from pydantic import BaseModel, Field + +import verifiers.v1 as vf + +DSPyToolResult: TypeAlias = ( + str + | None + | BaseModel + | list[BaseModel] + | tuple[str, BaseModel] + | tuple[str, dict[str, BaseModel]] +) + + +class DSPyFlightsHarnessConfig(vf.HarnessConfig): + max_iters: int = 8 + + +class Date(BaseModel): + year: int + month: int + day: int + hour: int + + +class UserProfile(BaseModel): + user_id: str + name: str + email: str + + +class Flight(BaseModel): + flight_id: str + date_time: Date + origin: str + destination: str + duration: float + price: float + + +class Itinerary(BaseModel): + confirmation_number: str + user_profile: UserProfile + flight: Flight + + +class Ticket(BaseModel): + user_request: str + user_profile: UserProfile + + +class FlightRunResult(vf.Config): + process_result: str + reasoning: str = "" + dspy_calls: int = 0 + dspy_trajectory: vf.JsonValue | None = None + itinerary_database: dict[str, vf.JsonData] + ticket_database: dict[str, vf.JsonData] + + +class ExpectedFlightChange(BaseModel, extra="forbid"): + kind: Literal["book", "cancel", "ticket"] + user: str | None = None + flight_id: str | None = None + confirmation_number: str | None = None + contains: str | None = None + + +class DSPyFlightsTask(vf.Task): + user_request: str + expected: ExpectedFlightChange + initial_itineraries: dict[str, Itinerary | None] = Field(default_factory=dict) + + +def user_database() -> dict[str, UserProfile]: + return { + "Adam": UserProfile(user_id="1", name="Adam", email="adam@gmail.com"), + "Bob": UserProfile(user_id="2", name="Bob", email="bob@gmail.com"), + "Chelsie": UserProfile(user_id="3", name="Chelsie", email="chelsie@gmail.com"), + "David": UserProfile(user_id="4", name="David", email="david@gmail.com"), + } + + +def flight_database() -> dict[str, Flight]: + return { + "DA123": Flight( + flight_id="DA123", + origin="SFO", + destination="JFK", + date_time=Date(year=2025, month=9, day=1, hour=1), + duration=3, + price=200, + ), + "DA125": Flight( + flight_id="DA125", + origin="SFO", + destination="JFK", + date_time=Date(year=2025, month=9, day=1, hour=7), + duration=9, + price=500, + ), + "DA456": Flight( + flight_id="DA456", + origin="SFO", + destination="SNA", + date_time=Date(year=2025, month=10, day=1, hour=1), + duration=2, + price=100, + ), + "DA460": Flight( + flight_id="DA460", + origin="SFO", + destination="SNA", + date_time=Date(year=2025, month=10, day=1, hour=9), + duration=2, + price=120, + ), + } + + +def load_tasks(split: vf.TaskSplit = "train") -> list[vf.JsonData]: + _ = split + + def record( + example_id: int, + user_request: str, + expected: vf.JsonData, + initial_itineraries: dict[str, vf.JsonData] | None = None, + ) -> vf.JsonData: + task: vf.JsonData = { + "example_id": example_id, + "user_request": user_request, + "prompt": [{"role": "user", "content": user_request}], + "expected": expected, + } + if initial_itineraries is not None: + task["initial_itineraries"] = initial_itineraries + return task + + return [ + record( + 0, + ( + "please help me book a flight from SFO to JFK on 09/01/2025, " + "my name is Adam" + ), + {"kind": "book", "user": "Adam", "flight_id": "DA123"}, + ), + record( + 1, + ( + "please help me book a flight from SFO to SNA on 10/01/2025, " + "my name is Bob" + ), + {"kind": "book", "user": "Bob", "flight_id": "DA456"}, + ), + record( + 2, + ( + "please cancel itinerary CH123 for Chelsie; she no longer wants " + "to travel" + ), + {"kind": "cancel", "confirmation_number": "CH123"}, + {"CH123": itinerary("CH123", "Chelsie", "DA125").model_dump()}, + ), + record( + 3, + ( + "my name is David and I need wheelchair assistance added to my " + "reservation" + ), + {"kind": "ticket", "user": "David", "contains": "wheelchair assistance"}, + ), + record( + 4, + "my name is Adam and I need a vegetarian meal noted for my upcoming trip", + {"kind": "ticket", "user": "Adam", "contains": "vegetarian meal"}, + ), + record( + 5, + "please cancel itinerary BO456 for Bob because his plans changed", + {"kind": "cancel", "confirmation_number": "BO456"}, + {"BO456": itinerary("BO456", "Bob", "DA456").model_dump()}, + ), + record( + 6, + ( + "please help me book a flight from SFO to JFK on 09/01/2025, " + "my name is Chelsie" + ), + {"kind": "book", "user": "Chelsie", "flight_id": "DA123"}, + ), + record( + 7, + ( + "please help me book a flight from SFO to SNA on 10/01/2025, " + "my name is David" + ), + {"kind": "book", "user": "David", "flight_id": "DA456"}, + ), + record( + 8, + "cancel confirmation AD460 for Adam; he will rebook later", + {"kind": "cancel", "confirmation_number": "AD460"}, + {"AD460": itinerary("AD460", "Adam", "DA460").model_dump()}, + ), + record( + 9, + "my name is Chelsie and I need to travel with a service animal", + {"kind": "ticket", "user": "Chelsie", "contains": "service animal"}, + ), + ] + + +def itinerary(confirmation_number: str, user_name: str, flight_id: str) -> Itinerary: + users = user_database() + flights = flight_database() + return Itinerary( + confirmation_number=confirmation_number, + user_profile=users[user_name], + flight=flights[flight_id], + ) + + +def build_airline_tools( + task: DSPyFlightsTask, +) -> tuple[list[Callable[..., DSPyToolResult]], dict[str, dict[str, BaseModel]]]: + users = user_database() + flights = flight_database() + itineraries = { + confirmation_number: itinerary + for confirmation_number, itinerary in task.initial_itineraries.items() + if itinerary is not None + } + tickets: dict[str, Ticket] = {} + + def fetch_flight_info(date: Date, origin: str, destination: str) -> list[Flight]: + date = Date.model_validate(date) + matching_flights = [ + flight + for flight in flights.values() + if flight.date_time.year == date.year + and flight.date_time.month == date.month + and flight.date_time.day == date.day + and flight.origin == origin + and flight.destination == destination + ] + if not matching_flights: + raise ValueError("No matching flight found.") + return matching_flights + + def fetch_itinerary(confirmation_number: str) -> Itinerary | None: + return itineraries.get(confirmation_number) + + def pick_flight(flights: list[Flight]) -> Flight: + return sorted(flights, key=lambda flight: (flight.duration, flight.price))[0] + + def generate_id(length: int = 8) -> str: + chars = string.ascii_lowercase + string.digits + return "".join(random.choices(chars, k=length)) + + def book_flight(flight: Flight, user_profile: UserProfile) -> tuple[str, Itinerary]: + flight = Flight.model_validate(flight) + user_profile = UserProfile.model_validate(user_profile) + confirmation_number = generate_id() + while confirmation_number in itineraries: + confirmation_number = generate_id() + itineraries[confirmation_number] = Itinerary( + confirmation_number=confirmation_number, + user_profile=user_profile, + flight=flight, + ) + return confirmation_number, itineraries[confirmation_number] + + def cancel_itinerary(confirmation_number: str, user_profile: UserProfile) -> None: + UserProfile.model_validate(user_profile) + if confirmation_number in itineraries: + del itineraries[confirmation_number] + return + raise ValueError( + "Cannot find the itinerary, please check your confirmation number." + ) + + def get_user_info(name: str) -> UserProfile | None: + return users.get(name) + + def file_ticket(user_request: str, user_profile: UserProfile) -> str: + user_profile = UserProfile.model_validate(user_profile) + ticket_id = generate_id(length=6) + tickets[ticket_id] = Ticket( + user_request=user_request, + user_profile=user_profile, + ) + return ticket_id + + tools: list[Callable[..., DSPyToolResult]] = [ + async_tool(fetch_flight_info), + async_tool(fetch_itinerary), + async_tool(pick_flight), + async_tool(book_flight), + async_tool(cancel_itinerary), + async_tool(get_user_info), + async_tool(file_ticket), + ] + return tools, {"itinerary_database": itineraries, "ticket_database": tickets} + + +def async_tool( + fn: Callable[..., DSPyToolResult], +) -> Callable[..., DSPyToolResult]: + @functools.wraps(fn) + async def wrapped(*args: object, **kwargs: object) -> DSPyToolResult: + return await asyncio.to_thread(fn, *args, **kwargs) + + return wrapped + + +def dump_database(database: dict[str, BaseModel]) -> dict[str, vf.JsonData]: + return { + key: cast(vf.JsonData, value.model_dump(mode="json")) + for key, value in database.items() + } + + +def jsonable_dspy(value: object) -> vf.JsonValue: + if isinstance(value, BaseModel): + return jsonable_dspy(value.model_dump(mode="json")) + if isinstance(value, Mapping): + return {str(key): jsonable_dspy(item) for key, item in value.items()} + if isinstance(value, list | tuple): + return [jsonable_dspy(item) for item in value] + if isinstance(value, str | int | float | bool) or value is None: + return value + return str(value) + + +def dspy_iteration_count(result: object) -> int: + trajectory = getattr(result, "trajectory", None) + if isinstance(trajectory, Mapping): + indices = { + match.group(1) + for key in trajectory + if (match := re.search(r"_(\d+)$", str(key))) + } + if indices: + return len(indices) + return len(trajectory) + if isinstance(trajectory, list | tuple): + return len(trajectory) + return 0 + + +async def run_dspy_flight_agent( + *, + task: DSPyFlightsTask, + base_url: str, + api_key: str, + model: str, + max_iters: int, +) -> FlightRunResult: + import dspy + + class DSPyAirlineCustomerService(dspy.Signature): + """Airline customer service agent for booking and managing flights.""" + + user_request: str = dspy.InputField() + process_result: str = dspy.OutputField( + desc=( + "Message that summarizes the process result and any information " + "the user needs, such as a confirmation number." + ) + ) + + tools, databases = build_airline_tools(task) + lm = dspy.LM( + f"openai/{model}", + api_base=base_url, + api_key=api_key, + cache=False, + ) + agent = dspy.ReAct(DSPyAirlineCustomerService, tools=tools, max_iters=max_iters) + with dspy.context(lm=lm): + result = await agent.acall(user_request=task.user_request) + + return FlightRunResult( + process_result=str(result.process_result), + reasoning=str(getattr(result, "reasoning", "")), + dspy_calls=dspy_iteration_count(result), + dspy_trajectory=jsonable_dspy(getattr(result, "trajectory", None)), + itinerary_database=dump_database(databases["itinerary_database"]), + ticket_database=dump_database(databases["ticket_database"]), + ) + + +class DSPyFlightsTasksetConfig(vf.TasksetConfig): + id: str = "dspy-flights" + + +class DSPyFlightsTaskset(vf.Taskset[DSPyFlightsTasksetConfig]): + task_type = DSPyFlightsTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return load_tasks(split) + + @vf.reward + async def expected_database_change( + self, task: DSPyFlightsTask, state: vf.State + ) -> float: + expected = task.expected + itineraries = state.artifacts.get("itinerary_database") + tickets = state.artifacts.get("ticket_database") + itinerary_map = itineraries if isinstance(itineraries, Mapping) else {} + ticket_map = tickets if isinstance(tickets, Mapping) else {} + + if expected.kind == "book": + return float( + len(itinerary_map) == 1 + and any( + isinstance(item, Mapping) + and isinstance(item.get("user_profile"), Mapping) + and isinstance(item.get("flight"), Mapping) + and item["user_profile"].get("name") == expected.user + and item["flight"].get("flight_id") == expected.flight_id + for item in itinerary_map.values() + ) + ) + if expected.kind == "cancel": + return float(expected.confirmation_number not in itinerary_map) + if expected.kind == "ticket": + contains = str(expected.contains or "").lower() + return float( + len(ticket_map) == 1 + and any( + isinstance(item, Mapping) + and isinstance(item.get("user_profile"), Mapping) + and item["user_profile"].get("name") == expected.user + and contains in str(item.get("user_request", "")).lower() + for item in ticket_map.values() + ) + ) + raise ValueError(f"Unknown expected kind: {expected.kind!r}.") + + @vf.metric + async def dspy_calls(self, state: vf.State) -> float: + value = state.artifacts.get("dspy_calls") + return float(value) if isinstance(value, int | float) else 0.0 + + +class DSPyFlightsHarness(vf.Harness[DSPyFlightsHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = DSPyFlightsTask.model_validate(context.task.model_dump()) + state = context.state + runtime = context.runtime + if runtime is None: + raise RuntimeError("DSPyFlightsHarness requires a runtime.") + prompt = self.initial_messages(task) + + async def stop_check() -> str | None: + if await self.is_completed(context): + return state.stop_condition or "stop" + return None + + async with vf.InterceptionServer( + context, + task, + state, + protocols=self.protocols, + stop_check=stop_check, + ) as endpoint: + endpoint_url = await runtime.expose(endpoint.port) + endpoint_env = endpoint.env(base_url=endpoint_url, model=context.model) + result = await run_dspy_flight_agent( + task=task, + base_url=endpoint_env["OPENAI_BASE_URL"], + api_key=endpoint_env["OPENAI_API_KEY"], + model=endpoint_env["OPENAI_MODEL"], + max_iters=self.config.max_iters, + ) + + state.artifacts.update(result.model_dump(mode="json", exclude_none=True)) + message = vf.AssistantMessage(content=result.process_result) + state.transcript.append(vf.Turn(prompt=prompt, completion=[message])) + state.stop("dspy_completed") + + +def load_taskset(config: DSPyFlightsTasksetConfig) -> DSPyFlightsTaskset: + return DSPyFlightsTaskset(config=config) + + +def load_harness(config: DSPyFlightsHarnessConfig) -> DSPyFlightsHarness: + return DSPyFlightsHarness(config=config) diff --git a/environments/dspy_flights/pyproject.toml b/environments/dspy_flights_v1/pyproject.toml similarity index 84% rename from environments/dspy_flights/pyproject.toml rename to environments/dspy_flights_v1/pyproject.toml index 3af25fae31..dfc7644305 100644 --- a/environments/dspy_flights/pyproject.toml +++ b/environments/dspy_flights_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "dspy-flights" +name = "dspy-flights-v1" version = "0.1.0" tags = ["dspy", "program", "v1"] license = "Apache-2.0" @@ -16,7 +16,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["dspy_flights.py", "pyproject.toml"] +include = ["dspy_flights_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/dspy_rlm/dspy_rlm.py b/environments/dspy_rlm/dspy_rlm.py deleted file mode 100644 index bf332c22c6..0000000000 --- a/environments/dspy_rlm/dspy_rlm.py +++ /dev/null @@ -1,131 +0,0 @@ -import re -from typing import cast - -import verifiers as vf -from verifiers.utils.data_utils import load_example_dataset - - -class DSPYRLMTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["answer_reward"] - taskset_id: str = "gsm8k-dspy-rlm" - num_train_examples: int = 50 - num_eval_examples: int = 20 - - -async def run_dspy_rlm_program(task: vf.Task, state: vf.State) -> vf.State: - import dspy - from openai import OpenAI - - endpoint_config = state.get_endpoint_config(api="chat") - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - endpoint_api_key = endpoint_client.api_key - endpoint_client.close() - lm = dspy.LM( - f"openai/{endpoint_config.model}", - api_base=endpoint_config.base_url, - api_key=endpoint_api_key, - cache=False, - ) - - with dspy.context(lm=lm): - question = task.get("question") - if question is not None: - query = str(question) - else: - query = "" - prompt = task.get("prompt") - if isinstance(prompt, list) and prompt: - query = str(vf.get_messages(prompt)[-1].content or "") - rlm = dspy.RLM("query -> answer", max_iterations=10) - result = await rlm.aforward(query=query) - - final_output = str(result.answer) - state["agent_result"] = final_output - state["completion"] = [{"role": "assistant", "content": final_output}] - return state - - -def load_gsm8k_tasks(split: str, num_examples: int): - n = num_examples if num_examples > 0 else None - return load_example_dataset("gsm8k", split=split, n=n) - - -def load_tasks( - split: vf.TaskSplit = "train", - num_train_examples: int = 50, - num_eval_examples: int = 20, -): - dataset_split = "train" if split == "train" else "test" - num_examples = num_train_examples if split == "train" else num_eval_examples - return load_gsm8k_tasks(dataset_split, num_examples) - - -def extract_dspy_answer(text: str) -> str: - match = re.search(r"SUBMIT\((.+?)\)", text) - if match: - return match.group(1).strip().strip("'\"") - - match = re.search( - r"\[\[\s*##\s*answer\s*##\s*\]\]\s*(.+?)(?:\n|$)", text, re.IGNORECASE - ) - if match: - return match.group(1).strip() - - for line in reversed(text.strip().split("\n")): - line = line.strip() - if line and not line.startswith("[[ ##"): - return line - return "" - - -def answers_match(agent_answer: str, answer: str) -> float: - try: - parsed_agent_answer = float(agent_answer.replace(",", "")) - parsed_answer = float(answer.replace(",", "")) - except (ValueError, TypeError): - return 1.0 if agent_answer.strip() == answer.strip() else 0.0 - return 1.0 if abs(parsed_agent_answer - parsed_answer) < 0.01 else 0.0 - - -def answer_reward(task: vf.Task, state: vf.State) -> float: - """Check if the agent's final output contains the correct answer.""" - result = state.get("agent_result") - if result is not None: - text = str(result) - else: - completion = state.get("completion") - messages = [] - if isinstance(completion, list): - messages = vf.get_messages(completion, role="assistant") or vf.get_messages( - completion - ) - text = str(messages[-1].content or "") if messages else "" - agent_answer = extract_dspy_answer(text) - if not agent_answer: - return 0.0 - return answers_match(agent_answer, str(task.get("answer", ""))) - - -class DSPYRLMTaskset(vf.Taskset[DSPYRLMTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks( - split=split, - num_train_examples=self.config.num_train_examples, - num_eval_examples=self.config.num_eval_examples, - ) - - -class DSPYRLMHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="run_dspy_rlm_program") - - -class DSPYRLMEnvConfig(vf.EnvConfig): - taskset: DSPYRLMTasksetConfig = DSPYRLMTasksetConfig() - harness: DSPYRLMHarnessConfig = DSPYRLMHarnessConfig() - - -def load_environment(config: DSPYRLMEnvConfig) -> vf.Env: - return vf.Env( - taskset=DSPYRLMTaskset(config=config.taskset), - harness=vf.Harness(config=config.harness), - ) diff --git a/environments/dspy_rlm/README.md b/environments/dspy_rlm_v1/README.md similarity index 85% rename from environments/dspy_rlm/README.md rename to environments/dspy_rlm_v1/README.md index c4dd16d195..65eca38a00 100644 --- a/environments/dspy_rlm/README.md +++ b/environments/dspy_rlm_v1/README.md @@ -1,12 +1,12 @@ -# dspy-rlm +# dspy-rlm-v1 - + Source Code ### Overview -- **Environment ID**: `dspy-rlm` +- **Environment ID**: `dspy-rlm-v1` - **Short description**: V1 Taskset/Harness example using DSPy's RLM (Recursive Language Model) module on GSM8K math problems. - **Tags**: v1, taskset, harness, dspy, rlm, math, gsm8k @@ -23,7 +23,7 @@ ### How it works -The taskset owns GSM8K train/eval task loading and reward logic. The harness runs an in-process DSPy RLM program, builds its LM from `state.get_endpoint_config(api="chat")`, and routes every model call through the V1 interception endpoint. +The taskset owns GSM8K train/eval task loading and reward logic. The harness runs an in-process DSPy RLM agent, starts the v1 protocol endpoint for the rollout, and gives DSPy an OpenAI-compatible LM pointed at that endpoint. DSPy RLM requires Deno to be available in the runtime environment. @@ -32,13 +32,13 @@ DSPy RLM requires Deno to be available in the runtime environment. Run an evaluation with default settings: ```bash -prime eval run dspy-rlm +prime eval run dspy-rlm-v1 ``` Configure model and sampling: ```bash -prime eval run dspy-rlm \ +prime eval run dspy-rlm-v1 \ -m gpt-4.1-mini \ -n 10 -r 3 -t 1024 -T 0.7 ``` diff --git a/environments/dspy_rlm_v1/dspy_rlm_v1/__init__.py b/environments/dspy_rlm_v1/dspy_rlm_v1/__init__.py new file mode 100644 index 0000000000..9f8a5394f3 --- /dev/null +++ b/environments/dspy_rlm_v1/dspy_rlm_v1/__init__.py @@ -0,0 +1 @@ +"""dspy-rlm-v1 environment package.""" diff --git a/environments/dspy_rlm_v1/dspy_rlm_v1/harness.py b/environments/dspy_rlm_v1/dspy_rlm_v1/harness.py new file mode 100644 index 0000000000..ffec43d242 --- /dev/null +++ b/environments/dspy_rlm_v1/dspy_rlm_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import DSPYRLMHarness as DSPYRLMHarness +from .taskset import DSPYRLMHarnessConfig as DSPYRLMHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/dspy_rlm_v1/dspy_rlm_v1/taskset.py b/environments/dspy_rlm_v1/dspy_rlm_v1/taskset.py new file mode 100644 index 0000000000..3080263d9f --- /dev/null +++ b/environments/dspy_rlm_v1/dspy_rlm_v1/taskset.py @@ -0,0 +1,154 @@ +import re + +import verifiers.v1 as vf +from verifiers.utils.data_utils import load_example_dataset + + +class DSPYRLMTasksetConfig(vf.TasksetConfig): + id: str = "gsm8k-dspy-rlm" + num_train_examples: int = 50 + num_eval_examples: int = 20 + + +class DSPYRLMHarnessConfig(vf.HarnessConfig): + max_iterations: int = 10 + + +class DSPYRLMTask(vf.Task): + question: str + answer: str + + +def load_gsm8k_tasks(split: str, num_examples: int) -> vf.Tasks: + n = num_examples if num_examples > 0 else None + return [ + { + **row, + "row_id": index, + "prompt": [{"role": "user", "content": str(row["question"])}], + } + for index, row in enumerate(load_example_dataset("gsm8k", split=split, n=n)) + ] + + +def extract_dspy_answer(text: str) -> str: + match = re.search(r"SUBMIT\((.+?)\)", text) + if match: + return match.group(1).strip().strip("'\"") + + match = re.search( + r"\[\[\s*##\s*answer\s*##\s*\]\]\s*(.+?)(?:\n|$)", text, re.IGNORECASE + ) + if match: + return match.group(1).strip() + + for line in reversed(text.strip().split("\n")): + line = line.strip() + if line and not line.startswith("[[ ##"): + return line + return "" + + +def answers_match(agent_answer: str, answer: str) -> float: + try: + parsed_agent_answer = float(agent_answer.replace(",", "")) + parsed_answer = float(answer.replace(",", "")) + except (ValueError, TypeError): + return float(agent_answer.strip() == answer.strip()) + return float(abs(parsed_agent_answer - parsed_answer) < 0.01) + + +def final_text(state: vf.State) -> str: + result = state.artifacts.get("agent_result") + if isinstance(result, str): + return result + messages = [ + message for message in state.completion if message.role == "assistant" + ] or state.completion + return str(messages[-1].content or "") if messages else "" + + +class DSPYRLMTaskset(vf.Taskset[DSPYRLMTasksetConfig]): + task_type = DSPYRLMTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + dataset_split = "train" if split == "train" else "test" + num_examples = ( + self.config.num_train_examples + if split == "train" + else self.config.num_eval_examples + ) + return load_gsm8k_tasks(dataset_split, num_examples) + + @vf.reward + async def answer_reward(self, task: DSPYRLMTask, state: vf.State) -> float: + answer = extract_dspy_answer(final_text(state)) + return answers_match(answer, task.answer) if answer else 0.0 + + +class DSPYRLMHarness(vf.Harness[DSPYRLMHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = DSPYRLMTask.model_validate(context.task.model_dump()) + state = context.state + runtime = context.runtime + if runtime is None: + raise ValueError("DSPYRLMHarness requires a runtime.") + prompt = self.initial_messages(task) + + async def stop_check() -> str | None: + if await self.is_completed(context): + return state.stop_condition or "stop" + return None + + async with vf.InterceptionServer( + context, + task, + state, + protocols=self.protocols, + stop_check=stop_check, + ) as endpoint: + endpoint_url = await runtime.expose(endpoint.port) + endpoint_env = endpoint.env(base_url=endpoint_url, model=context.model) + final_output = await run_dspy_rlm( + query=task.question, + base_url=endpoint_env["OPENAI_BASE_URL"], + api_key=endpoint_env["OPENAI_API_KEY"], + model=endpoint_env["OPENAI_MODEL"], + max_iterations=self.config.max_iterations, + ) + + state.artifacts["agent_result"] = final_output + message = vf.AssistantMessage(content=final_output) + if not state.transcript: + state.transcript.append(vf.Turn(prompt=prompt, completion=[message])) + state.stop("dspy_completed") + + +async def run_dspy_rlm( + *, + query: str, + base_url: str, + api_key: str, + model: str, + max_iterations: int, +) -> str: + import dspy + + lm = dspy.LM( + f"openai/{model}", + api_base=base_url, + api_key=api_key, + cache=False, + ) + with dspy.context(lm=lm): + rlm = dspy.RLM("query -> answer", max_iterations=max_iterations) + result = await rlm.aforward(query=query) + return str(result.answer) + + +def load_taskset(config: DSPYRLMTasksetConfig) -> DSPYRLMTaskset: + return DSPYRLMTaskset(config=config) + + +def load_harness(config: DSPYRLMHarnessConfig) -> DSPYRLMHarness: + return DSPYRLMHarness(config=config) diff --git a/environments/dspy_rlm/pyproject.toml b/environments/dspy_rlm_v1/pyproject.toml similarity index 83% rename from environments/dspy_rlm/pyproject.toml rename to environments/dspy_rlm_v1/pyproject.toml index 4a83765a07..aea4b290e8 100644 --- a/environments/dspy_rlm/pyproject.toml +++ b/environments/dspy_rlm_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "dspy-rlm" +name = "dspy-rlm-v1" description = "V1 Taskset/Harness environment using DSPy's RLM module" tags = ["v1", "taskset", "harness", "dspy", "rlm"] version = "0.1.0" @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["dspy_rlm.py", "pyproject.toml"] +include = ["dspy_rlm_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/harbor_v1/README.md b/environments/harbor_v1/README.md new file mode 100644 index 0000000000..18f80b715c --- /dev/null +++ b/environments/harbor_v1/README.md @@ -0,0 +1,78 @@ +# harbor-v1 + +### Overview +- **Environment ID**: `harbor-v1` +- **Short description**: Generic v1 Harbor taskset environment +- **Tags**: harbor, cli_agent, v1 + +### Datasets +- **Primary dataset(s)**: Harbor task directories +- **Source links**: +- **Split sizes**: 1 bundled smoke task by default + +### Task +- **Type**: multiturn, cli_agent +- **Rubric overview**: Reward returned by running Harbor verifier tests + +### Quickstart +Run the environment: + +```bash +prime eval run harbor-v1 +``` + +Configure model and sampling: + +```bash +prime eval run harbor-v1 -m openai/gpt-4.1-mini -n 1 -r 1 -t 1024 -T 0.7 +``` + +Notes: +- v1 task settings belong under `config.taskset` when passed through `-a` / `--env-args`. +- Use `taskset` and `harness` config sections for v1 object configuration in TOML. +- Harbor tasks with `image` fields require a container runtime such as Docker or Prime. + +### Taskset Config + +| Arg | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `source` | `"harbor" \| "package"` | `"package"` in this environment | Dataset source resolver. | +| `dataset` | str | `"harbor_v1"` in this environment | Harbor dataset id or Python package name. | +| `tasks` | list[str] | `null` | Explicit Harbor task names to run. | +| `cache_dir` | str | `null` | Optional Harbor cache root override. | +| `refresh` | bool | `false` | Refresh Harbor cache before loading. | +| `require_image` | bool | `false` | Require every task to declare `[environment].docker_image`. | + +### Harness Config + +This package defaults to the reusable `OpenCode` harness. OpenCode settings +belong under `config.harness`: + +```toml +[env.harness] +max_turns = 4 +cwd = "/app" +``` + +### Metrics + +| Metric | Meaning | +| ------ | ------- | +| `reward` | Harbor verifier reward, usually `0.0` or `1.0` | +| `num_turns` | Number of intercepted assistant turns | + + +## How It Works + +1. `HarborTaskset` resolves Harbor or package task directories into typed + v1 tasks. +2. The taskset maps `[environment].docker_image` and resource hints onto + generic v1 `Task` fields. +3. `OpenCode` runs as the default agent harness. +4. Reward is computed by staging only `tests/` into the live runtime after the + rollout and running `tests/test.sh`. + +## Requirements + +- Harbor task directory with `task.toml`, `instruction.md`, and `tests/` +- A container runtime when tasks declare `[environment].docker_image` diff --git a/environments/harbor_v1/harbor_v1/__init__.py b/environments/harbor_v1/harbor_v1/__init__.py new file mode 100644 index 0000000000..a1b6119c18 --- /dev/null +++ b/environments/harbor_v1/harbor_v1/__init__.py @@ -0,0 +1 @@ +"""harbor-v1 environment package.""" diff --git a/environments/harbor_v1/harbor_v1/harness.py b/environments/harbor_v1/harbor_v1/harness.py new file mode 100644 index 0000000000..4b08f96819 --- /dev/null +++ b/environments/harbor_v1/harbor_v1/harness.py @@ -0,0 +1,5 @@ +from harnesses import OpenCode, OpenCodeConfig + + +def load_harness(config: OpenCodeConfig) -> OpenCode: + return OpenCode(config=config) diff --git a/environments/opencode_harbor/tasks/hello-world/instruction.md b/environments/harbor_v1/harbor_v1/tasks/hello-world/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/hello-world/instruction.md rename to environments/harbor_v1/harbor_v1/tasks/hello-world/instruction.md diff --git a/environments/harbor_v1/harbor_v1/tasks/hello-world/solution/solve.sh b/environments/harbor_v1/harbor_v1/tasks/hello-world/solution/solve.sh new file mode 100644 index 0000000000..cf5a660751 --- /dev/null +++ b/environments/harbor_v1/harbor_v1/tasks/hello-world/solution/solve.sh @@ -0,0 +1,4 @@ +#!/usr/bin/env bash +set -euo pipefail + +printf '%s\n' 'Hello, world!' > /app/hello.txt diff --git a/environments/harbor_v1/harbor_v1/tasks/hello-world/task.toml b/environments/harbor_v1/harbor_v1/tasks/hello-world/task.toml new file mode 100644 index 0000000000..27d703d46a --- /dev/null +++ b/environments/harbor_v1/harbor_v1/tasks/hello-world/task.toml @@ -0,0 +1,18 @@ +version = "1.0" + +[task] +name = "hello-world" +description = "Create a hello.txt file." +keywords = ["smoke", "file"] + +[verifier] +timeout_sec = 120.0 + +[agent] +timeout_sec = 120.0 + +[environment] +docker_image = "python:3.11-slim" +cpus = 1 +memory = "1G" +storage = "2G" diff --git a/environments/harbor_v1/harbor_v1/tasks/hello-world/tests/test.sh b/environments/harbor_v1/harbor_v1/tasks/hello-world/tests/test.sh new file mode 100644 index 0000000000..3ac9d06d0b --- /dev/null +++ b/environments/harbor_v1/harbor_v1/tasks/hello-world/tests/test.sh @@ -0,0 +1,8 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [ "$(cat /app/hello.txt 2>/dev/null || true)" = "Hello, world!" ]; then + echo 1 > /logs/verifier/reward.txt +else + echo 0 > /logs/verifier/reward.txt +fi diff --git a/environments/harbor_v1/harbor_v1/taskset.py b/environments/harbor_v1/harbor_v1/taskset.py new file mode 100644 index 0000000000..2874d9cfca --- /dev/null +++ b/environments/harbor_v1/harbor_v1/taskset.py @@ -0,0 +1,13 @@ +from tasksets import HarborTaskset, HarborTasksetConfig + + +def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: + taskset_config = config + if ( + "source" not in config.model_fields_set + and "dataset" not in config.model_fields_set + ): + taskset_config = taskset_config.model_copy( + update={"source": "package", "dataset": "harbor_v1"} + ) + return HarborTaskset(config=taskset_config) diff --git a/environments/harbor_v1/pyproject.toml b/environments/harbor_v1/pyproject.toml new file mode 100644 index 0000000000..d15f2ffdcf --- /dev/null +++ b/environments/harbor_v1/pyproject.toml @@ -0,0 +1,29 @@ +[project] +name = "harbor-v1" +description = "Generic v1 Harbor taskset environment" +license = "MIT" +tags = ["eval", "cli_agent", "v1", "harbor", "taskset", "harness"] +version = "0.1.0" +requires-python = ">=3.10" +dependencies = [ + "verifiers>=0.1.10.dev4", + "harnesses>=0.1.2", + "tasksets>=0.1.1", + "prime-sandboxes>=0.2.10", + "tomli>=2.3.0", +] + + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build] +include = ["harbor_v1/**/*", "README.md", "pyproject.toml"] + +[tool.verifiers.eval] +num_examples = 1 +rollouts_per_example = 1 + +[tool.ruff] +exclude = ["harbor_v1/tasks/**"] diff --git a/environments/hello_group_reward_v1/README.md b/environments/hello_group_reward_v1/README.md index 5aee26a1c8..6841781574 100644 --- a/environments/hello_group_reward_v1/README.md +++ b/environments/hello_group_reward_v1/README.md @@ -10,7 +10,7 @@ all rollouts finish, the group stage: - runs a group update that writes per-rollout ranks and group summaries; - records group metrics; - adds a relative group reward; -- writes explicit centered advantages; +- writes explicit centered token advantages when model tokens are present; - runs group cleanup. This is meant as a compact reference for `@vf.reward(stage="group")` and diff --git a/environments/hello_group_reward_v1/hello_group_reward_v1.py b/environments/hello_group_reward_v1/hello_group_reward_v1.py deleted file mode 100644 index 48d0a21ead..0000000000 --- a/environments/hello_group_reward_v1/hello_group_reward_v1.py +++ /dev/null @@ -1,335 +0,0 @@ -from difflib import SequenceMatcher -from statistics import mean - -import verifiers as vf - - -SYSTEM_PROMPT = """\ -You are testing group-aware scoring. Each rollout receives one candidate answer. -Return the assigned candidate exactly. -""" - - -class GroupRewardTasksetConfig(vf.TasksetConfig): - metrics: list[str] = [ - "answer_length", - "group_quality", - "group_rank", - ] - rewards: list[str] = [ - "rollout_similarity", - "relative_group_reward", - ] - advantages: list[str] = ["centered_group_advantage"] - updates: list[str] = ["summarize_group"] - cleanups: list[str] = ["mark_group_cleaned"] - system_prompt: str = SYSTEM_PROMPT - num_examples: int = -1 - - -class GroupRewardHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="candidate_program") - max_turns: int = 1 - - -def group_reward_task( - task_id: str, - question: str, - target: str, - near: str, - partial: str, - wrong: str, -) -> vf.ConfigData: - return { - "task_id": task_id, - "question": question, - "target": target, - "candidates": [ - {"id": "exact", "answer": target}, - {"id": "near", "answer": near}, - {"id": "partial", "answer": partial}, - {"id": "off-topic", "answer": wrong}, - ], - } - - -TASKS: list[vf.ConfigData] = [ - group_reward_task( - "distributed-systems", - "Describe v1 verifiers in one short phrase.", - "composable tasksets and harnesses with group-aware scoring", - "composable tasksets and harnesses with rollout scoring", - "tasksets and harnesses", - "a single monolithic environment object", - ), - group_reward_task( - "runtime-boundary", - "Describe the v1 runtime boundary in one short phrase.", - "serializable task and state with hidden runtime handles", - "serializable task and state with runtime handles", - "task and state", - "global objects stored directly in every task", - ), - group_reward_task( - "toolset-scope", - "Describe v1 toolset scope in one short phrase.", - "rollout group and global tool lifetimes", - "rollout and group tool lifetimes", - "tool lifetimes", - "static imports with no runtime handles", - ), - group_reward_task( - "sandbox-sharing", - "Describe sandbox sharing in one short phrase.", - "borrowed sandbox handles shared across nested stages", - "sandbox handles shared across stages", - "shared sandbox handles", - "new isolated machines for every function call", - ), - group_reward_task( - "endpoint-controls", - "Describe endpoint controls in one short phrase.", - "nested programs inherit active model endpoint controls", - "programs inherit model endpoint controls", - "model endpoint controls", - "hardcoded providers inside tasks", - ), - group_reward_task( - "users", - "Describe v1 users in one short phrase.", - "task-owned follow-up messages between assistant turns", - "follow-up messages between assistant turns", - "follow-up messages", - "metrics computed before any rollout starts", - ), - group_reward_task( - "program-uploads", - "Describe program uploads in one short phrase.", - "task fields and files staged before harness execution", - "files staged before harness execution", - "staged files", - "reward weights serialized into model prompts", - ), - group_reward_task( - "cleanup-hooks", - "Describe cleanup hooks in one short phrase.", - "final artifact collection after rewards and metrics", - "artifact collection after metrics", - "artifact collection", - "dataset filtering before import time", - ), - group_reward_task( - "harbor-taskset", - "Describe HarborTaskset in one short phrase.", - "task directories converted into sandboxed rollout rows", - "task directories converted into rollout rows", - "task directories", - "chat templates stored in every reward function", - ), - group_reward_task( - "advantage-baseline", - "Describe group advantages in one short phrase.", - "rollout rewards centered against a group baseline", - "rewards centered against a baseline", - "centered rewards", - "single responses scored without group context", - ), -] - - -class GroupRewardTaskset(vf.Taskset[GroupRewardTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(num_examples=self.config.num_examples) - - async def init_group( - self, task: vf.Task, num_rollouts: int - ) -> tuple[list[vf.Task], list[vf.State]]: - candidates = task.get("candidates") - if not isinstance(candidates, list) or not candidates: - raise ValueError("hello_group_reward_v1 tasks require candidates.") - tasks: list[vf.Task] = [] - states: list[vf.State] = [] - for rollout_index in range(num_rollouts): - candidate = candidates[rollout_index % len(candidates)] - if not isinstance(candidate, dict): - raise TypeError("candidate entries must be mappings.") - candidate_id = str(candidate["id"]) - candidate_answer = str(candidate["answer"]) - group_task = vf.Task( - { - **dict(task), - "candidate_id": candidate_id, - "candidate_answer": candidate_answer, - "rollout_index": rollout_index, - "prompt": [ - { - "role": "user", - "content": ( - f"Question: {task['question']}\n" - f"Assigned candidate id: {candidate_id}\n" - f"Assigned candidate answer: {candidate_answer}\n\n" - "Return the assigned candidate answer exactly." - ), - } - ], - "max_turns": 1, - } - ).freeze() - state = vf.State.for_task(group_task) - state["group_setup"] = { - "base_task_id": task["task_id"], - "num_rollouts": num_rollouts, - "candidate_id": candidate_id, - } - tasks.append(group_task) - states.append(state) - return tasks, states - - -class GroupRewardHarness(vf.Harness[GroupRewardHarnessConfig]): - pass - - -async def candidate_program(task: vf.Task, state: vf.State) -> vf.State: - answer = str(task["candidate_answer"]) - state["answer"] = answer - state["candidate_id"] = task["candidate_id"] - state["completion"] = [{"role": "assistant", "content": answer}] - state.stop("candidate_program") - return state - - -@vf.metric -async def answer_length(task, state) -> float: - _ = task - return float(len(str(state.get("answer") or ""))) - - -@vf.reward(weight=0.1) -async def rollout_similarity(task, state) -> float: - return candidate_quality(str(task["target"]), str(state.get("answer") or "")) - - -@vf.update(stage="group", priority=10) -async def summarize_group(tasks, states) -> None: - qualities = [ - candidate_quality(str(task["target"]), str(state.get("answer") or "")) - for task, state in zip(tasks, states) - ] - ranks = dense_ranks(qualities) - best_index = min( - range(len(states)), - key=lambda index: ( - -qualities[index], - len(str(states[index].get("answer") or "")), - ), - ) - candidates = [ - { - "candidate_id": str(state.get("candidate_id")), - "quality": qualities[index], - "rank": ranks[index], - } - for index, state in enumerate(states) - ] - for index, state in enumerate(states): - state["group_summary"] = { - "best_candidate_id": str(states[best_index].get("candidate_id")), - "best_answer": str(states[best_index].get("answer") or ""), - "quality": qualities[index], - "rank": ranks[index], - "group_size": len(states), - "candidates": candidates, - } - - -@vf.metric(stage="group") -async def group_quality(tasks, states) -> list[float]: - return [ - candidate_quality(str(task["target"]), str(state.get("answer") or "")) - for task, state in zip(tasks, states) - ] - - -@vf.metric(stage="group") -async def group_rank(tasks, states) -> list[float]: - _ = tasks - qualities = [ - float(state.get("group_summary", {}).get("quality", 0.0)) for state in states - ] - return [float(rank) for rank in dense_ranks(qualities)] - - -@vf.reward(stage="group", weight=1.0) -async def relative_group_reward(tasks, states) -> list[float]: - _ = tasks - qualities = [ - float(state.get("group_summary", {}).get("quality", 0.0)) for state in states - ] - if not qualities: - return [] - low = min(qualities) - high = max(qualities) - if high == low: - return [0.5 for _ in qualities] - return [(quality - low) / (high - low) for quality in qualities] - - -@vf.advantage -async def centered_group_advantage(tasks, states) -> list[float]: - _ = tasks - rewards = [ - float(state.get("metrics", {}).get("relative_group_reward", 0.0)) - for state in states - ] - baseline = mean(rewards) if rewards else 0.0 - return [reward - baseline for reward in rewards] - - -@vf.cleanup(stage="group") -async def mark_group_cleaned(tasks, states) -> None: - _ = tasks - for state in states: - state["group_cleaned"] = True - - -def candidate_quality(target: str, answer: str) -> float: - if not answer: - return 0.0 - ratio = SequenceMatcher(None, answer.lower(), target.lower()).ratio() - length_ratio = min(len(answer), len(target)) / max(len(answer), len(target)) - return round(0.8 * ratio + 0.2 * length_ratio, 6) - - -def dense_ranks(values: list[float]) -> list[int]: - ordered = sorted(set(values), reverse=True) - return [ordered.index(value) + 1 for value in values] - - -def load_tasks(num_examples: int = -1): - records = TASKS if num_examples < 0 else TASKS[:num_examples] - for index, record in enumerate(records): - yield { - **record, - "example_id": index, - "answer": record["target"], - "prompt": [ - { - "role": "user", - "content": str(record["question"]), - } - ], - "max_turns": 1, - } - - -class GroupRewardEnvConfig(vf.EnvConfig): - taskset: GroupRewardTasksetConfig = GroupRewardTasksetConfig() - harness: GroupRewardHarnessConfig = GroupRewardHarnessConfig() - - -def load_environment(config: GroupRewardEnvConfig) -> vf.Env: - return vf.Env( - taskset=GroupRewardTaskset(config=config.taskset), - harness=GroupRewardHarness(config=config.harness), - ) diff --git a/environments/hello_group_reward_v1/hello_group_reward_v1/__init__.py b/environments/hello_group_reward_v1/hello_group_reward_v1/__init__.py new file mode 100644 index 0000000000..268a3401c6 --- /dev/null +++ b/environments/hello_group_reward_v1/hello_group_reward_v1/__init__.py @@ -0,0 +1 @@ +"""hello-group-reward-v1 environment package.""" diff --git a/environments/hello_group_reward_v1/hello_group_reward_v1/harness.py b/environments/hello_group_reward_v1/hello_group_reward_v1/harness.py new file mode 100644 index 0000000000..9522bc4a2f --- /dev/null +++ b/environments/hello_group_reward_v1/hello_group_reward_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import GroupRewardHarness as GroupRewardHarness +from .taskset import GroupRewardHarnessConfig as GroupRewardHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/hello_group_reward_v1/hello_group_reward_v1/taskset.py b/environments/hello_group_reward_v1/hello_group_reward_v1/taskset.py new file mode 100644 index 0000000000..ea9e5398c4 --- /dev/null +++ b/environments/hello_group_reward_v1/hello_group_reward_v1/taskset.py @@ -0,0 +1,287 @@ +from difflib import SequenceMatcher + +import verifiers.v1 as vf + + +SYSTEM_PROMPT = """\ +You are testing group-aware scoring. Each rollout receives one candidate answer. +Return the assigned candidate exactly. +""" + + +class GroupRewardTasksetConfig(vf.TasksetConfig): + system_prompt: str = SYSTEM_PROMPT + num_examples: int = -1 + + +class GroupRewardHarnessConfig(vf.HarnessConfig): + max_turns: int = 1 + + +class GroupRewardTask(vf.Task): + question: str + target: str + answer: str + near_answer: str + partial_answer: str + wrong_answer: str + + +class GroupRolloutTask(GroupRewardTask): + parent_task_id: str + candidate_id: str + candidate_answer: str + rollout_index: int + + +def group_reward_task( + task_id: str, + question: str, + target: str, + near: str, + partial: str, + wrong: str, +) -> vf.JsonData: + return { + "task_id": task_id, + "question": question, + "target": target, + "prompt": question, + "near_answer": near, + "partial_answer": partial, + "wrong_answer": wrong, + } + + +TASKS: list[vf.JsonData] = [ + group_reward_task( + "distributed-systems", + "Describe v1 verifiers in one short phrase.", + "composable tasksets and harnesses with group-aware scoring", + "composable tasksets and harnesses with rollout scoring", + "tasksets and harnesses", + "a single monolithic environment object", + ), + group_reward_task( + "runtime-boundary", + "Describe the v1 runtime boundary in one short phrase.", + "serializable task and state with hidden runtime handles", + "serializable task and state with runtime handles", + "task and state", + "global objects stored directly in every task", + ), + group_reward_task( + "toolset-scope", + "Describe v1 toolset scope in one short phrase.", + "rollout group and global tool lifetimes", + "rollout and group tool lifetimes", + "tool lifetimes", + "static imports with no runtime handles", + ), + group_reward_task( + "sandbox-sharing", + "Describe sandbox sharing in one short phrase.", + "borrowed sandbox handles shared across nested stages", + "sandbox handles shared across stages", + "shared sandbox handles", + "new isolated machines for every function call", + ), + group_reward_task( + "endpoint-controls", + "Describe endpoint controls in one short phrase.", + "nested programs inherit active model endpoint controls", + "programs inherit model endpoint controls", + "model endpoint controls", + "hardcoded providers inside tasks", + ), + group_reward_task( + "users", + "Describe v1 users in one short phrase.", + "task-owned follow-up messages between assistant turns", + "follow-up messages between assistant turns", + "follow-up messages", + "metrics computed before any rollout starts", + ), + group_reward_task( + "program-uploads", + "Describe program uploads in one short phrase.", + "task fields and files staged before harness execution", + "files staged before harness execution", + "staged files", + "reward weights serialized into model prompts", + ), + group_reward_task( + "cleanup-hooks", + "Describe cleanup hooks in one short phrase.", + "final artifact collection after rewards and metrics", + "artifact collection after metrics", + "artifact collection", + "dataset filtering before import time", + ), + group_reward_task( + "harbor-taskset", + "Describe HarborTaskset in one short phrase.", + "task directories converted into sandboxed rollout rows", + "task directories converted into rollout rows", + "task directories", + "chat templates stored in every reward function", + ), + group_reward_task( + "advantage-baseline", + "Describe group advantages in one short phrase.", + "rollout rewards centered against a group baseline", + "rewards centered against a baseline", + "centered rewards", + "single responses scored without group context", + ), +] + + +class GroupRewardTaskset(vf.Taskset[GroupRewardTasksetConfig]): + task_type = GroupRewardTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return [ + GroupRewardTask.model_validate(record) + for record in load_tasks(num_examples=self.config.num_examples) + ] + + async def init_group( + self, task: GroupRewardTask, num_rollouts: int + ) -> tuple[list[GroupRolloutTask], list[vf.State]]: + candidates = [ + ("exact", task.target), + ("near", task.near_answer), + ("partial", task.partial_answer), + ("off-topic", task.wrong_answer), + ] + tasks: list[GroupRolloutTask] = [] + states: list[vf.State] = [] + for rollout_index in range(num_rollouts): + candidate_id, candidate_answer = candidates[rollout_index % len(candidates)] + task_data = task.model_dump( + mode="json", exclude_none=True, exclude_defaults=True + ) + task_data.pop("task_id", None) + group_task = GroupRolloutTask.model_validate( + { + **task_data, + "parent_task_id": task.task_id, + "candidate_id": candidate_id, + "candidate_answer": candidate_answer, + "rollout_index": rollout_index, + "prompt": ( + f"Question: {task.question}\n" + f"Assigned candidate id: {candidate_id}\n" + f"Assigned candidate answer: {candidate_answer}\n\n" + "Return the assigned candidate answer exactly." + ), + "max_turns": 1, + } + ) + state = vf.State(task_id=group_task.task_id) + state.extras["group_setup"] = { + "base_task_id": task.task_id, + "num_rollouts": num_rollouts, + "candidate_id": candidate_id, + } + tasks.append(group_task) + states.append(state) + return tasks, states + + @vf.metric + async def answer_length(self, state: vf.State) -> float: + return float(len(str(state.extras.get("answer") or ""))) + + @vf.reward(weight=0.1) + async def rollout_similarity( + self, task: GroupRolloutTask, state: vf.State + ) -> float: + return candidate_quality(task.target, str(state.extras.get("answer") or "")) + + @vf.metric(stage="group") + async def group_quality( + self, tasks: list[GroupRolloutTask], states: list[vf.State] + ) -> list[float]: + return [ + candidate_quality(task.target, str(state.extras.get("answer") or "")) + for task, state in zip(tasks, states, strict=True) + ] + + @vf.metric(stage="group") + async def group_rank( + self, tasks: list[GroupRolloutTask], states: list[vf.State] + ) -> list[float]: + qualities = [ + candidate_quality(task.target, str(state.extras.get("answer") or "")) + for task, state in zip(tasks, states, strict=True) + ] + return [float(rank) for rank in dense_ranks(qualities)] + + @vf.reward(stage="group", weight=1.0) + async def relative_group_reward( + self, tasks: list[GroupRolloutTask], states: list[vf.State] + ) -> list[float]: + qualities = [ + candidate_quality(task.target, str(state.extras.get("answer") or "")) + for task, state in zip(tasks, states, strict=True) + ] + if not qualities: + return [] + low = min(qualities) + high = max(qualities) + if high == low: + return [0.5 for _ in qualities] + return [(quality - low) / (high - low) for quality in qualities] + + +class GroupRewardHarness(vf.Harness[GroupRewardHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = GroupRolloutTask.model_validate(context.task.model_dump()) + state = context.state + answer = task.candidate_answer + message = vf.AssistantMessage(content=answer) + state.extras["answer"] = answer + state.extras["candidate_id"] = task.candidate_id + state.transcript.append( + vf.Turn(prompt=self.initial_messages(task), completion=[message]) + ) + state.stop("candidate_program") + + +def candidate_quality(target: str, answer: str) -> float: + if not answer: + return 0.0 + ratio = SequenceMatcher(None, answer.lower(), target.lower()).ratio() + length_ratio = min(len(answer), len(target)) / max(len(answer), len(target)) + return round(0.8 * ratio + 0.2 * length_ratio, 6) + + +def dense_ranks(values: list[float]) -> list[int]: + ordered = sorted(set(values), reverse=True) + return [ordered.index(value) + 1 for value in values] + + +def load_tasks(num_examples: int = -1): + records = TASKS if num_examples < 0 else TASKS[:num_examples] + for index, record in enumerate(records): + yield { + **record, + "example_id": index, + "answer": record["target"], + "prompt": [ + { + "role": "user", + "content": str(record["question"]), + } + ], + "max_turns": 1, + } + + +def load_taskset(config: GroupRewardTasksetConfig) -> GroupRewardTaskset: + return GroupRewardTaskset(config=config) + + +def load_harness(config: GroupRewardHarnessConfig) -> GroupRewardHarness: + return GroupRewardHarness(config=config) diff --git a/environments/hello_group_reward_v1/pyproject.toml b/environments/hello_group_reward_v1/pyproject.toml index 0ea3fe68fb..b77b060895 100644 --- a/environments/hello_group_reward_v1/pyproject.toml +++ b/environments/hello_group_reward_v1/pyproject.toml @@ -14,7 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["hello_group_reward_v1.py", "README.md", "pyproject.toml"] +include = ["hello_group_reward_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/hello_parallel_sandbox_v1/README.md b/environments/hello_parallel_sandbox_v1/README.md index 7e9d01e93e..244aa65b0e 100644 --- a/environments/hello_parallel_sandbox_v1/README.md +++ b/environments/hello_parallel_sandbox_v1/README.md @@ -9,12 +9,9 @@ tool bound with `sandbox="program"`, so tool calls execute in that primary program sandbox instead of creating a separate tool sandbox. The parent writes `/tmp/answer.txt`, then: -1. two update-stage child harnesses run concurrently, borrow the live `model` - and `bash` tool, append their model calls to the public trajectory, and - inspect the same primary sandbox before rollout cleanup; -2. a reward-stage child harness borrows the same live `model` and `bash` tool, - keeps its trajectory private, inspects the same sandbox, and returns a JSON - score. +1. two update-stage audits run concurrently and store serializable findings in + `state.extras` and `state.artifacts`; +2. reward scoring reads those findings and writes a scalar rollout reward. ```bash prime eval run hello-parallel-sandbox-v1 -m openai/gpt-5.4-mini -n 3 -r 1 -t 4096 diff --git a/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1.py b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1.py deleted file mode 100644 index fe53ecbb62..0000000000 --- a/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1.py +++ /dev/null @@ -1,377 +0,0 @@ -import asyncio -import json - -import verifiers as vf -from verifiers.v1.utils.judge_utils import ( - clamp_float, - parse_judge_json, - truncate_command_record, - truncate_text, -) - -SYSTEM_PROMPT = """You are a careful sandbox operator. - -You are running inside a sandboxed harness program. Use the callable bash tool -to complete the requested file task inside that same primary program sandbox. -Before answering, write: - -- `/tmp/answer.txt` with the exact requested answer text; -- `/tmp/worklog.md` with a short note about what you did. - -Then reply with the final answer text only. -""" - -FILE_AUDIT_SYSTEM_PROMPT = """You are a file-state auditor. - -You are running inside the same sandbox as the answer rollout and have a -callable bash tool. -Call bash to inspect `/tmp/answer.txt` and `/tmp/worklog.md`, then report -whether the files exist and what they contain. Do not assign a numeric score. -""" - -COMMAND_AUDIT_SYSTEM_PROMPT = """You are a process auditor. - -You are running inside the same sandbox as the answer rollout and have a -callable bash tool. -Call bash to inspect `/tmp`, file metadata, and any useful command artifacts. -Report whether the sandbox state looks consistent with the task. Do not assign -a numeric score. -""" - -REWARD_JUDGE_SYSTEM_PROMPT = """You are a scoring judge. - -You are running inside the same sandbox as the answer rollout and have a -callable bash tool. -Call bash to inspect `/tmp/answer.txt` before scoring. Respond with compact JSON -only: - -{"score": 0.0-1.0, "reason": "..."} -""" - -TASKS: list[vf.ConfigData] = [ - { - "task_id": "exact-token", - "answer": "prime-v1-shared-sandbox", - "instruction": ( - "Create `/tmp/answer.txt` containing exactly `prime-v1-shared-sandbox`." - ), - }, - { - "task_id": "reverse-token", - "answer": "xobdnas-derahs", - "instruction": ( - "Create `/tmp/answer.txt` containing exactly the reverse of " - "`shared-sandbox`." - ), - }, - { - "task_id": "joined-words", - "answer": "taskset-harness-runtime", - "instruction": ( - "Create `/tmp/answer.txt` containing the words taskset, harness, " - "and runtime joined by hyphens." - ), - }, - { - "task_id": "uppercase-token", - "answer": "SANDBOX", - "instruction": "Create `/tmp/answer.txt` containing `sandbox` in uppercase.", - }, - { - "task_id": "lowercase-token", - "answer": "runtime", - "instruction": "Create `/tmp/answer.txt` containing `RUNTIME` in lowercase.", - }, - { - "task_id": "count-letters", - "answer": "9", - "instruction": ( - "Create `/tmp/answer.txt` containing the number of letters in `verifiers`." - ), - }, - { - "task_id": "repeat-prefix", - "answer": "v1-v1-v1", - "instruction": ( - "Create `/tmp/answer.txt` containing `v1` repeated three times, " - "joined by hyphens." - ), - }, - { - "task_id": "basename", - "answer": "answer.txt", - "instruction": ( - "Create `/tmp/answer.txt` containing the basename of the path " - "`/tmp/answer.txt`." - ), - }, - { - "task_id": "math-sum", - "answer": "42", - "instruction": "Create `/tmp/answer.txt` containing the sum of 19 and 23.", - }, - { - "task_id": "sorted-words", - "answer": "alpha,beta,gamma", - "instruction": ( - "Create `/tmp/answer.txt` containing alpha, beta, and gamma in " - "alphabetical order, separated by commas and no spaces." - ), - }, -] - - -PROGRAM_SANDBOX = { - "image": "python:3.11-slim", - "scope": "rollout", - "network_access": True, - "timeout_minutes": 20, - "command_timeout": 120, -} - - -class ParallelSandboxTasksetConfig(vf.TasksetConfig): - toolsets: dict[str, vf.ToolsetConfig] = { - "bash": vf.ToolsetConfig( - tools=["bash"], - write=True, - sandbox="program", - ) - } - updates: list[str] = ["parallel_sandbox_audit"] - rewards: list[str] = ["sandbox_stage_score"] - metrics: list[str] = ["bash_calls", "update_audits"] - cleanups: list[str] = ["collect_program_sandbox_commands"] - system_prompt: str = SYSTEM_PROMPT - num_examples: int = -1 - - -class ParallelSandboxHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(sandbox=True, channels="callable") - sandbox: vf.SandboxConfig = vf.SandboxConfig(**PROGRAM_SANDBOX) - max_turns: int = 4 - - -async def bash(command: str, sandbox, state) -> str: - """Run a bash command in the active program sandbox.""" - result = await sandbox.execute(command, timeout=120, working_dir="/tmp") - output = { - "exit_code": int(getattr(result, "exit_code", 0)), - "stdout": truncate_text(str(getattr(result, "stdout", "") or "")), - "stderr": truncate_text(str(getattr(result, "stderr", "") or "")), - } - state.setdefault("bash_tool_outputs", []).append(output) - return json.dumps(output, ensure_ascii=False) - - -@vf.update(priority=10) -async def parallel_sandbox_audit(task, state) -> None: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - audit_specs = [ - ( - "file_audit", - FILE_AUDIT_SYSTEM_PROMPT, - file_audit_prompt(task, response), - ), - ( - "command_audit", - COMMAND_AUDIT_SYSTEM_PROMPT, - command_audit_prompt(task), - ), - ] - - async def run_audit( - label: str, system_prompt: str, prompt: str - ) -> tuple[str, vf.State]: - audit_task = vf.Task( - { - "prompt": [{"role": "user", "content": prompt}], - "max_turns": 2, - } - ).freeze() - audit_state = state.for_task( - audit_task, - borrow=["model", "sandbox"], - tools="bash", - transcript="append", - ) - audit_state = await vf.Harness( - config=vf.HarnessConfig( - system_prompt=system_prompt, - max_turns=2, - ) - ).run(audit_task, audit_state) - return label, audit_state - - audit_states = await asyncio.gather( - *( - run_audit(label, system_prompt, prompt) - for label, system_prompt, prompt in audit_specs - ) - ) - state["parallel_audits"] = [] - for label, audit_state in audit_states: - messages = vf.get_messages( - audit_state.get("completion") or [], role="assistant" - ) - findings = str(messages[-1].content or "") if messages else "" - state["parallel_audits"].append( - { - "name": label, - "findings": findings, - "trajectory_id": audit_state.get("trajectory_id"), - } - ) - - -@vf.reward(weight=1.0) -async def sandbox_stage_score(task, state) -> float: - judge_task = vf.Task( - { - "prompt": [ - { - "role": "user", - "content": reward_prompt(task, state), - } - ], - "max_turns": 2, - } - ).freeze() - judge_state = state.for_task(judge_task, borrow=["model", "sandbox"], tools="bash") - judge_state = await vf.Harness( - config=vf.HarnessConfig( - system_prompt=REWARD_JUDGE_SYSTEM_PROMPT, - max_turns=2, - ) - ).run(judge_task, judge_state) - messages = vf.get_messages(judge_state.get("completion") or [], role="assistant") - judge_text = str(messages[-1].content or "") if messages else "" - parsed = parse_judge_json(judge_text) - score = clamp_float(parsed.get("score", 0.0)) - state["reward_judge"] = { - "score": score, - "reason": str(parsed.get("reason", "")), - "raw": judge_text, - } - return score - - -@vf.cleanup(priority=10) -async def collect_program_sandbox_commands(task, state) -> None: - _ = task - state["program_sandbox_commands"] = [ - truncate_command_record(record) for record in state.get("sandbox_commands", []) - ] - state.pop("sandbox_commands", None) - - -@vf.metric(priority=-10) -async def bash_calls(task, state) -> float: - _ = task - return float(len(state.get("sandbox_commands", []))) - - -@vf.metric -async def update_audits(task, state) -> float: - _ = task - audits = state.get("parallel_audits", []) - return float(len(audits) if isinstance(audits, list) else 0) - - -def file_audit_prompt(task: vf.Task, response: str) -> str: - return ( - "Task instruction:\n" - f"{task['instruction']}\n\n" - "Expected answer text:\n" - f"{task['answer']}\n\n" - "Assistant final answer:\n" - f"{response}\n\n" - "Call bash to inspect the sandbox. A good command is:\n" - "python - <<'PY'\n" - "from pathlib import Path\n" - "import json\n" - "paths = ['/tmp/answer.txt', '/tmp/worklog.md']\n" - "print(json.dumps({p: {\n" - " 'exists': Path(p).exists(),\n" - " 'text': Path(p).read_text(errors='replace') if Path(p).exists() else '',\n" - "} for p in paths}))\n" - "PY\n" - ) - - -def command_audit_prompt(task: vf.Task) -> str: - return ( - "Task instruction:\n" - f"{task['instruction']}\n\n" - "Call bash to inspect the sandbox file layout and metadata. A good " - "command is:\n" - "python - <<'PY'\n" - "from pathlib import Path\n" - "import json\n" - "payload = {\n" - " 'tmp_files': sorted(str(p) for p in Path('/tmp').glob('*')),\n" - " 'answer_size': Path('/tmp/answer.txt').stat().st_size if Path('/tmp/answer.txt').exists() else None,\n" - " 'worklog_size': Path('/tmp/worklog.md').stat().st_size if Path('/tmp/worklog.md').exists() else None,\n" - "}\n" - "print(json.dumps(payload))\n" - "PY\n" - ) - - -def reward_prompt(task: vf.Task, state: vf.State) -> str: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - return ( - "Task instruction:\n" - f"{task['instruction']}\n\n" - "Expected answer text:\n" - f"{task['answer']}\n\n" - "Assistant final answer:\n" - f"{response}\n\n" - "Update-stage audit findings:\n" - f"{json.dumps(state.get('parallel_audits', []), indent=2)}\n\n" - "Call bash to inspect `/tmp/answer.txt` directly, then score whether " - "the sandbox state and final answer satisfy the task." - ) - - -def load_tasks(num_examples: int = -1): - rows = TASKS if num_examples < 0 else TASKS[:num_examples] - for index, row in enumerate(rows): - yield { - **row, - "example_id": index, - "prompt": [ - { - "role": "user", - "content": ( - f"{row['instruction']}\n\n" - "Use the bash tool for the file operation, write " - "`/tmp/worklog.md`, then answer with the requested text only." - ), - } - ], - "max_turns": 4, - } - - -class ParallelSandboxTaskset(vf.Taskset[ParallelSandboxTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(num_examples=self.config.num_examples) - - -class ParallelSandboxHarness(vf.Harness[ParallelSandboxHarnessConfig]): - pass - - -class ParallelSandboxEnvConfig(vf.EnvConfig): - taskset: ParallelSandboxTasksetConfig = ParallelSandboxTasksetConfig() - harness: ParallelSandboxHarnessConfig = ParallelSandboxHarnessConfig() - - -def load_environment(config: ParallelSandboxEnvConfig) -> vf.Env: - return vf.Env( - taskset=ParallelSandboxTaskset(config=config.taskset), - harness=ParallelSandboxHarness(config=config.harness), - ) diff --git a/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/__init__.py b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/__init__.py new file mode 100644 index 0000000000..81b58042eb --- /dev/null +++ b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/__init__.py @@ -0,0 +1 @@ +"""hello-parallel-sandbox-v1 environment package.""" diff --git a/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/harness.py b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/harness.py new file mode 100644 index 0000000000..4afd1b79a5 --- /dev/null +++ b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/harness.py @@ -0,0 +1,2 @@ +from .taskset import ParallelSandboxHarnessConfig as ParallelSandboxHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/taskset.py b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/taskset.py new file mode 100644 index 0000000000..39b9ed9a58 --- /dev/null +++ b/environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1/taskset.py @@ -0,0 +1,118 @@ +import asyncio + +import verifiers.v1 as vf + +SYSTEM_PROMPT = """Reply with the requested answer text only.""" + +TASKS: list[vf.JsonData] = [ + { + "task_id": "exact-token", + "answer": "prime-v1-shared-sandbox", + "instruction": "Return exactly `prime-v1-shared-sandbox`.", + }, + { + "task_id": "reverse-token", + "answer": "xobdnas-derahs", + "instruction": "Return exactly the reverse of `shared-sandbox`.", + }, + { + "task_id": "joined-words", + "answer": "taskset-harness-runtime", + "instruction": "Return taskset, harness, and runtime joined by hyphens.", + }, +] + + +class ParallelSandboxTasksetConfig(vf.TasksetConfig): + system_prompt: str = SYSTEM_PROMPT + num_examples: int = -1 + + +class ParallelSandboxHarnessConfig(vf.HarnessConfig): + max_turns: int = 1 + + +class ParallelSandboxTask(vf.Task): + answer: str + instruction: str + + +class ParallelSandboxTaskset(vf.Taskset[ParallelSandboxTasksetConfig]): + task_type = ParallelSandboxTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + rows = ( + TASKS if self.config.num_examples < 0 else TASKS[: self.config.num_examples] + ) + return [ + { + **row, + "row_id": index, + "prompt": [{"role": "user", "content": str(row["instruction"])}], + "max_turns": 1, + } + for index, row in enumerate(rows) + ] + + @vf.update(priority=10) + async def parallel_audit(self, task: ParallelSandboxTask, state: vf.State) -> None: + response = assistant_text(state) + file_audit, command_audit = await asyncio.gather( + audit_exact_answer(task, response), + audit_shape(task, response), + ) + audits = [file_audit, command_audit] + state.extras["parallel_audits"] = audits + state.artifacts["parallel_audits"] = audits + + @vf.metric + async def update_audits(self, state: vf.State) -> float: + audits = state.extras.get("parallel_audits") + return float(len(audits) if isinstance(audits, list) else 0) + + @vf.reward(weight=1.0) + async def sandbox_stage_score(self, state: vf.State) -> float: + audits = state.extras.get("parallel_audits") + if not isinstance(audits, list) or not audits: + return 0.0 + passed = [ + bool(audit.get("passed")) for audit in audits if isinstance(audit, dict) + ] + return sum(float(item) for item in passed) / len(passed) + + +async def audit_exact_answer(task: ParallelSandboxTask, response: str) -> vf.JsonData: + await asyncio.sleep(0) + expected = task.answer + return { + "name": "exact_answer", + "passed": response.strip() == expected, + "expected": expected, + "observed": response.strip(), + } + + +async def audit_shape(task: ParallelSandboxTask, response: str) -> vf.JsonData: + await asyncio.sleep(0) + expected = task.answer + return { + "name": "shape", + "passed": bool(response.strip()) and "\n" not in response.strip(), + "expected_length": len(expected), + "observed_length": len(response.strip()), + } + + +def assistant_text(state: vf.State) -> str: + messages = [message for message in state.completion if message.role == "assistant"] + return str(messages[-1].content or "") if messages else "" + + +def load_taskset(config: ParallelSandboxTasksetConfig) -> ParallelSandboxTaskset: + return ParallelSandboxTaskset(config=config) + + +def load_harness(config: ParallelSandboxHarnessConfig) -> vf.Harness: + return vf.Harness(config=config) diff --git a/environments/hello_parallel_sandbox_v1/pyproject.toml b/environments/hello_parallel_sandbox_v1/pyproject.toml index 1ed5e9cbae..50b5e808c8 100644 --- a/environments/hello_parallel_sandbox_v1/pyproject.toml +++ b/environments/hello_parallel_sandbox_v1/pyproject.toml @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["hello_parallel_sandbox_v1.py", "README.md", "pyproject.toml"] +include = ["hello_parallel_sandbox_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/hello_rlm_v1/hello_rlm_v1.py b/environments/hello_rlm_v1/hello_rlm_v1.py deleted file mode 100644 index 55ee1971ac..0000000000 --- a/environments/hello_rlm_v1/hello_rlm_v1.py +++ /dev/null @@ -1,78 +0,0 @@ -import verifiers as vf -from harnesses import RLM, RLMConfig - - -@vf.reward(weight=1.0) -async def exact_answer(task, state) -> float: - stdout = str(state.get("command", {}).get("stdout") or "") - return float(str(task["answer"]).lower() in stdout.lower()) - - -def load_tasks(split: vf.TaskSplit = "train"): - _ = split - return [ - { - "question": "Reply with exactly hello rlm.", - "answer": "hello rlm", - }, - { - "question": "Reply with exactly taskset harness.", - "answer": "taskset harness", - }, - { - "question": "Reply with exactly runtime boundary.", - "answer": "runtime boundary", - }, - { - "question": "Reply with exactly sandbox lease.", - "answer": "sandbox lease", - }, - { - "question": "Reply with exactly toolset scope.", - "answer": "toolset scope", - }, - { - "question": "Reply with exactly group reward.", - "answer": "group reward", - }, - { - "question": "Reply with exactly endpoint proxy.", - "answer": "endpoint proxy", - }, - { - "question": "Reply with exactly cleanup signal.", - "answer": "cleanup signal", - }, - { - "question": "Reply with exactly harbor task.", - "answer": "harbor task", - }, - { - "question": "Reply with exactly recursive model.", - "answer": "recursive model", - }, - ] - - -class HelloRLMTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["exact_answer"] - - -class HelloRLMTaskset(vf.Taskset[HelloRLMTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -def load_taskset(config: HelloRLMTasksetConfig) -> HelloRLMTaskset: - return HelloRLMTaskset(config=config) - - -def load_harness(config: RLMConfig) -> RLM: - return RLM(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/hello_rlm_v1/hello_rlm_v1/__init__.py b/environments/hello_rlm_v1/hello_rlm_v1/__init__.py new file mode 100644 index 0000000000..3d905ab1d5 --- /dev/null +++ b/environments/hello_rlm_v1/hello_rlm_v1/__init__.py @@ -0,0 +1 @@ +"""hello-rlm-v1 environment package.""" diff --git a/environments/hello_rlm_v1/hello_rlm_v1/harness.py b/environments/hello_rlm_v1/hello_rlm_v1/harness.py new file mode 100644 index 0000000000..09d2d98091 --- /dev/null +++ b/environments/hello_rlm_v1/hello_rlm_v1/harness.py @@ -0,0 +1,2 @@ +from .taskset import HelloRLMHarnessConfig as HelloRLMHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/hello_rlm_v1/hello_rlm_v1/taskset.py b/environments/hello_rlm_v1/hello_rlm_v1/taskset.py new file mode 100644 index 0000000000..ebaffae58a --- /dev/null +++ b/environments/hello_rlm_v1/hello_rlm_v1/taskset.py @@ -0,0 +1,53 @@ +import verifiers.v1 as vf +from harnesses import RLM, RLMConfig + + +def load_tasks(split: vf.TaskSplit = "train"): + _ = split + return [ + { + "prompt": "Reply with exactly hello rlm.", + "answer": "hello rlm", + }, + { + "prompt": "Reply with exactly taskset harness.", + "answer": "taskset harness", + }, + { + "prompt": "Reply with exactly runtime boundary.", + "answer": "runtime boundary", + }, + ] + + +class HelloRLMTask(vf.Task): + answer: str + + +class HelloRLMTasksetConfig(vf.TasksetConfig): + pass + + +class HelloRLMTaskset(vf.Taskset[HelloRLMTasksetConfig]): + task_type = HelloRLMTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return [HelloRLMTask.model_validate(record) for record in load_tasks(split)] + + @vf.reward(weight=1.0) + async def exact_answer(self, task: HelloRLMTask, state: vf.State) -> float: + command = state.artifacts.get("command") + stdout = str(command.get("stdout") if isinstance(command, dict) else "") + return float(task.answer.lower() in stdout.lower()) + + +class HelloRLMHarnessConfig(RLMConfig): + cwd: str | None = None + + +def load_taskset(config: HelloRLMTasksetConfig) -> HelloRLMTaskset: + return HelloRLMTaskset(config=config) + + +def load_harness(config: HelloRLMHarnessConfig) -> RLM: + return RLM(config=config) diff --git a/environments/hello_rlm_v1/pyproject.toml b/environments/hello_rlm_v1/pyproject.toml index 9c05bb1a93..a00f2962b0 100644 --- a/environments/hello_rlm_v1/pyproject.toml +++ b/environments/hello_rlm_v1/pyproject.toml @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["hello_rlm_v1.py", "pyproject.toml"] +include = ["hello_rlm_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/hello_self_judge_v1/README.md b/environments/hello_self_judge_v1/README.md index ccd7ca7d9a..6734ab2038 100644 --- a/environments/hello_self_judge_v1/README.md +++ b/environments/hello_self_judge_v1/README.md @@ -1,20 +1,14 @@ # hello-self-judge-v1 -V1 example where the answer rollout and the judge rollout share live runtime -resources. +V1 example where the answer rollout is reviewed by taskset-owned update logic. -The answer harness runs the base loop locally. The taskset contributes a -rollout-scoped, sandbox-backed `bash` tool. Each task asks the model to fetch -public web pages, write `/tmp/evidence.md`, and answer from the evidence. The -taskset then: +The answer harness runs the base loop. Each task asks the model to answer with a +sources line. The taskset then: -1. runs an update-stage judge harness that borrows the live `model` and - `bash` tool, appends to the public trajectory, and uses the same - tool-owned sandbox to inspect `/tmp/evidence.md`; -2. stores the update judge's findings under `state["update_judge"]`; -3. runs a reward-stage judge harness that borrows only the live `model`, keeps - its trajectory private, stores JSON under `state["judge"]`, and returns its - score. +1. stores judge findings under `state.extras["judge"]` and + `state.artifacts["judge_findings"]`; +2. reports source-mention metrics; +3. computes the reward from the serialized findings. ```bash prime eval run hello-self-judge-v1 -m openai/gpt-5.4-mini -n 3 -r 1 -t 4096 diff --git a/environments/hello_self_judge_v1/hello_self_judge_v1.py b/environments/hello_self_judge_v1/hello_self_judge_v1.py deleted file mode 100644 index 8ed347957e..0000000000 --- a/environments/hello_self_judge_v1/hello_self_judge_v1.py +++ /dev/null @@ -1,365 +0,0 @@ -import json - -import verifiers as vf -from verifiers.v1.utils.judge_utils import ( - clamp_float, - parse_judge_json, - truncate_command_record, - truncate_text, -) - -SYSTEM_PROMPT = """You are a web evidence assistant running in an isolated sandbox. - -Use the bash tool to fetch the requested public web pages and inspect their -contents. Before your final answer, write `/tmp/evidence.md` with: - -- the source URL or URLs you used; -- short copied excerpts or facts from those sources; -- notes explaining how the evidence supports your answer. - -Then give a concise answer with a `Sources:` line. Do not answer from memory -alone. -""" - -UPDATE_JUDGE_SYSTEM_PROMPT = """You are a strict evidence reviewer. - -Review whether the assistant answer is supported by the sandbox evidence. You -have a bash tool connected to the same sandbox used by the answering model. - -First call bash to inspect `/tmp/evidence.md` and any other relevant sandbox -state. Then write plain-language findings for a later scoring judge. Mention -strengths, gaps, and any contradiction. Do not assign a numeric score. -""" - -REWARD_JUDGE_SYSTEM_PROMPT = """You convert evidence-review findings into a score. - -Grade whether the answer is supported by the sandbox evidence, not whether it -matches a hidden ground truth. Respond with compact JSON only: - -{"score": 0.0-1.0, "reason": "..."} - -Use these criteria: -- 1.0: answer is clearly supported by fetched evidence and cites sources; -- 0.7: mostly supported but missing a detail or source citation; -- 0.4: plausible but weakly supported by the sandbox state; -- 0.0: no meaningful evidence, no answer, or answer contradicts evidence. -""" - - -TASKS: list[vf.ConfigData] = [ - { - "task_id": "example-domains", - "question": ( - "Fetch https://www.iana.org/domains/reserved and explain what the " - "example domains are reserved for." - ), - "seed_urls": ["https://www.iana.org/domains/reserved"], - }, - { - "task_id": "rfc-9110-404", - "question": ( - "Fetch https://www.rfc-editor.org/rfc/rfc9110.txt and explain what " - "HTTP status code 404 means." - ), - "seed_urls": ["https://www.rfc-editor.org/rfc/rfc9110.txt"], - }, - { - "task_id": "python-venv", - "question": ( - "Fetch the Python venv documentation at " - "https://docs.python.org/3/library/venv.html and summarize what a " - "virtual environment is used for." - ), - "seed_urls": ["https://docs.python.org/3/library/venv.html"], - }, - { - "task_id": "robots-txt", - "question": ( - "Fetch https://www.wikipedia.org/robots.txt and report one rule or " - "section that appears in the file." - ), - "seed_urls": ["https://www.wikipedia.org/robots.txt"], - }, - { - "task_id": "gnu-gpl", - "question": ( - "Fetch https://www.gnu.org/licenses/gpl-3.0.txt and summarize one " - "permission and one condition from the GPLv3 license text." - ), - "seed_urls": ["https://www.gnu.org/licenses/gpl-3.0.txt"], - }, - { - "task_id": "python-json", - "question": ( - "Fetch https://docs.python.org/3/library/json.html and summarize " - "what json.dumps does." - ), - "seed_urls": ["https://docs.python.org/3/library/json.html"], - }, - { - "task_id": "iana-time-zones", - "question": ( - "Fetch https://www.iana.org/time-zones and explain what kind of " - "resource the Time Zone Database is." - ), - "seed_urls": ["https://www.iana.org/time-zones"], - }, - { - "task_id": "mozilla-http", - "question": ( - "Fetch https://developer.mozilla.org/en-US/docs/Web/HTTP/Guides/" - "Overview and summarize what HTTP is used for." - ), - "seed_urls": [ - "https://developer.mozilla.org/en-US/docs/Web/HTTP/Guides/Overview" - ], - }, - { - "task_id": "w3c-html", - "question": ( - "Fetch https://www.w3.org/TR/html52/introduction.html and summarize " - "what HTML is for." - ), - "seed_urls": ["https://www.w3.org/TR/html52/introduction.html"], - }, - { - "task_id": "sqlite-about", - "question": ( - "Fetch https://www.sqlite.org/about.html and summarize two design " - "properties SQLite claims about itself." - ), - "seed_urls": ["https://www.sqlite.org/about.html"], - }, - { - "task_id": "iana-root-zone", - "question": ( - "Fetch https://www.iana.org/domains/root and explain what the root " - "zone database contains." - ), - "seed_urls": ["https://www.iana.org/domains/root"], - }, - { - "task_id": "python-pathlib", - "question": ( - "Fetch https://docs.python.org/3/library/pathlib.html and summarize " - "why someone would use pathlib." - ), - "seed_urls": ["https://docs.python.org/3/library/pathlib.html"], - }, -] - - -class SelfJudgeTasksetConfig(vf.TasksetConfig): - toolsets: dict[str, dict[str, str]] = {"bash": {"fn": "load_bash_toolset"}} - updates: list[str] = ["sandbox_judge"] - rewards: list[str] = ["self_consistency_score"] - metrics: list[str] = ["bash_calls"] - system_prompt: str = SYSTEM_PROMPT - num_examples: int = -1 - - -class SelfJudgeHarnessConfig(vf.HarnessConfig): - max_turns: int = 8 - - -async def bash(command: str, sandbox, state) -> str: - """Run a bash command in the rollout sandbox and return stdout/stderr.""" - result = await sandbox.execute(command, timeout=120, working_dir="/tmp") - output = { - "exit_code": int(getattr(result, "exit_code", 0)), - "stdout": truncate_text(str(getattr(result, "stdout", "") or "")), - "stderr": truncate_text(str(getattr(result, "stderr", "") or "")), - } - state.setdefault("bash_tool_outputs", []).append(output) - return json.dumps(output, ensure_ascii=False) - - -@vf.cleanup(priority=10) -async def collect_bash_commands(task, state) -> None: - _ = task - state["bash_commands"] = [ - truncate_command_record(record) for record in state.get("sandbox_commands", []) - ] - state.pop("sandbox_commands", None) - - -@vf.metric -async def bash_calls(task, state) -> float: - _ = task - return float(len(state.get("bash_tool_outputs", []))) - - -@vf.reward(weight=1.0) -async def self_consistency_score(task, state) -> float: - updated = state.get("update_judge") - if not isinstance(updated, dict): - return 0.0 - findings = str(updated.get("findings") or "") - if not findings: - return 0.0 - - judge_task = vf.Task( - { - "prompt": [ - { - "role": "user", - "content": score_prompt(task, findings), - } - ], - "max_turns": 1, - } - ).freeze() - judge_state = state.for_task(judge_task, borrow="model") - judge_state = await vf.Harness( - config=vf.HarnessConfig( - system_prompt=REWARD_JUDGE_SYSTEM_PROMPT, - max_turns=1, - ) - ).run(judge_task, judge_state) - - messages = vf.get_messages(judge_state.get("completion") or [], role="assistant") - judge_text = str(messages[-1].content or "") if messages else "" - parsed = parse_judge_json(judge_text) - score = clamp_float(parsed.get("score", 0.0)) - state["judge"] = { - "score": score, - "reason": str(parsed.get("reason", "")), - "raw": judge_text, - } - return score - - -@vf.update(priority=10) -async def sandbox_judge(task, state) -> None: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - judge_task = vf.Task( - { - "prompt": [ - { - "role": "user", - "content": update_prompt(task, response), - } - ], - "max_turns": 3, - } - ).freeze() - judge_state = state.for_task( - judge_task, - borrow="model", - tools="bash", - transcript="append", - ) - bash_output_start = len(state.get("bash_tool_outputs", [])) - judge_state = await vf.Harness( - config=vf.HarnessConfig( - system_prompt=UPDATE_JUDGE_SYSTEM_PROMPT, - max_turns=3, - ) - ).run(judge_task, judge_state) - judge_bash_outputs = state.get("bash_tool_outputs", [])[bash_output_start:] - - messages = vf.get_messages(judge_state.get("completion") or [], role="assistant") - findings = str(messages[-1].content or "") if messages else "" - state["update_judge"] = { - "findings": findings, - "trajectory_id": judge_state["trajectory_id"], - "bash_calls": len(judge_bash_outputs), - } - state["sandbox_report"] = judge_bash_outputs - - -def update_prompt(task: vf.Task, response: str) -> str: - return ( - "Task:\n" - f"{task['question']}\n\n" - "Expected seed URLs:\n" - f"{json.dumps(task.get('seed_urls', []), indent=2)}\n\n" - "Assistant answer:\n" - f"{response}\n\n" - "Call bash before writing findings. A useful inspection command is:\n" - "python - <<'PY'\n" - "from pathlib import Path\n" - "import json\n" - "path = Path('/tmp/evidence.md')\n" - "payload = {\n" - " 'evidence_exists': path.exists(),\n" - " 'evidence_bytes': path.stat().st_size if path.exists() else 0,\n" - " 'evidence_preview': path.read_text(errors='replace')[:6000] if path.exists() else '',\n" - " 'tmp_files': sorted(str(p) for p in Path('/tmp').glob('*'))[:100],\n" - "}\n" - "print(json.dumps(payload))\n" - "PY\n\n" - "After inspecting the sandbox, write findings in words. Do not output " - "JSON and do not assign a score." - ) - - -def score_prompt(task: vf.Task, findings: str) -> str: - return ( - "Task:\n" - f"{task['question']}\n\n" - "Evidence-review findings from the update stage:\n" - f"{findings}\n\n" - "Convert the findings to a calibrated JSON score. You cannot inspect " - "the sandbox directly; score only from the findings above." - ) - - -def load_tasks(num_examples: int = -1): - rows = TASKS if num_examples < 0 else TASKS[:num_examples] - for index, row in enumerate(rows): - question = str(row["question"]) - yield { - **row, - "example_id": index, - "prompt": [ - { - "role": "user", - "content": ( - f"{question}\n\n" - "Use bash to fetch the source material. Save your " - "evidence and notes to `/tmp/evidence.md` before " - "answering." - ), - } - ], - "max_turns": 8, - } - - -def load_bash_toolset() -> vf.Toolset: - return vf.Toolset( - tools=[bash], - write=True, - scope="rollout", - sandbox=vf.SandboxConfig( - image="python:3.11-slim", - scope="rollout", - network_access=True, - timeout_minutes=30, - command_timeout=120, - ), - cleanups=[collect_bash_commands], - ) - - -class SelfJudgeTaskset(vf.Taskset[SelfJudgeTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(num_examples=self.config.num_examples) - - -class SelfJudgeHarness(vf.Harness[SelfJudgeHarnessConfig]): - pass - - -class SelfJudgeEnvConfig(vf.EnvConfig): - taskset: SelfJudgeTasksetConfig = SelfJudgeTasksetConfig() - harness: SelfJudgeHarnessConfig = SelfJudgeHarnessConfig() - - -def load_environment(config: SelfJudgeEnvConfig) -> vf.Env: - return vf.Env( - taskset=SelfJudgeTaskset(config=config.taskset), - harness=SelfJudgeHarness(config=config.harness), - ) diff --git a/environments/hello_self_judge_v1/hello_self_judge_v1/__init__.py b/environments/hello_self_judge_v1/hello_self_judge_v1/__init__.py new file mode 100644 index 0000000000..da5732bc9b --- /dev/null +++ b/environments/hello_self_judge_v1/hello_self_judge_v1/__init__.py @@ -0,0 +1 @@ +"""hello-self-judge-v1 environment package.""" diff --git a/environments/hello_self_judge_v1/hello_self_judge_v1/harness.py b/environments/hello_self_judge_v1/hello_self_judge_v1/harness.py new file mode 100644 index 0000000000..cc066d64f9 --- /dev/null +++ b/environments/hello_self_judge_v1/hello_self_judge_v1/harness.py @@ -0,0 +1,2 @@ +from .taskset import SelfJudgeHarnessConfig as SelfJudgeHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/hello_self_judge_v1/hello_self_judge_v1/taskset.py b/environments/hello_self_judge_v1/hello_self_judge_v1/taskset.py new file mode 100644 index 0000000000..7d0170fb20 --- /dev/null +++ b/environments/hello_self_judge_v1/hello_self_judge_v1/taskset.py @@ -0,0 +1,104 @@ +import verifiers.v1 as vf + +SYSTEM_PROMPT = """Answer the question concisely and include a Sources: line.""" + +TASKS: list[vf.JsonData] = [ + { + "task_id": "example-domains", + "question": "Explain what the example domains are reserved for.", + "seed_urls": ["https://www.iana.org/domains/reserved"], + "answer_hint": "reserved for use in documentation and examples", + }, + { + "task_id": "rfc-9110-404", + "question": "Explain what HTTP status code 404 means.", + "seed_urls": ["https://www.rfc-editor.org/rfc/rfc9110.txt"], + "answer_hint": "target resource was not found", + }, + { + "task_id": "python-json", + "question": "Summarize what json.dumps does.", + "seed_urls": ["https://docs.python.org/3/library/json.html"], + "answer_hint": "serializes an object to a JSON formatted string", + }, +] + + +class SelfJudgeTasksetConfig(vf.TasksetConfig): + system_prompt: str = SYSTEM_PROMPT + num_examples: int = -1 + + +class SelfJudgeHarnessConfig(vf.HarnessConfig): + max_turns: int = 1 + + +class SelfJudgeTask(vf.Task): + question: str + seed_urls: list[str] + answer_hint: str + + +class SelfJudgeTaskset(vf.Taskset[SelfJudgeTasksetConfig]): + task_type = SelfJudgeTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + rows = ( + TASKS if self.config.num_examples < 0 else TASKS[: self.config.num_examples] + ) + return [ + { + **row, + "row_id": index, + "prompt": [{"role": "user", "content": str(row["question"])}], + "max_turns": 1, + } + for index, row in enumerate(rows) + ] + + @vf.update(priority=10) + async def evidence_review(self, task: SelfJudgeTask, state: vf.State) -> None: + response = assistant_text(state) + urls = [str(url) for url in task.seed_urls] + findings = { + "has_answer": bool(response.strip()), + "mentions_source": any(url in response for url in urls), + "mentions_sources_line": "sources:" in response.lower(), + "mentions_hint": task.answer_hint.lower() in response.lower(), + } + state.extras["judge"] = findings + state.artifacts["judge_findings"] = findings + + @vf.metric + async def source_mentions(self, state: vf.State) -> float: + judge = state.extras.get("judge") + if not isinstance(judge, dict): + return 0.0 + return float(bool(judge.get("mentions_source"))) + + @vf.reward(weight=1.0) + async def self_consistency_score(self, state: vf.State) -> float: + judge = state.extras.get("judge") + if not isinstance(judge, dict): + return 0.0 + checks = [ + bool(judge.get("has_answer")), + bool(judge.get("mentions_sources_line")), + bool(judge.get("mentions_hint")), + ] + return sum(float(check) for check in checks) / len(checks) + + +def assistant_text(state: vf.State) -> str: + messages = [message for message in state.completion if message.role == "assistant"] + return str(messages[-1].content or "") if messages else "" + + +def load_taskset(config: SelfJudgeTasksetConfig) -> SelfJudgeTaskset: + return SelfJudgeTaskset(config=config) + + +def load_harness(config: SelfJudgeHarnessConfig) -> vf.Harness: + return vf.Harness(config=config) diff --git a/environments/hello_self_judge_v1/pyproject.toml b/environments/hello_self_judge_v1/pyproject.toml index 03ad9627fe..65be6ce442 100644 --- a/environments/hello_self_judge_v1/pyproject.toml +++ b/environments/hello_self_judge_v1/pyproject.toml @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["hello_self_judge_v1.py", "README.md", "pyproject.toml"] +include = ["hello_self_judge_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/hello_subagent_v1/hello_subagent_v1.py b/environments/hello_subagent_v1/hello_subagent_v1.py deleted file mode 100644 index 4d26002507..0000000000 --- a/environments/hello_subagent_v1/hello_subagent_v1.py +++ /dev/null @@ -1,110 +0,0 @@ -import verifiers as vf - - -async def child_program( - task: vf.Task, state: vf.State -) -> dict[str, list[dict[str, str]]]: - _ = state - name = str(task["name"]) - return {"completion": [{"role": "assistant", "content": f"hello {name}"}]} - - -class ChildHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="child_program") - - -async def ask_subagent(name: str, state) -> str: - """Ask a child harness to produce the greeting for one name.""" - harness = vf.Harness(config=ChildHarnessConfig()) - task = vf.Task( - { - "name": name, - "system_prompt": ( - "You are a child subagent. Reply with exactly " - f"`hello {name}` and no extra text." - ), - "prompt": [ - {"role": "user", "content": f"Say hello to {name}."}, - ], - } - ).freeze() - child_state = state.for_task(task, borrow="model") - child_state = await harness.run(task, child_state) - messages = vf.get_messages(child_state.get("completion") or [], role="assistant") - answer = str(messages[-1].content or "").strip() if messages else "" - state.setdefault("subagent_calls", []).append({"name": name, "answer": answer}) - return answer - - -@vf.metric -async def subagent_calls(task, state) -> float: - return float(len(state.get("subagent_calls", []))) - - -@vf.reward(weight=1.0) -async def exact_answer(task, state) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - answer = str(messages[-1].content or "").strip() if messages else "" - return float(answer == task["answer"]) - - -NAME_GROUPS = [ - ["world"], - ["prime", "verifiers"], - ["taskset", "harness", "runtime"], - ["sandbox"], - ["alpha", "beta"], - ["delta", "epsilon", "zeta"], - ["tools", "users"], - ["group", "reward", "advantage"], - ["mcp", "search"], - ["open", "superintelligence", "stack"], -] - - -def load_tasks(split: vf.TaskSplit = "train"): - _ = split - return [ - { - "names": names, - "prompt": [{"role": "user", "content": f"Names: {', '.join(names)}"}], - "answer": ", ".join(f"hello {name}" for name in names), - } - for names in NAME_GROUPS - ] - - -class SubagentTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["exact_answer"] - system_prompt: str = ( - "You are a parent coordinator. You must call ask_subagent once for " - "each requested name. After all tool results are available, join " - "the child answers with ', ' and output only that final joined text." - ) - - -class SubagentHarnessConfig(vf.HarnessConfig): - metrics: list[str] = ["subagent_calls"] - - -class SubagentTaskset(vf.Taskset[SubagentTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -class SubagentHarness(vf.Harness[SubagentHarnessConfig]): - def load_toolsets(self, config: SubagentHarnessConfig) -> vf.Toolsets: - _ = config - return {"subagent": vf.Toolset(tools=[ask_subagent], scope="rollout")} - - -class SubagentEnvConfig(vf.EnvConfig): - taskset: SubagentTasksetConfig = SubagentTasksetConfig() - harness: SubagentHarnessConfig = SubagentHarnessConfig() - - -def load_environment(config: SubagentEnvConfig) -> vf.Env: - return vf.Env( - taskset=SubagentTaskset(config=config.taskset), - harness=SubagentHarness(config=config.harness), - ) diff --git a/environments/hello_subagent_v1/hello_subagent_v1/__init__.py b/environments/hello_subagent_v1/hello_subagent_v1/__init__.py new file mode 100644 index 0000000000..6a55372d33 --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/__init__.py @@ -0,0 +1 @@ +"""hello-subagent-v1 environment package.""" diff --git a/environments/hello_subagent_v1/hello_subagent_v1/harness.py b/environments/hello_subagent_v1/hello_subagent_v1/harness.py new file mode 100644 index 0000000000..f97ef0848a --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/harness.py @@ -0,0 +1,2 @@ +from .taskset import SubagentHarnessConfig as SubagentHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/hello_subagent_v1/hello_subagent_v1/servers/__init__.py b/environments/hello_subagent_v1/hello_subagent_v1/servers/__init__.py new file mode 100644 index 0000000000..b75ee16e1b --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for hello-subagent-v1.""" diff --git a/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/__init__.py b/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/__init__.py new file mode 100644 index 0000000000..aa6efc028a --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/__init__.py @@ -0,0 +1,3 @@ +from .config import SubagentToolsetConfig + +__all__ = ["SubagentToolsetConfig"] diff --git a/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/config.py b/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/config.py new file mode 100644 index 0000000000..564d5d8590 --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/config.py @@ -0,0 +1,5 @@ +import verifiers.v1 as vf + + +class SubagentToolsetConfig(vf.ToolsetConfig): + pass diff --git a/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/toolset.py b/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/toolset.py new file mode 100644 index 0000000000..6cf3f25262 --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/servers/subagent/toolset.py @@ -0,0 +1,13 @@ +import verifiers.v1 as vf + +from .config import SubagentToolsetConfig + + +class SubagentToolset(vf.Toolset[SubagentToolsetConfig]): + @vf.tool(extends={"subagent_calls": "state.extras.subagent_calls"}) + def ask_subagent(self, name: str) -> dict: + answer = f"hello {name}" + return { + "content": answer, + "subagent_calls": [{"name": name, "answer": answer}], + } diff --git a/environments/hello_subagent_v1/hello_subagent_v1/taskset.py b/environments/hello_subagent_v1/hello_subagent_v1/taskset.py new file mode 100644 index 0000000000..435ce00c95 --- /dev/null +++ b/environments/hello_subagent_v1/hello_subagent_v1/taskset.py @@ -0,0 +1,71 @@ +import verifiers.v1 as vf + +from .servers.subagent import SubagentToolsetConfig + +NAME_GROUPS = [ + ["world"], + ["prime", "verifiers"], + ["taskset", "harness", "runtime"], + ["sandbox"], + ["alpha", "beta"], + ["delta", "epsilon", "zeta"], + ["tools", "users"], + ["group", "reward", "advantage"], + ["mcp", "search"], + ["open", "superintelligence", "stack"], +] + + +class SubagentTasksetConfig(vf.TasksetConfig): + system_prompt: str = ( + "You are a parent coordinator. Call subagent_ask_subagent once for each " + "requested name. After all tool results are available, join the child " + "answers with ', ' and output only that final joined text." + ) + toolsets: vf.ToolsetConfigs = {"subagent": SubagentToolsetConfig()} + + +class SubagentHarnessConfig(vf.HarnessConfig): + max_turns: int = 8 + + +class SubagentTask(vf.Task): + names: list[str] + answer: str + + +class SubagentTaskset(vf.Taskset[SubagentTasksetConfig]): + task_type = SubagentTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [ + { + "names": names, + "prompt": [{"role": "user", "content": f"Names: {', '.join(names)}"}], + "answer": ", ".join(f"hello {name}" for name in names), + } + for names in NAME_GROUPS + ] + + @vf.metric + async def subagent_calls(self, state: vf.State) -> float: + calls = state.extras.get("subagent_calls") + return float(len(calls) if isinstance(calls, list) else 0) + + @vf.reward(weight=1.0) + async def exact_answer(self, task: SubagentTask, state: vf.State) -> float: + messages = [ + message for message in state.completion if message.role == "assistant" + ] + answer = str(messages[-1].content or "").strip() if messages else "" + return float(answer == task.answer) + + +def load_taskset(config: SubagentTasksetConfig) -> SubagentTaskset: + return SubagentTaskset(config=config) + + +def load_harness(config: SubagentHarnessConfig) -> vf.Harness: + return vf.Harness(config=config) diff --git a/environments/hello_subagent_v1/pyproject.toml b/environments/hello_subagent_v1/pyproject.toml index 3dd8bde8d2..e2c689a456 100644 --- a/environments/hello_subagent_v1/pyproject.toml +++ b/environments/hello_subagent_v1/pyproject.toml @@ -14,7 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["hello_subagent_v1.py", "pyproject.toml"] +include = ["hello_subagent_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/langchain_deep_agents_wikispeedia/langchain_deep_agents_wikispeedia.py b/environments/langchain_deep_agents_wikispeedia/langchain_deep_agents_wikispeedia.py deleted file mode 100644 index 51bcfe25f1..0000000000 --- a/environments/langchain_deep_agents_wikispeedia/langchain_deep_agents_wikispeedia.py +++ /dev/null @@ -1,591 +0,0 @@ -import asyncio -import json -from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence -from typing import Protocol, cast - -from datasets import Dataset - -import verifiers as vf -from verifiers.v1.utils.prompt_utils import normalize_system_prompt - -if __package__: - from .wiki_graph import WikiGraph, WikiPair, load_wiki_graph -else: - from wiki_graph import WikiGraph, WikiPair, load_wiki_graph - - -class AgentMessage(Protocol): - role: str - content: object - - -def system_prompt(allow_go_back: bool = True) -> str: - backtracking = ( - "Use `go_back` to undo your last click." - if allow_go_back - else "Backtracking is disabled, so choose each link carefully." - ) - return f"""\ -This game is easy and fun: - -You are given two Wikipedia articles. Starting from the first article, your goal \ -is to reach the second one, exclusively by following links in the articles you \ -encounter. (For the articles you are given this is always possible.) - -Each article ends with a list of `Available links: ...` — those are the only \ -links you can follow. Use the `click_link` tool to navigate to one. \ -{backtracking} - -You also have access to deep-agent scaffolding tools (`write_todos`, \ -`write_file`, `read_file`, `ls`, `edit_file`, `task`). Use them when they help: \ -sketch a plan with `write_todos`, jot promising bridge concepts or dead-ends \ -in a file, and call `task` to spawn a focused sub-agent for a sub-search. They \ -are entirely optional. - -Try to be quick — think about which broader concepts connect the source to \ -the target, and aim for the article that most likely lists your destination \ -among its links. - -When you reach the target the system will say `TARGET REACHED`. Stop calling \ -tools at that point and reply with a brief confirmation.""" - - -SYSTEM_PROMPT = system_prompt() - - -class WikispeediaTasksetConfig(vf.TasksetConfig): - cache_dir: str | None = None - min_path_length: int = 3 - max_path_length: int = 6 - train_size: int = 50_000 - eval_size: int = 1_000 - eval_target_fraction: float = 0.1 - split_seed: int = 0 - links_only: bool = False - allow_go_back: bool = True - max_turns: int = 50 - efficiency_weight: float = 0.0 - stratify_path_length: bool = True - - -class WikispeediaHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig( - fn="run_langchain_deep_agents_wikispeedia_program" - ) - max_turns: int = 50 - timeout_seconds: float = 1200.0 - - -class WikispeediaTaskset(vf.Taskset[WikispeediaTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(self.config, split=split) - - -class WikispeediaHarness(vf.Harness[WikispeediaHarnessConfig]): - pass - - -def format_article(wiki: WikiGraph, article: str, links_only: bool = False) -> str: - links = wiki.get_links(article) - links_str = ", ".join(links) if links else "(no outgoing links)" - if links_only: - return f"# {article}\n\nAvailable links: {links_str}" - text = wiki.get_text(article) - return f"# {article}\n\n{text}\n\n---\nAvailable links: {links_str}" - - -async def click_link(article: str, wiki: WikiGraph, state: vf.State) -> str: - """Navigate to a linked Wikipedia article.""" - links_only = bool(state.get("links_only", False)) - current = state["current_article"] - available = wiki.get_links(current) - normalized = wiki.normalize_name(article) - if normalized is None or normalized not in available: - avail_str = ", ".join(available) if available else "(none)" - return ( - f"'{article}' is not a valid link from '{current}'.\n" - f"Available links: {avail_str}" - ) - state["current_article"] = normalized - state["path"].append(normalized) - if normalized == state["info"]["target"]: - state["reached_target"] = True - state.stop("target_reached") - return ( - f"TARGET REACHED: {normalized}\n\n" - "You successfully navigated to the target. Stop calling tools " - "and reply briefly to confirm." - ) - return format_article(wiki, normalized, links_only=links_only) - - -async def go_back(wiki: WikiGraph, state: vf.State) -> str: - """Undo the last click_link and return to the previous article.""" - path = state["path"] - if len(path) <= 1: - return "You are already at the starting article. Cannot go back." - path.pop() - state["current_article"] = path[-1] - return format_article( - wiki, path[-1], links_only=bool(state.get("links_only", False)) - ) - - -DEEP_AGENT_TOOLS = { - "write_todos", - "write_file", - "read_file", - "ls", - "edit_file", - "grep", - "task", -} -WIKISPEEDIA_TOOLS = {"click_link", "go_back"} - - -async def reached_target(task: vf.Task, state: vf.State) -> float: - return 1.0 if state.get("reached_target", False) else 0.0 - - -async def path_efficiency(task: vf.Task, state: vf.State) -> float: - if not state.get("reached_target", False): - return 0.0 - shortest = float(state["info"]["shortest_path"]) - actual = max(len(state.get("path", [])) - 1, 1) - return min(1.0, shortest / actual) - - -async def path_length(task: vf.Task, state: vf.State) -> float: - return float(max(len(state.get("path", [])) - 1, 0)) - - -async def shortest_path(task: vf.Task, state: vf.State) -> float: - return float(state.get("info", {}).get("shortest_path", 0)) - - -async def agent_timeout(task: vf.Task, state: vf.State) -> float: - return 1.0 if state.get("agent_timeout", False) else 0.0 - - -def iter_tool_calls(state: vf.State) -> Iterator[str]: - completion = state.get("completion") or [] - messages = ( - vf.get_messages(completion, role="assistant") - if isinstance(completion, list) - else [] - ) - for msg in messages: - tool_calls = msg.tool_calls - if not isinstance(tool_calls, list): - continue - for tool_call in tool_calls: - yield tool_call.name - - -def count_tool_calls(state: vf.State, name: str | None = None) -> int: - if name is None: - return sum(1 for _ in iter_tool_calls(state)) - return sum(1 for tool_name in iter_tool_calls(state) if tool_name == name) - - -def make_tool_count_metric( - name: str, -) -> Callable[[vf.Task, vf.State], Awaitable[float]]: - async def metric(task: vf.Task, state: vf.State) -> float: - return float(count_tool_calls(state, name)) - - metric.__name__ = f"calls_{name}" - return metric - - -def load_toolset( - cache_dir: str | None = None, - allow_go_back: bool = True, - config: vf.ToolsetConfig | None = None, -) -> vf.Toolset: - wiki_graph: WikiGraph | None = None - - def wiki() -> WikiGraph: - nonlocal wiki_graph - if wiki_graph is None: - wiki_graph = load_wiki_graph(cache_dir) - return wiki_graph - - async def click_link_tool(article: str, state: vf.State) -> str: - return await click_link(article, wiki(), state) - - click_link_tool.__name__ = "click_link" - click_link_tool.__doc__ = click_link.__doc__ - - tools: list[vf.Handler] = [click_link_tool] - if allow_go_back: - - async def go_back_tool(state: vf.State) -> str: - return await go_back(wiki(), state) - - go_back_tool.__name__ = "go_back" - go_back_tool.__doc__ = go_back.__doc__ - tools.append(go_back_tool) - return vf.Toolset( - tools=tools, - config=config, - ) - - -async def total_tool_calls(task: vf.Task, state: vf.State) -> float: - return float(count_tool_calls(state)) - - -async def assistant_turns(task: vf.Task, state: vf.State) -> float: - completion = state.get("completion") or [] - return float( - len(vf.get_messages(completion, role="assistant")) - if isinstance(completion, list) - else 0 - ) - - -async def invalid_link_rate(task: vf.Task, state: vf.State) -> float: - clicks = 0 - invalid = 0 - completion = state.get("completion") or [] - if not isinstance(completion, list): - return 0.0 - - transcript = vf.get_messages(completion) - id_to_name: dict[str, str] = {} - for msg in transcript: - if msg.role == "assistant": - tool_calls = msg.tool_calls - if tool_calls: - for tc in tool_calls: - id_to_name[tc.id] = tc.name - - for msg in transcript: - if msg.role != "tool": - continue - tool_name = id_to_name.get(msg.tool_call_id) - if tool_name is None: - extra = msg.get("name") - tool_name = extra if isinstance(extra, str) else None - if tool_name != "click_link": - continue - clicks += 1 - content = msg.content - if isinstance(content, str) and "is not a valid link" in content: - invalid += 1 - return float(invalid / clicks) if clicks else 0.0 - - -@vf.update(priority=-200) -async def restore_agent_completion(task: vf.Task, state: vf.State) -> None: - agent_completion = state.get("agent_completion") - if isinstance(agent_completion, list): - state["completion"] = agent_completion - - -def build_dataset( - wiki: WikiGraph, - pairs: list[WikiPair], - links_only: bool, - max_turns: int, -) -> Dataset: - records = [] - for source, target, dist in pairs: - starting = format_article(wiki, source, links_only=links_only) - prompt_text = ( - f"Your mission: {source} >> {target}\n\n" - f"Here is the starting article:\n\n{starting}" - ) - info: vf.ConfigData = { - "source": source, - "target": target, - "shortest_path": dist, - } - human = wiki.get_human_stats(source, target) - if human is not None: - info.update(human) - records.append( - { - "task_id": f"{source}->{target}", - "prompt": [{"role": "user", "content": prompt_text}], - "answer": target, - "info": info, - "links_only": links_only, - "max_turns": max_turns, - } - ) - return Dataset.from_list(records) - - -def split_pairs( - config: WikispeediaTasksetConfig, -) -> tuple[list[WikiPair], list[WikiPair]]: - return load_wiki_graph(config.cache_dir).split_pairs( - train_size=config.train_size, - eval_size=config.eval_size, - min_dist=config.min_path_length, - max_dist=config.max_path_length, - eval_target_fraction=config.eval_target_fraction, - seed=config.split_seed, - stratify=config.stratify_path_length, - ) - - -def load_tasks( - config: WikispeediaTasksetConfig, split: vf.TaskSplit = "train" -) -> Dataset: - train, eval_ = split_pairs(config) - return build_dataset( - load_wiki_graph(config.cache_dir), - train if split == "train" else eval_, - links_only=config.links_only, - max_turns=config.max_turns, - ) - - -def serialize_agent_completion( - messages: Sequence[AgentMessage | vf.JsonData], -) -> list[vf.ConfigData]: - role_aliases = { - "human": "user", - "ai": "assistant", - "tool": "tool", - "system": "system", - } - call_names: dict[str, str] = {} - serialized: list[vf.ConfigData] = [] - for message in messages: - if isinstance(message, Mapping): - payload = dict(message) - else: - model_dump = getattr(message, "model_dump", None) - payload = ( - model_dump(mode="json", exclude_none=True) - if callable(model_dump) - else { - "role": getattr(message, "role", None) - or getattr(message, "type", "assistant"), - "content": getattr(message, "content", str(message)), - "name": getattr(message, "name", None), - "tool_call_id": getattr(message, "tool_call_id", None), - "tool_calls": getattr(message, "tool_calls", None), - } - ) - raw_role = payload.get("role") or payload.get("type") or "assistant" - role = role_aliases.get(str(raw_role), str(raw_role)) - item: vf.ConfigData = { - "role": role, - "content": payload.get("content", ""), - } - tool_calls = payload.get("tool_calls") - if isinstance(tool_calls, list) and tool_calls: - normalized_tool_calls = [] - for tool_call in tool_calls: - if not isinstance(tool_call, Mapping): - continue - tool_call_payload = dict(tool_call) - name = tool_call_payload.get("name") - tool_id = tool_call_payload.get("id") or tool_call_payload.get( - "tool_call_id" - ) - if isinstance(tool_id, str) and isinstance(name, str): - call_names[tool_id] = name - arguments = tool_call_payload.get("arguments") - if not isinstance(arguments, str): - args = tool_call_payload.get("args", {}) - try: - arguments = json.dumps(args if args is not None else {}) - except (TypeError, ValueError): - arguments = str(args) - tool_call_payload["arguments"] = arguments - normalized_tool_calls.append(tool_call_payload) - item["tool_calls"] = normalized_tool_calls - name = payload.get("name") - if isinstance(name, str): - item["name"] = name - tool_call_id = payload.get("tool_call_id") - if isinstance(tool_call_id, str): - item["tool_call_id"] = tool_call_id - if item["role"] == "tool" and "name" not in item: - name = call_names.get(tool_call_id) - if name is not None: - item["name"] = name - serialized.append(item) - if serialized and serialized[0].get("role") == "user": - return serialized[1:] - return serialized - - -def langchain_navigation_tools(runtime_tools): - from langchain_core.tools import tool - - nav_tools = [] - if "click_link" in runtime_tools: - click_link_tool = runtime_tools["click_link"] - - @tool - async def click_link(article: str) -> str: - """Navigate to a linked Wikipedia article.""" - return str(await click_link_tool(article=article)) - - nav_tools.append(click_link) - if "go_back" in runtime_tools: - go_back_tool = runtime_tools["go_back"] - - @tool - async def go_back() -> str: - """Undo the last click_link and return to the previous article.""" - return str(await go_back_tool()) - - nav_tools.append(go_back) - return nav_tools - - -def make_langchain_deep_agents_program( - max_turns: int, - timeout_seconds: float, -) -> Callable[[vf.Task, vf.State], Awaitable[vf.State]]: - async def run_langchain_deep_agents_wikispeedia_program( - task: vf.Task, state: vf.State - ) -> vf.State: - from deepagents import create_deep_agent - from langchain_openai import ChatOpenAI - from langgraph.errors import GraphRecursionError - from openai import OpenAI - - state["current_article"] = state["info"]["source"] - state["path"] = [state["info"]["source"]] - state["reached_target"] = False - state["agent_timeout"] = False - state["links_only"] = bool(task.get("links_only", False)) - - endpoint_config = state.get_endpoint_config(api="chat") - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - endpoint_api_key = endpoint_client.api_key - endpoint_client.close() - model = ChatOpenAI( - model=endpoint_config.model, - base_url=endpoint_config.base_url, - api_key=endpoint_api_key, - ) - runtime_tools = state.get_tools() - nav_tools = langchain_navigation_tools(runtime_tools) - state_system_prompt = "" - system_prompt_messages = state.get("system_prompt") - if isinstance(system_prompt_messages, list): - state_system_prompt = "\n\n".join( - str(message.content or "") - for message in vf.get_messages(system_prompt_messages) - ) - agent = create_deep_agent( - model=model, - tools=nav_tools, - system_prompt=state_system_prompt or SYSTEM_PROMPT, - ) - prompt = str(cast(list[vf.ConfigData], state["prompt"])[-1]["content"]) - recursion_limit = state.get_max_turns(max_turns) - invoke_config = ( - {"recursion_limit": recursion_limit} if recursion_limit > 0 else None - ) - invoke = agent.ainvoke( - {"messages": [{"role": "user", "content": prompt}]}, - config=invoke_config, - ) - try: - result = await asyncio.wait_for(invoke, timeout=timeout_seconds) - except (TimeoutError, GraphRecursionError) as exc: - state["agent_timeout"] = True - state.stop( - "agent_timeout" - if isinstance(exc, TimeoutError) - else "agent_recursion_limit" - ) - state.setdefault("agent_completion", []) - return state - - messages = result.get("messages", []) if isinstance(result, Mapping) else [] - completion = serialize_agent_completion(messages) - state["agent_completion"] = completion - state["completion"] = completion - if completion: - state["agent_result"] = str(completion[-1].get("content") or "") - return state - - return run_langchain_deep_agents_wikispeedia_program - - -async def run_langchain_deep_agents_wikispeedia_program( - task: vf.Task, state: vf.State, harness: WikispeediaHarness -) -> vf.State: - return await make_langchain_deep_agents_program( - max_turns=harness.config.max_turns, - timeout_seconds=harness.config.timeout_seconds, - )(task, state) - - -def load_taskset( - config: WikispeediaTasksetConfig, -) -> WikispeediaTaskset: - rewards = [reached_target] - metrics = [ - path_length, - shortest_path, - agent_timeout, - total_tool_calls, - assistant_turns, - invalid_link_rate, - *[ - make_tool_count_metric(name) - for name in sorted(DEEP_AGENT_TOOLS | WIKISPEEDIA_TOOLS) - ], - ] - if config.efficiency_weight > 0: - - async def weighted_path_efficiency(task: vf.Task, state: vf.State) -> float: - return await path_efficiency(task, state) - - weighted_path_efficiency.__name__ = "path_efficiency" - rewards.append( - vf.reward(weight=config.efficiency_weight)(weighted_path_efficiency) - ) - else: - metrics.insert(0, path_efficiency) - - taskset = WikispeediaTaskset(config=config) - taskset.taskset_id = "langchain-deep-agents-wikispeedia" - taskset.system_prompt = normalize_system_prompt( - system_prompt(allow_go_back=config.allow_go_back), - field_name="taskset.system_prompt", - ) - taskset.add_toolset( - load_toolset( - cache_dir=config.cache_dir, - allow_go_back=config.allow_go_back, - ) - ) - for reward in rewards: - taskset.add_reward(reward) - for metric in metrics: - taskset.add_metric(metric) - return taskset - - -def load_harness( - config: WikispeediaHarnessConfig, -) -> WikispeediaHarness: - harness = WikispeediaHarness(config=config) - harness.add_update(restore_agent_completion) - return harness - - -class WikispeediaEnvConfig(vf.EnvConfig): - taskset: WikispeediaTasksetConfig = WikispeediaTasksetConfig() - harness: WikispeediaHarnessConfig = WikispeediaHarnessConfig() - - -def load_environment(config: WikispeediaEnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/langchain_deep_agents_wikispeedia/README.md b/environments/langchain_deep_agents_wikispeedia_v1/README.md similarity index 94% rename from environments/langchain_deep_agents_wikispeedia/README.md rename to environments/langchain_deep_agents_wikispeedia_v1/README.md index 3732f6d84f..e502918fb6 100644 --- a/environments/langchain_deep_agents_wikispeedia/README.md +++ b/environments/langchain_deep_agents_wikispeedia_v1/README.md @@ -1,9 +1,9 @@ -# langchain-deep-agents-wikispeedia +# langchain-deep-agents-wikispeedia-v1 LangChain deep-agents trained on Wikispeedia navigation through a v1 `Taskset`/`Harness`. ### Overview -- **Environment ID**: `langchain-deep-agents-wikispeedia` +- **Environment ID**: `langchain-deep-agents-wikispeedia-v1` - **Short description**: Multi-turn navigation through the Wikispeedia article graph with LangChain `create_deep_agent` (todos, virtual files, sub-agents) plus two task tools (`click_link`, `go_back`). - **Tags**: v1, taskset, harness, multi-turn, tool-use, langchain, deep-agents, wikispeedia, navigation @@ -22,12 +22,12 @@ LangChain deep-agents trained on Wikispeedia navigation through a v1 `Taskset`/` Run an evaluation with default settings: ```bash -prime eval run langchain-deep-agents-wikispeedia +prime eval run langchain-deep-agents-wikispeedia-v1 ``` Configure model and difficulty band: ```bash -prime eval run langchain-deep-agents-wikispeedia \ +prime eval run langchain-deep-agents-wikispeedia-v1 \ -m openai/gpt-4.1-mini \ -n 20 -r 3 -t 4096 -T 0.7 \ -a '{"config": {"taskset": {"min_path_length": 4, "max_path_length": 6, "max_turns": 40}}}' @@ -35,7 +35,7 @@ prime eval run langchain-deep-agents-wikispeedia \ Disable `go_back` (force planning over backtracking): ```bash -prime eval run langchain-deep-agents-wikispeedia \ +prime eval run langchain-deep-agents-wikispeedia-v1 \ -m openai/gpt-4.1-mini -n 20 -r 3 \ -a '{"config": {"taskset": {"allow_go_back": false}}}' ``` @@ -79,7 +79,7 @@ Notes: | `agent_timeout` | 1.0 if rollout hit `timeout_seconds` | | `calls_click_link`, `calls_go_back` | navigation tool counts (zero-weight) | | `calls_write_todos`, `calls_write_file`, `calls_read_file`, `calls_ls`, `calls_edit_file`, `calls_grep`, `calls_task` | deep-agent tool counts (zero-weight) | -| `total_tool_calls`, `assistant_turns` | trajectory shape (zero-weight) | +| `total_tool_calls`, `assistant_turns` | transcript shape (zero-weight) | | `invalid_link_rate` | fraction of `click_link` calls that named a non-existent link (hallucination canary, zero-weight) | ### Notes diff --git a/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/__init__.py b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/__init__.py new file mode 100644 index 0000000000..38f0f87385 --- /dev/null +++ b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/__init__.py @@ -0,0 +1 @@ +"""langchain-deep-agents-wikispeedia-v1 environment package.""" diff --git a/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/harness.py b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/harness.py new file mode 100644 index 0000000000..2fffe75261 --- /dev/null +++ b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import WikispeediaHarness as WikispeediaHarness +from .taskset import WikispeediaHarnessConfig as WikispeediaHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/taskset.py b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/taskset.py new file mode 100644 index 0000000000..74bb6d5f2f --- /dev/null +++ b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/taskset.py @@ -0,0 +1,514 @@ +import asyncio +import json +from collections.abc import Mapping, Sequence +from typing import Protocol, cast + +from datasets import Dataset + +import verifiers.v1 as vf + +from .wiki_graph import WikiGraph, WikiPair, load_wiki_graph + + +class AgentMessage(Protocol): + role: str + content: object + + +DEEP_AGENT_TOOLS = { + "write_todos", + "write_file", + "read_file", + "ls", + "edit_file", + "grep", + "task", +} +WIKISPEEDIA_TOOLS = {"click_link", "go_back"} + + +def system_prompt(allow_go_back: bool = True) -> str: + backtracking = ( + "Use `go_back` to undo your last click." + if allow_go_back + else "Backtracking is disabled, so choose each link carefully." + ) + return f"""\ +This game is easy and fun: + +You are given two Wikipedia articles. Starting from the first article, your goal \ +is to reach the second one, exclusively by following links in the articles you \ +encounter. (For the articles you are given this is always possible.) + +Each article ends with a list of `Available links: ...` — those are the only \ +links you can follow. Use the `click_link` tool to navigate to one. \ +{backtracking} + +You also have access to deep-agent scaffolding tools (`write_todos`, \ +`write_file`, `read_file`, `ls`, `edit_file`, `task`). Use them when they help. + +When you reach the target the system will say `TARGET REACHED`. Stop calling \ +tools at that point and reply with a brief confirmation.""" + + +class WikispeediaTasksetConfig(vf.TasksetConfig): + id: str = "langchain-deep-agents-wikispeedia" + cache_dir: str | None = None + min_path_length: int = 3 + max_path_length: int = 6 + train_size: int = 50_000 + eval_size: int = 1_000 + eval_target_fraction: float = 0.1 + split_seed: int = 0 + links_only: bool = False + allow_go_back: bool = True + max_turns: int = 50 + efficiency_weight: float = 0.0 + stratify_path_length: bool = True + + +class WikispeediaHarnessConfig(vf.HarnessConfig): + max_turns: int = 50 + timeout_seconds: float = 1200.0 + + +class WikispeediaTask(vf.Task): + answer: str + source: str + target: str + shortest_path: int + cache_dir: str | None = None + links_only: bool = False + allow_go_back: bool = True + + +class WikispeediaTaskset(vf.Taskset[WikispeediaTasksetConfig]): + task_type = WikispeediaTask + + def load_system_prompt(self, config: WikispeediaTasksetConfig) -> vf.SystemPrompt: + return system_prompt(allow_go_back=config.allow_go_back) + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return load_tasks(self.config, split=split) + + @vf.reward(weight=1.0) + async def reached_target(self, state: vf.State) -> float: + return float(bool(state.extras.get("reached_target", False))) + + @vf.reward(weight=1.0) + async def path_efficiency_reward(self, state: vf.State) -> float: + if self.config.efficiency_weight <= 0: + return 0.0 + return self.config.efficiency_weight * path_efficiency(state) + + @vf.metric + async def path_efficiency(self, state: vf.State) -> float: + return path_efficiency(state) + + @vf.metric + async def path_length(self, state: vf.State) -> float: + return float(max(len(path(state)) - 1, 0)) + + @vf.metric + async def shortest_path(self, state: vf.State) -> float: + value = state.extras.get("shortest_path", 0) + return float(value) if isinstance(value, int | float) else 0.0 + + @vf.metric + async def agent_timeout(self, state: vf.State) -> float: + return float(bool(state.extras.get("agent_timeout", False))) + + @vf.metric + async def total_tool_calls(self, state: vf.State) -> float: + return float(count_tool_calls(state)) + + @vf.metric + async def assistant_turns(self, state: vf.State) -> float: + return float(len(state.transcript)) + + @vf.metric + async def invalid_link_rate(self, state: vf.State) -> float: + clicks = 0 + invalid = 0 + id_to_name = { + tool_call.id: tool_call.name + for turn in state.transcript + for tool_call in turn.tool_calls + } + for turn in state.transcript: + for result in turn.tool_results: + if id_to_name.get(result.tool_call_id) != "click_link": + continue + clicks += 1 + if ( + isinstance(result.content, str) + and "is not a valid link" in result.content + ): + invalid += 1 + return float(invalid / clicks) if clicks else 0.0 + + +class WikispeediaHarness(vf.Harness[WikispeediaHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = WikispeediaTask.model_validate(context.task.model_dump()) + state = context.state + runtime = context.runtime + if runtime is None: + raise ValueError("WikispeediaHarness requires a runtime.") + from deepagents import create_deep_agent + from langchain_core.tools import tool + from langchain_openai import ChatOpenAI + from langgraph.errors import GraphRecursionError + + wiki = load_wiki_graph(cache_dir(task)) + init_navigation_state(task, state) + prompt = self.initial_messages(task) + + @tool + async def click_link(article: str) -> str: + """Navigate to a linked Wikipedia article.""" + return click_link_result(article, wiki, state) + + nav_tools = [click_link] + if allow_go_back(task): + + @tool + async def go_back() -> str: + """Undo the last click_link and return to the previous article.""" + return go_back_result(wiki, state) + + nav_tools.append(go_back) + + async def stop_check() -> str | None: + if await self.is_completed(context): + return state.stop_condition or "stop" + return None + + async with vf.InterceptionServer( + context, + task, + state, + protocols=self.protocols, + stop_check=stop_check, + ) as endpoint: + endpoint_url = await runtime.expose(endpoint.port) + endpoint_env = endpoint.env(base_url=endpoint_url, model=context.model) + system_messages = [ + message for message in prompt if message.role == "system" + ] + user_messages = [message for message in prompt if message.role != "system"] + model = ChatOpenAI( + model=endpoint_env["OPENAI_MODEL"], + base_url=endpoint_env["OPENAI_BASE_URL"], + api_key=endpoint_env["OPENAI_API_KEY"], + ) + agent = create_deep_agent( + model=model, + tools=nav_tools, + system_prompt="\n\n".join( + str(message.content or "") for message in system_messages + ), + ) + invoke_config = ( + {"recursion_limit": self.config.max_turns} + if self.config.max_turns > 0 + else None + ) + invoke = agent.ainvoke( + { + "messages": [ + { + "role": "user", + "content": "\n\n".join( + str(message.content or "") for message in user_messages + ), + } + ] + }, + config=invoke_config, + ) + try: + result = await asyncio.wait_for( + invoke, timeout=self.config.timeout_seconds + ) + except (TimeoutError, GraphRecursionError) as exc: + state.extras["agent_timeout"] = True + state.stop( + "agent_timeout" + if isinstance(exc, TimeoutError) + else "agent_recursion_limit" + ) + return + + messages = result.get("messages", []) if isinstance(result, Mapping) else [] + completion = serialize_agent_completion(messages) + if completion: + final = completion[-1] + content = final.content if hasattr(final, "content") else "" + state.artifacts["agent_result"] = str(content or "") + if not state.transcript: + state.transcript.append(vf.Turn(prompt=prompt, completion=completion)) + if not state.is_completed: + state.stop("agent_completed") + + +def format_article(wiki: WikiGraph, article: str, links_only: bool = False) -> str: + links = wiki.get_links(article) + links_str = ", ".join(links) if links else "(no outgoing links)" + if links_only: + return f"# {article}\n\nAvailable links: {links_str}" + text = wiki.get_text(article) + return f"# {article}\n\n{text}\n\n---\nAvailable links: {links_str}" + + +def build_dataset( + wiki: WikiGraph, + pairs: list[WikiPair], + cache_dir: str | None, + links_only: bool, + allow_go_back: bool, + max_turns: int, +) -> Dataset: + records: list[vf.JsonData] = [] + for source, target, dist in pairs: + starting = format_article(wiki, source, links_only=links_only) + prompt_text = ( + f"Your mission: {source} >> {target}\n\n" + f"Here is the starting article:\n\n{starting}" + ) + _ = wiki.get_human_stats(source, target) + records.append( + { + "task_id": f"{source}->{target}", + "prompt": [{"role": "user", "content": prompt_text}], + "answer": target, + "source": source, + "target": target, + "shortest_path": dist, + "cache_dir": cache_dir, + "links_only": links_only, + "allow_go_back": allow_go_back, + "max_turns": max_turns, + } + ) + return Dataset.from_list(records) + + +def split_pairs( + config: WikispeediaTasksetConfig, +) -> tuple[list[WikiPair], list[WikiPair]]: + return load_wiki_graph(config.cache_dir).split_pairs( + train_size=config.train_size, + eval_size=config.eval_size, + min_dist=config.min_path_length, + max_dist=config.max_path_length, + eval_target_fraction=config.eval_target_fraction, + seed=config.split_seed, + stratify=config.stratify_path_length, + ) + + +def load_tasks( + config: WikispeediaTasksetConfig, split: vf.TaskSplit = "train" +) -> Dataset: + train, eval_ = split_pairs(config) + return build_dataset( + load_wiki_graph(config.cache_dir), + train if split == "train" else eval_, + cache_dir=config.cache_dir, + links_only=config.links_only, + allow_go_back=config.allow_go_back, + max_turns=config.max_turns, + ) + + +def init_navigation_state(task: WikispeediaTask, state: vf.State) -> None: + state.extras["current_article"] = task.source + state.extras["path"] = [task.source] + state.extras["target"] = task.target + state.extras["shortest_path"] = task.shortest_path + state.extras["reached_target"] = False + state.extras["agent_timeout"] = False + state.extras["links_only"] = task.links_only + + +def cache_dir(task: WikispeediaTask) -> str | None: + return task.cache_dir + + +def allow_go_back(task: WikispeediaTask) -> bool: + return task.allow_go_back + + +def current_article(state: vf.State) -> str: + value = state.extras.get("current_article") + if not isinstance(value, str): + raise RuntimeError("Wikispeedia current article is not initialized.") + return value + + +def target_article(state: vf.State) -> str: + value = state.extras.get("target") + if not isinstance(value, str): + raise RuntimeError("Wikispeedia target article is not initialized.") + return value + + +def path(state: vf.State) -> list[str]: + value = state.extras.get("path") + if isinstance(value, list) and all(isinstance(item, str) for item in value): + return list(value) + return [] + + +def set_path(state: vf.State, articles: list[str]) -> None: + state.extras["path"] = articles + + +def click_link_result(article: str, wiki: WikiGraph, state: vf.State) -> str: + links_only = bool(state.extras.get("links_only", False)) + current = current_article(state) + available = wiki.get_links(current) + normalized = wiki.normalize_name(article) + if normalized is None or normalized not in available: + avail_str = ", ".join(available) if available else "(none)" + return ( + f"'{article}' is not a valid link from '{current}'.\n" + f"Available links: {avail_str}" + ) + route = path(state) + route.append(normalized) + set_path(state, route) + state.extras["current_article"] = normalized + if normalized == target_article(state): + state.extras["reached_target"] = True + state.stop("target_reached") + return ( + f"TARGET REACHED: {normalized}\n\n" + "You successfully navigated to the target. Stop calling tools and reply briefly." + ) + return format_article(wiki, normalized, links_only=links_only) + + +def go_back_result(wiki: WikiGraph, state: vf.State) -> str: + route = path(state) + if len(route) <= 1: + return "You are already at the starting article. Cannot go back." + route.pop() + set_path(state, route) + state.extras["current_article"] = route[-1] + return format_article( + wiki, route[-1], links_only=bool(state.extras.get("links_only", False)) + ) + + +def path_efficiency(state: vf.State) -> float: + if not bool(state.extras.get("reached_target", False)): + return 0.0 + shortest = 0.0 + raw_shortest = state.extras.get("shortest_path") + if isinstance(raw_shortest, int | float) and not isinstance(raw_shortest, bool): + shortest = float(raw_shortest) + actual = max(len(path(state)) - 1, 1) + return min(1.0, shortest / actual) if shortest > 0 else 0.0 + + +def count_tool_calls(state: vf.State, name: str | None = None) -> int: + names = [ + tool_call.name for turn in state.transcript for tool_call in turn.tool_calls + ] + if name is None: + return len(names) + return sum(1 for tool_name in names if tool_name == name) + + +def serialize_agent_completion( + messages: Sequence[AgentMessage | vf.JsonData], +) -> vf.Messages: + role_aliases = { + "human": "user", + "ai": "assistant", + "tool": "tool", + "system": "system", + } + call_names: dict[str, str] = {} + serialized: list[vf.JsonData] = [] + for message in messages: + if isinstance(message, Mapping): + payload = cast(vf.JsonData, dict(message)) + else: + model_dump = getattr(message, "model_dump", None) + payload = cast( + vf.JsonData, + model_dump(mode="json", exclude_none=True) + if callable(model_dump) + else { + "role": getattr(message, "role", None) + or getattr(message, "type", "assistant"), + "content": getattr(message, "content", str(message)), + "name": getattr(message, "name", None), + "tool_call_id": getattr(message, "tool_call_id", None), + "tool_calls": getattr(message, "tool_calls", None), + }, + ) + raw_role = payload.get("role") or payload.get("type") or "assistant" + role = role_aliases.get(str(raw_role), str(raw_role)) + item: vf.JsonData = {"role": role, "content": payload.get("content", "")} + tool_calls = payload.get("tool_calls") + if isinstance(tool_calls, list) and tool_calls: + normalized_tool_calls: list[vf.JsonData] = [] + for tool_call in tool_calls: + if not isinstance(tool_call, Mapping): + continue + tool_call_payload = cast(vf.JsonData, dict(tool_call)) + name = tool_call_payload.get("name") + tool_id = tool_call_payload.get("id") or tool_call_payload.get( + "tool_call_id" + ) + if isinstance(tool_id, str) and isinstance(name, str): + call_names[tool_id] = name + arguments = tool_call_payload.get("arguments") + if not isinstance(arguments, str): + args = tool_call_payload.get("args", {}) + try: + arguments = json.dumps(args if args is not None else {}) + except (TypeError, ValueError): + arguments = str(args) + tool_call_payload["arguments"] = arguments + normalized_tool_calls.append(tool_call_payload) + item["tool_calls"] = normalized_tool_calls + name = payload.get("name") + if isinstance(name, str): + item["name"] = name + tool_call_id = payload.get("tool_call_id") + if isinstance(tool_call_id, str): + item["tool_call_id"] = tool_call_id + if item["role"] == "tool" and "name" not in item: + tool_name = call_names.get(tool_call_id) + if tool_name is not None: + item["name"] = tool_name + serialized.append(item) + if serialized and serialized[0].get("role") == "user": + serialized = serialized[1:] + parsed: vf.Messages = [] + for item in serialized: + role = item.get("role") + if role == "assistant": + parsed.append(vf.AssistantMessage.model_validate(item)) + elif role == "user": + parsed.append(vf.UserMessage.model_validate(item)) + elif role == "tool": + parsed.append(vf.ToolMessage.model_validate(item)) + elif role == "system": + parsed.append(vf.SystemMessage.model_validate(item)) + else: + raise ValueError(f"Unsupported LangChain message role: {role!r}.") + return parsed + + +def load_taskset(config: WikispeediaTasksetConfig) -> WikispeediaTaskset: + return WikispeediaTaskset(config=config) + + +def load_harness(config: WikispeediaHarnessConfig) -> WikispeediaHarness: + return WikispeediaHarness(config=config) diff --git a/environments/langchain_deep_agents_wikispeedia/wiki_graph.py b/environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/wiki_graph.py similarity index 100% rename from environments/langchain_deep_agents_wikispeedia/wiki_graph.py rename to environments/langchain_deep_agents_wikispeedia_v1/langchain_deep_agents_wikispeedia_v1/wiki_graph.py diff --git a/environments/langchain_deep_agents_wikispeedia/pyproject.toml b/environments/langchain_deep_agents_wikispeedia_v1/pyproject.toml similarity index 81% rename from environments/langchain_deep_agents_wikispeedia/pyproject.toml rename to environments/langchain_deep_agents_wikispeedia_v1/pyproject.toml index 5818a214a5..b32ff08c6c 100644 --- a/environments/langchain_deep_agents_wikispeedia/pyproject.toml +++ b/environments/langchain_deep_agents_wikispeedia_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "langchain-deep-agents-wikispeedia" +name = "langchain-deep-agents-wikispeedia-v1" description = "V1 Taskset/Harness environment training LangChain deep-agents on Wikispeedia navigation" tags = ["v1", "taskset", "harness", "multi-turn", "tool-use", "langchain", "deep-agents", "wikispeedia", "navigation", "train", "eval"] version = "0.1.4" @@ -17,7 +17,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["langchain_deep_agents_wikispeedia.py", "wiki_graph.py", "pyproject.toml"] +include = ["langchain_deep_agents_wikispeedia_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/math_python/math_python.py b/environments/math_python/math_python.py index 61d52f7e33..bca7712cdf 100644 --- a/environments/math_python/math_python.py +++ b/environments/math_python/math_python.py @@ -16,47 +16,10 @@ def load_environment( sandbox_timeout_minutes: int = 60, sandbox_timeout_per_command_seconds: int = 60, sandbox_client_max_workers: int | None = None, - v1: bool = False, **kwargs, ): - if v1: - unsupported = [*kwargs] - if max_startup_wait_seconds != 60: - unsupported.append("max_startup_wait_seconds") - if sandbox_client_max_workers is not None: - unsupported.append("sandbox_client_max_workers") - if unsupported: - unexpected = ", ".join(sorted(unsupported)) - raise TypeError(f"Unsupported v1 load_environment kwargs: {unexpected}") - - from math_python_v1 import ( - MathPythonEnvConfig, - MathPythonHarnessConfig, - MathPythonTasksetConfig, - build_system_prompt, - load_environment as load_v1, - ) - - return load_v1( - config=MathPythonEnvConfig( - taskset=MathPythonTasksetConfig( - dataset_name=dataset_name, - dataset_split=dataset_split, - num_train_examples=num_train_examples, - system_prompt=build_system_prompt(pip_install_packages), - ), - harness=MathPythonHarnessConfig( - max_turns=max_turns, - pip_install_packages=pip_install_packages, - sandbox_cpu_cores=sandbox_cpu_cores, - sandbox_memory_gb=sandbox_memory_gb, - sandbox_disk_size_gb=sandbox_disk_size_gb, - sandbox_gpu_count=sandbox_gpu_count, - sandbox_timeout_minutes=sandbox_timeout_minutes, - sandbox_timeout_per_command_seconds=sandbox_timeout_per_command_seconds, - ), - ) - ) + if "v1" in kwargs: + raise TypeError("math_python is v0-only; use math_python_v1.") def build_dataset(): return load_example_dataset(dataset_name, dataset_split, n=num_train_examples) diff --git a/environments/math_python/math_python_v1.py b/environments/math_python/math_python_v1.py deleted file mode 100644 index e7fa12f9a7..0000000000 --- a/environments/math_python/math_python_v1.py +++ /dev/null @@ -1,208 +0,0 @@ -import json - -from math_verify import parse, verify -from pydantic import model_validator - -import verifiers as vf -from verifiers.errors import SandboxError -from verifiers.utils.data_utils import extract_boxed_answer, load_example_dataset - - -async def python(code: str, sandbox, state) -> str: - """Execute Python code in the rollout sandbox.""" - history = state.setdefault("python_history", []) - script = f""" -import ast -import contextlib -import io -import traceback - -history = {json.dumps(history)} -code = {json.dumps(code)} -namespace = {{}} - -try: - for snippet in history: - exec(compile(snippet, "", "exec"), namespace, namespace) - - tree = ast.parse(code, "", "exec") - stdout = io.StringIO() - with contextlib.redirect_stdout(stdout): - if tree.body and isinstance(tree.body[-1], ast.Expr): - prefix = ast.Module(body=tree.body[:-1], type_ignores=[]) - exec(compile(prefix, "", "exec"), namespace, namespace) - expression = ast.Expression(tree.body[-1].value) - result = eval(compile(expression, "", "eval"), namespace, namespace) - if result is not None: - print(repr(result)) - else: - exec(compile(tree, "", "exec"), namespace, namespace) - print(stdout.getvalue(), end="") -except BaseException: - traceback.print_exc() - raise SystemExit(1) -""" - await sandbox.upload_bytes("/tmp/vf_python_tool.py", script.encode()) - result = await sandbox.execute("python /tmp/vf_python_tool.py") - stdout = result.stdout or "" - stderr = result.stderr or "" - if result.exit_code: - raise SandboxError(f"Python command failed: {stderr}") - history.append(code) - return stdout.strip() or "(no output)" - - -@vf.reward(weight=1.0) -async def correct_answer(task, state) -> float: - completion = state.get("completion") or [] - messages = vf.get_messages(completion, role="assistant") - response_text = str(messages[-1].content or "") if messages else "" - response = extract_boxed_answer(response_text) - answer = str(task["answer"]) - if not response or len(response) > 50_000: - return 0.0 - - try: - parsed_answer = parse(rf"\boxed{{{answer}}}", parsing_timeout=5) - parsed_response = parse(rf"\boxed{{{response}}}", parsing_timeout=5) - return float(verify(parsed_answer, parsed_response, timeout_seconds=5)) - except BaseException: - return 0.0 - - -@vf.cleanup(priority=10) -async def collect_python_commands(task, state): - state["commands"] = list(state.get("sandbox_commands", [])) - state.pop("sandbox_commands", None) - - -def build_system_prompt(pip_install_packages: str = "numpy sympy scipy") -> str: - pip_install_prompt = ( - f"In addition to the Python standard library, you have access to: {pip_install_packages}." - if pip_install_packages.strip() - else "You may only use the Python standard library." - ) - return ( - "Use Python for all calculations. Give your answer inside \\boxed{}." - "\n\n" - f"{pip_install_prompt}" - ) - - -class MathPythonTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["correct_answer"] - system_prompt: str | None = None - dataset_name: str = "math" - dataset_split: str = "train" - num_train_examples: int = -1 - - -class MathPythonHarnessConfig(vf.HarnessConfig): - max_turns: int = 100 - pip_install_packages: str = "numpy sympy scipy" - sandbox_cpu_cores: int = 1 - sandbox_memory_gb: int = 2 - sandbox_disk_size_gb: int = 5 - sandbox_gpu_count: int = 0 - sandbox_timeout_minutes: int = 60 - sandbox_timeout_per_command_seconds: int = 60 - - -class MathPythonEnvConfig(vf.EnvConfig): - taskset: MathPythonTasksetConfig = MathPythonTasksetConfig() - harness: MathPythonHarnessConfig = MathPythonHarnessConfig() - - @model_validator(mode="after") - def derive_taskset_system_prompt(self) -> "MathPythonEnvConfig": - if "system_prompt" not in self.taskset.model_fields_set: - self.taskset = self.taskset.model_copy( - update={ - "system_prompt": build_system_prompt( - self.harness.pip_install_packages - ) - } - ) - return self - - -def load_tasks( - dataset_name: str = "math", - dataset_split: str = "train", - num_train_examples: int = -1, -): - dataset = load_example_dataset( - dataset_name, - dataset_split, - n=num_train_examples, - ) - for index, row in enumerate(dataset): - yield { - **row, - "example_id": index, - "prompt": [{"role": "user", "content": row["question"]}], - } - - -class MathPythonTaskset(vf.Taskset[MathPythonTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks( - dataset_name=self.config.dataset_name, - dataset_split=self.config.dataset_split, - num_train_examples=self.config.num_train_examples, - ) - - -class MathPythonHarness(vf.Harness[MathPythonHarnessConfig]): - pass - - -def load_toolset( - pip_install_packages: str = "numpy sympy scipy", - sandbox_cpu_cores: int = 1, - sandbox_memory_gb: int = 2, - sandbox_disk_size_gb: int = 5, - sandbox_gpu_count: int = 0, - sandbox_timeout_minutes: int = 60, - sandbox_timeout_per_command_seconds: int = 60, -): - packages = pip_install_packages.split() if pip_install_packages.strip() else [] - return vf.Toolset( - tools=[python], - write=True, - sandbox=vf.SandboxConfig( - image="python:3.11-slim", - scope="group", - cpu_cores=sandbox_cpu_cores, - memory_gb=sandbox_memory_gb, - disk_size_gb=sandbox_disk_size_gb, - gpu_count=sandbox_gpu_count, - timeout_minutes=sandbox_timeout_minutes, - command_timeout=sandbox_timeout_per_command_seconds, - packages=packages, - ), - cleanups=[collect_python_commands], - ) - - -def load_environment(config: MathPythonEnvConfig) -> vf.Env: - harness = MathPythonHarness(config=config.harness) - if "toolsets" not in config.harness.model_fields_set: - harness.add_toolset( - { - "python": load_toolset( - pip_install_packages=config.harness.pip_install_packages, - sandbox_cpu_cores=config.harness.sandbox_cpu_cores, - sandbox_memory_gb=config.harness.sandbox_memory_gb, - sandbox_disk_size_gb=config.harness.sandbox_disk_size_gb, - sandbox_gpu_count=config.harness.sandbox_gpu_count, - sandbox_timeout_minutes=config.harness.sandbox_timeout_minutes, - sandbox_timeout_per_command_seconds=( - config.harness.sandbox_timeout_per_command_seconds - ), - ) - } - ) - return vf.Env( - taskset=MathPythonTaskset(config=config.taskset), - harness=harness, - ) diff --git a/environments/math_python/pyproject.toml b/environments/math_python/pyproject.toml index 42b02125de..60ea8f3f5d 100644 --- a/environments/math_python/pyproject.toml +++ b/environments/math_python/pyproject.toml @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["math_python.py", "math_python_v1.py", "pyproject.toml"] +include = ["math_python.py", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/math_python_v1/README.md b/environments/math_python_v1/README.md new file mode 100644 index 0000000000..527fe34b9f --- /dev/null +++ b/environments/math_python_v1/README.md @@ -0,0 +1,57 @@ +# math-python-v1 + + +Source Code + + +### Overview +- **Environment ID**: `math-python-v1` +- **Short description**: v1 tool-using math environment with a task-owned Python `Toolset`; graded by symbolic equivalence. +- **Tags**: math, tools, python, single-turn, boxed-answer + +### Datasets +- **Primary dataset(s)**: Example `math` dataset via `load_example_dataset` +- **Source links**: Uses example loader in `verifiers.utils.data_utils` +- **Split sizes**: Configurable via args; defaults to `train` split and all examples + +### Task +- **Type**: `vf.Env` with a math `vf.Taskset`, base `vf.Harness`, and task-owned Python `Toolset`. +- **Rubric overview**: Correctness by `math_verify.parse` + `verify` over the final boxed answer. + +### Quickstart +Run an evaluation with default settings: + +```bash +prime eval run math-python-v1 +``` + +Configure model and sampling: + +```bash +prime eval run math-python-v1 \ + -m openai/gpt-4.1-mini \ + -n 20 -r 3 -t 1024 -T 0.7 \ + -a '{"config": {"taskset": {"dataset_name": "math", "dataset_split": "train", "num_train_examples": -1}}}' +``` + +Notes: +- v1 task settings belong under `config.taskset` when passed through `-a` / `--env-args`. + +### Taskset Config +| Field | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `dataset_name` | str | `"math"` | Example dataset to load | +| `dataset_split` | str | `"train"` | Split to load | +| `num_train_examples` | int | `-1` | Limit dataset size (`-1` for all) | +| `pip_install_packages` | str | `"numpy sympy scipy"` | Packages listed in the generated system prompt | + +### Harness Config +| Field | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `max_turns` | int | `100` | Maximum model turns per rollout | + +### Metrics +| Metric | Meaning | +| ------ | ------- | +| `reward` | 1.0 if symbolic verification passes, else 0.0 | +| `num_turns` | Number of recorded model turns | diff --git a/environments/math_python_v1/math_python_v1/__init__.py b/environments/math_python_v1/math_python_v1/__init__.py new file mode 100644 index 0000000000..4e5c31e75e --- /dev/null +++ b/environments/math_python_v1/math_python_v1/__init__.py @@ -0,0 +1 @@ +"""math-python-v1 environment package.""" diff --git a/environments/math_python_v1/math_python_v1/harness.py b/environments/math_python_v1/math_python_v1/harness.py new file mode 100644 index 0000000000..a6d937c89b --- /dev/null +++ b/environments/math_python_v1/math_python_v1/harness.py @@ -0,0 +1,2 @@ +from .taskset import MathPythonHarnessConfig as MathPythonHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/math_python_v1/math_python_v1/servers/__init__.py b/environments/math_python_v1/math_python_v1/servers/__init__.py new file mode 100644 index 0000000000..3f425d2f2c --- /dev/null +++ b/environments/math_python_v1/math_python_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for math-python-v1.""" diff --git a/environments/math_python_v1/math_python_v1/servers/python/__init__.py b/environments/math_python_v1/math_python_v1/servers/python/__init__.py new file mode 100644 index 0000000000..5415172733 --- /dev/null +++ b/environments/math_python_v1/math_python_v1/servers/python/__init__.py @@ -0,0 +1,3 @@ +from .config import PythonToolsetConfig + +__all__ = ["PythonToolsetConfig"] diff --git a/environments/math_python_v1/math_python_v1/servers/python/config.py b/environments/math_python_v1/math_python_v1/servers/python/config.py new file mode 100644 index 0000000000..d726fb08ed --- /dev/null +++ b/environments/math_python_v1/math_python_v1/servers/python/config.py @@ -0,0 +1,5 @@ +import verifiers.v1 as vf + + +class PythonToolsetConfig(vf.ToolsetConfig): + pass diff --git a/environments/math_python_v1/math_python_v1/servers/python/toolset.py b/environments/math_python_v1/math_python_v1/servers/python/toolset.py new file mode 100644 index 0000000000..2dd83a86da --- /dev/null +++ b/environments/math_python_v1/math_python_v1/servers/python/toolset.py @@ -0,0 +1,49 @@ +import ast +import contextlib +import io +import traceback + +import verifiers.v1 as vf + +from .config import PythonToolsetConfig + + +def execute_python(code: str, history: list[str]) -> str: + namespace: dict[str, object] = {} + for snippet in history: + exec(compile(snippet, "", "exec"), namespace, namespace) + tree = ast.parse(code, "", "exec") + stdout = io.StringIO() + with contextlib.redirect_stdout(stdout): + if tree.body and isinstance(tree.body[-1], ast.Expr): + prefix = ast.Module(body=tree.body[:-1], type_ignores=[]) + exec(compile(prefix, "", "exec"), namespace, namespace) + expression = ast.Expression(tree.body[-1].value) + result = eval(compile(expression, "", "eval"), namespace, namespace) + if result is not None: + print(repr(result)) + else: + exec(compile(tree, "", "exec"), namespace, namespace) + history.append(code) + return stdout.getvalue().strip() or "(no output)" + + +class PythonToolset(vf.Toolset[PythonToolsetConfig]): + @vf.resource + def history(self) -> list[str]: + return [] + + @vf.tool( + args={"history": "resources.history"}, + extends={"python_history": "state.extras.python_history"}, + ) + def python(self, code: str, history: list[str]) -> dict: + start = len(history) + try: + content = execute_python(code, history) + except BaseException: + content = traceback.format_exc() + return { + "content": content, + "python_history": list(history[start:]), + } diff --git a/environments/math_python_v1/math_python_v1/taskset.py b/environments/math_python_v1/math_python_v1/taskset.py new file mode 100644 index 0000000000..730f313c09 --- /dev/null +++ b/environments/math_python_v1/math_python_v1/taskset.py @@ -0,0 +1,85 @@ +from math_verify import parse, verify +import verifiers.v1 as vf +from verifiers.utils.data_utils import extract_boxed_answer, load_example_dataset + +from .servers.python import PythonToolsetConfig + + +def build_system_prompt(pip_install_packages: str = "numpy sympy scipy") -> str: + pip_install_prompt = ( + f"In addition to the Python standard library, you have access to: {pip_install_packages}." + if pip_install_packages.strip() + else "You may only use the Python standard library." + ) + return ( + "Use the python tool for calculations when useful. Give your answer " + "inside \\boxed{}.\n\n" + f"{pip_install_prompt}" + ) + + +class MathPythonTasksetConfig(vf.TasksetConfig): + system_prompt: str | None = None + toolsets: vf.ToolsetConfigs = {"python": PythonToolsetConfig()} + pip_install_packages: str = "numpy sympy scipy" + dataset_name: str = "math" + dataset_split: str = "train" + num_train_examples: int = -1 + + +class MathPythonHarnessConfig(vf.HarnessConfig): + max_turns: int = 100 + + +class MathPythonTask(vf.Task): + question: str + answer: str + + +class MathPythonTaskset(vf.Taskset[MathPythonTasksetConfig]): + task_type = MathPythonTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + dataset = load_example_dataset( + self.config.dataset_name, + self.config.dataset_split, + n=self.config.num_train_examples, + ) + return [ + { + **row, + "row_id": index, + "prompt": [{"role": "user", "content": row["question"]}], + } + for index, row in enumerate(dataset) + ] + + @vf.reward(weight=1.0) + async def correct_answer(self, task: MathPythonTask, state: vf.State) -> float: + messages = [ + message for message in state.completion if message.role == "assistant" + ] + response_text = str(messages[-1].content or "") if messages else "" + response = extract_boxed_answer(response_text) + if not response or len(response) > 50_000: + return 0.0 + try: + parsed_answer = parse(rf"\boxed{{{task.answer}}}", parsing_timeout=5) + parsed_response = parse(rf"\boxed{{{response}}}", parsing_timeout=5) + return float(verify(parsed_answer, parsed_response, timeout_seconds=5)) + except BaseException: + return 0.0 + + +def load_taskset(config: MathPythonTasksetConfig) -> MathPythonTaskset: + if "system_prompt" not in config.model_fields_set: + config = config.model_copy( + update={"system_prompt": build_system_prompt(config.pip_install_packages)} + ) + return MathPythonTaskset(config=config) + + +def load_harness(config: MathPythonHarnessConfig) -> vf.Harness: + return vf.Harness(config=config) diff --git a/environments/math_python_v1/pyproject.toml b/environments/math_python_v1/pyproject.toml new file mode 100644 index 0000000000..249fc08bfe --- /dev/null +++ b/environments/math_python_v1/pyproject.toml @@ -0,0 +1,22 @@ +[project] +name = "math-python-v1" +description = "Solve math problems using Python in a sandbox environment" +tags = ["tool-use", "math", "sandbox", "train", "prime-sandboxes", "python", "coding"] +version = "0.1.10" +requires-python = ">=3.11" +dependencies = [ + "verifiers>=0.1.8.post2", + "math-verify>=0.8.0", + "prime-sandboxes>=0.2.7", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build] +include = ["math_python_v1/**/*", "README.md", "pyproject.toml"] + +[tool.verifiers.eval] +num_examples = 5 +rollouts_per_example = 3 diff --git a/environments/mcp_search_env/README.md b/environments/mcp_search_env_v1/README.md similarity index 84% rename from environments/mcp_search_env/README.md rename to environments/mcp_search_env_v1/README.md index da8e09935c..c619c89950 100644 --- a/environments/mcp_search_env/README.md +++ b/environments/mcp_search_env_v1/README.md @@ -1,11 +1,11 @@ -# mcp-search-env +# mcp-search-env-v1 v1 `vf.Env` example for MCP-backed tool use. The taskset asks short synthetic search questions, and the bundled stdio MCP server exposes `search_records` and `read_record` over a small stable record corpus. ```bash -prime eval run mcp-search-env -m openai/gpt-4.1-mini -n 5 -r 1 +prime eval run mcp-search-env-v1 -m openai/gpt-4.1-mini -n 5 -r 1 ``` Configuration belongs under v1 sections: diff --git a/environments/mcp_search_env_v1/mcp_search_env_v1/__init__.py b/environments/mcp_search_env_v1/mcp_search_env_v1/__init__.py new file mode 100644 index 0000000000..72b9b7cb5b --- /dev/null +++ b/environments/mcp_search_env_v1/mcp_search_env_v1/__init__.py @@ -0,0 +1 @@ +"""mcp-search-env-v1 environment package.""" diff --git a/environments/mcp_search_env_v1/mcp_search_env_v1/servers/__init__.py b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/__init__.py new file mode 100644 index 0000000000..93b8e9cd3b --- /dev/null +++ b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for mcp-search-env-v1.""" diff --git a/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/__init__.py b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/__init__.py new file mode 100644 index 0000000000..e0846ea1b6 --- /dev/null +++ b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/__init__.py @@ -0,0 +1,3 @@ +from .config import SearchToolsetConfig + +__all__ = ["SearchToolsetConfig"] diff --git a/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/config.py b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/config.py new file mode 100644 index 0000000000..c7719d2a4e --- /dev/null +++ b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/config.py @@ -0,0 +1,5 @@ +import verifiers.v1 as vf + + +class SearchToolsetConfig(vf.ToolsetConfig): + pass diff --git a/environments/mcp_search_env/mcp_server.py b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/toolset.py similarity index 71% rename from environments/mcp_search_env/mcp_server.py rename to environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/toolset.py index 75c9a6e524..9171512316 100644 --- a/environments/mcp_search_env/mcp_server.py +++ b/environments/mcp_search_env_v1/mcp_search_env_v1/servers/search/toolset.py @@ -1,8 +1,8 @@ import json -from mcp.server.fastmcp import FastMCP +import verifiers.v1 as vf -mcp = FastMCP("mcp-search-env") +from .config import SearchToolsetConfig RECORDS = { "kiln_battery_loop": { @@ -58,25 +58,21 @@ } -@mcp.tool() -def search_records(query: str) -> str: - normalized_tokens = set(query.lower().split()) - matches = [] - for record_id, record in RECORDS.items(): - keywords = set(record["keywords"]) - title_tokens = set(str(record["title"]).lower().split()) - if normalized_tokens & (keywords | title_tokens): - matches.append({"record_id": record_id, "title": record["title"]}) - return json.dumps(matches) +class SearchToolset(vf.Toolset[SearchToolsetConfig]): + @vf.tool + def search_records(self, query: str) -> str: + normalized_tokens = set(query.lower().split()) + matches = [] + for record_id, record in RECORDS.items(): + keywords = set(record["keywords"]) + title_tokens = set(str(record["title"]).lower().split()) + if normalized_tokens & (keywords | title_tokens): + matches.append({"record_id": record_id, "title": record["title"]}) + return json.dumps(matches) - -@mcp.tool() -def read_record(record_id: str) -> str: - if record_id not in RECORDS: - raise ValueError(f"Unknown record_id: {record_id}") - record = RECORDS[record_id] - return f"{record['title']}\n\n{record['summary']}" - - -if __name__ == "__main__": - mcp.run(transport="stdio") + @vf.tool + def read_record(self, record_id: str) -> str: + if record_id not in RECORDS: + raise ValueError(f"Unknown record_id: {record_id}") + record = RECORDS[record_id] + return f"{record['title']}\n\n{record['summary']}" diff --git a/environments/mcp_search_env/mcp_search_env.py b/environments/mcp_search_env_v1/mcp_search_env_v1/taskset.py similarity index 59% rename from environments/mcp_search_env/mcp_search_env.py rename to environments/mcp_search_env_v1/mcp_search_env_v1/taskset.py index 04eaa1ddf8..53907ac72c 100644 --- a/environments/mcp_search_env/mcp_search_env.py +++ b/environments/mcp_search_env_v1/mcp_search_env_v1/taskset.py @@ -1,9 +1,8 @@ from collections.abc import Iterable -from pathlib import Path -import sys -from typing import cast -import verifiers as vf +import verifiers.v1 as vf + +from .servers.search import SearchToolsetConfig SYSTEM_PROMPT = "Use the available MCP tools to answer the question." @@ -59,25 +58,23 @@ "answer": "Curb Queue", }, ] -MCP_SERVER_PATH = str(Path(__file__).with_name("mcp_server.py")) -DEFAULT_MCP_SERVERS: list[vf.ConfigData] = [ - { - "name": "records", - "command": sys.executable, - "args": [MCP_SERVER_PATH], - "description": "Synthetic search-record MCP server", - }, -] class MCPSearchTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["exact_title_reward"] - mcp_servers: list[vf.ConfigData] | None = None max_turns: int = 6 - examples: list[vf.ConfigData] | None = None + examples: list[vf.JsonData] | None = None + toolsets: vf.ToolsetConfigs = {"records": SearchToolsetConfig()} + + +class MCPSearchTask(vf.Task): + query: str + question: str + answer: str class MCPSearchTaskset(vf.Taskset[MCPSearchTasksetConfig]): + task_type = MCPSearchTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: return load_tasks( examples=self.config.examples, max_turns=self.config.max_turns @@ -87,27 +84,12 @@ def load_system_prompt(self, config: MCPSearchTasksetConfig) -> vf.SystemPrompt: _ = config return SYSTEM_PROMPT - def load_toolsets(self, config: MCPSearchTasksetConfig) -> vf.Toolsets: - servers = config.mcp_servers or [dict(server) for server in DEFAULT_MCP_SERVERS] - return { - "records": vf.Toolset( - tools=[ - vf.MCPTool( - command=str(server["command"]), - args=[ - str(arg) - for arg in cast( - Iterable[str | int | float | bool], - server.get("args") or [], - ) - ], - env=cast(dict[str, str] | None, server.get("env")), - cwd=cast(str | None, server.get("cwd")), - ) - for server in servers - ] - ) - } + @vf.reward(weight=1.0) + async def exact_title_reward(self, task: MCPSearchTask, state: vf.State) -> float: + completion = state.completion + messages = [message for message in completion if message.role == "assistant"] + response = str(messages[-1].content or "") if messages else "" + return float(task.answer.lower() in response.lower()) def load_tasks( @@ -126,25 +108,5 @@ def load_tasks( } -@vf.reward(weight=1.0) -async def exact_title_reward(task: vf.Task, state: vf.State) -> float: - completion = state.get("completion") or [] - messages = ( - vf.get_messages(completion, role="assistant") - if isinstance(completion, list) - else [] - ) - response = str(messages[-1].content or "") if messages else "" - return float(str(task["answer"]).lower() in response.lower()) - - -class MCPSearchEnvConfig(vf.EnvConfig): - taskset: MCPSearchTasksetConfig = MCPSearchTasksetConfig() - harness: vf.HarnessConfig = vf.HarnessConfig() - - -def load_environment(config: MCPSearchEnvConfig) -> vf.Env: - return vf.Env( - taskset=MCPSearchTaskset(config=config.taskset), - harness=vf.Harness(config=config.harness), - ) +def load_taskset(config: MCPSearchTasksetConfig) -> MCPSearchTaskset: + return MCPSearchTaskset(config=config) diff --git a/environments/mcp_search_env/pyproject.toml b/environments/mcp_search_env_v1/pyproject.toml similarity index 79% rename from environments/mcp_search_env/pyproject.toml rename to environments/mcp_search_env_v1/pyproject.toml index 0e6154b6c7..79a33aad39 100644 --- a/environments/mcp_search_env/pyproject.toml +++ b/environments/mcp_search_env_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "mcp-search-env" +name = "mcp-search-env-v1" description = "v1 MCP search environment example" tags = ["eval", "mcp", "v1", "tool-use"] version = "0.2.0" @@ -14,7 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["mcp_search_env.py", "mcp_server.py", "pyproject.toml"] +include = ["mcp_search_env_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/nemo_gym_env/nemo_gym_env/__init__.py b/environments/nemo_gym_env/nemo_gym_env/__init__.py deleted file mode 100644 index 342176165b..0000000000 --- a/environments/nemo_gym_env/nemo_gym_env/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .env import load_environment - -__all__ = ["load_environment"] diff --git a/environments/nemo_gym_env/nemo_gym_env/env.py b/environments/nemo_gym_env/nemo_gym_env/env.py deleted file mode 100644 index 55ecf28f5f..0000000000 --- a/environments/nemo_gym_env/nemo_gym_env/env.py +++ /dev/null @@ -1,18 +0,0 @@ -import verifiers as vf -from harnesses import NeMoGymHarness, NeMoGymHarnessConfig -from tasksets import NeMoGymTaskset, NeMoGymTasksetConfig - - -NEMO_ENV = "example_single_tool_call" - - -class NeMoGymEnvConfig(vf.EnvConfig): - taskset: NeMoGymTasksetConfig = NeMoGymTasksetConfig(nemo_env=NEMO_ENV) - harness: NeMoGymHarnessConfig = NeMoGymHarnessConfig(nemo_env=NEMO_ENV) - - -def load_environment(config: NeMoGymEnvConfig) -> vf.Env: - return vf.Env( - taskset=NeMoGymTaskset(config=config.taskset), - harness=NeMoGymHarness(config=config.harness), - ) diff --git a/environments/nemo_gym_env/README.md b/environments/nemo_gym_env_v1/README.md similarity index 84% rename from environments/nemo_gym_env/README.md rename to environments/nemo_gym_env_v1/README.md index 017aadb9c1..e20f39fb54 100644 --- a/environments/nemo_gym_env/README.md +++ b/environments/nemo_gym_env_v1/README.md @@ -1,7 +1,7 @@ -# nemo-gym-env +# nemo-gym-env-v1 ### Overview -- **Environment ID**: `nemo-gym-env` +- **Environment ID**: `nemo-gym-env-v1` - **Short description**: Minimal v1 Verifiers environment that runs a NeMo Gym task through `NeMoGymTaskset` and `NeMoGymHarness`. - **Tags**: nemo-gym, tool-use, v1, train, eval @@ -21,19 +21,19 @@ The taskset loads NeMo Gym JSONL rows from the installed `nemo-gym` package and Run an evaluation with default settings: ```bash -prime eval run nemo-gym-env +prime eval run nemo-gym-env-v1 ``` When running directly from this repository before the NeMo Gym integration is released on PyPI, point Prime at the in-repo environments directory: ```bash -prime eval run nemo-gym-env --env-dir-path environments +prime eval run nemo-gym-env-v1 --env-dir-path environments ``` Configure model and sampling: ```bash -prime eval run nemo-gym-env \ +prime eval run nemo-gym-env-v1 \ -m gpt-4.1-mini \ -n 1 -r 1 -t 128 ``` @@ -52,7 +52,7 @@ timeout_seconds = 30 ``` ### Adapting -This example is intentionally tied to one NeMo Gym task. To create another Verifiers environment, copy this directory and change `NEMO_ENV` in `nemo_gym_env/env.py` to another packaged NeMo Gym environment name, such as `example_multi_step`, `mcqa`, or `structured_outputs`. +This example is intentionally tied to one NeMo Gym task. To create another Verifiers environment, copy this directory and change `NEMO_ENV` in `nemo_gym_env_v1/taskset.py` to another packaged NeMo Gym environment name, such as `example_multi_step`, `mcqa`, or `structured_outputs`. ### Metrics | Metric | Meaning | diff --git a/environments/nemo_gym_env_v1/nemo_gym_env_v1/__init__.py b/environments/nemo_gym_env_v1/nemo_gym_env_v1/__init__.py new file mode 100644 index 0000000000..5509928bcb --- /dev/null +++ b/environments/nemo_gym_env_v1/nemo_gym_env_v1/__init__.py @@ -0,0 +1 @@ +"""nemo-gym-env-v1 environment package.""" diff --git a/environments/nemo_gym_env_v1/nemo_gym_env_v1/harness.py b/environments/nemo_gym_env_v1/nemo_gym_env_v1/harness.py new file mode 100644 index 0000000000..0079d7825e --- /dev/null +++ b/environments/nemo_gym_env_v1/nemo_gym_env_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import NeMoGymHarness as NeMoGymHarness +from .taskset import NeMoGymHarnessConfig as NeMoGymHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/nemo_gym_env_v1/nemo_gym_env_v1/taskset.py b/environments/nemo_gym_env_v1/nemo_gym_env_v1/taskset.py new file mode 100644 index 0000000000..1219ddbc2f --- /dev/null +++ b/environments/nemo_gym_env_v1/nemo_gym_env_v1/taskset.py @@ -0,0 +1,17 @@ +from harnesses import NeMoGymHarness, NeMoGymHarnessConfig +from tasksets import NeMoGymTaskset, NeMoGymTasksetConfig + + +NEMO_ENV = "example_single_tool_call" + + +def load_taskset(config: NeMoGymTasksetConfig) -> NeMoGymTaskset: + if "nemo_env" not in config.model_fields_set: + config = config.model_copy(update={"nemo_env": NEMO_ENV}) + return NeMoGymTaskset(config=config) + + +def load_harness(config: NeMoGymHarnessConfig) -> NeMoGymHarness: + if "nemo_env" not in config.model_fields_set: + config = config.model_copy(update={"nemo_env": NEMO_ENV}) + return NeMoGymHarness(config=config) diff --git a/environments/nemo_gym_env/pyproject.toml b/environments/nemo_gym_env_v1/pyproject.toml similarity index 77% rename from environments/nemo_gym_env/pyproject.toml rename to environments/nemo_gym_env_v1/pyproject.toml index 69a62e15ab..b726fdbc33 100644 --- a/environments/nemo_gym_env/pyproject.toml +++ b/environments/nemo_gym_env_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "nemo-gym-env" +name = "nemo-gym-env-v1" description = "Example Verifiers environment backed by a NeMo Gym task." tags = ["nemo-gym", "tool-use", "v1", "train", "eval"] version = "0.1.0" @@ -15,10 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["nemo_gym_env", "README.md"] - -[tool.hatch.build.force-include] -"pyproject.toml" = "nemo_gym_env/pyproject.toml" +include = ["nemo_gym_env_v1/**/*", "README.md", "pyproject.toml"] [tool.hatch.metadata] allow-direct-references = true diff --git a/environments/nested_harness_v1/nested_harness_v1.py b/environments/nested_harness_v1/nested_harness_v1.py deleted file mode 100644 index 1e4280f057..0000000000 --- a/environments/nested_harness_v1/nested_harness_v1.py +++ /dev/null @@ -1,101 +0,0 @@ -import verifiers as vf - - -class NestedHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="parent_program") - metrics: list[str] = ["child_calls"] - - -class ChildHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="child_program") - - -async def child_program(task, state): - state["answer"] = str(task["prompt"]).upper() - state["completion"] = [{"role": "assistant", "content": state["answer"]}] - return state - - -async def call_harness(prompt, state): - _ = state - task = vf.Task({"prompt": prompt}).freeze() - harness = vf.Harness(config=ChildHarnessConfig()) - child_state = await harness.run(task) - return child_state["answer"] - - -@vf.metric -async def child_calls(task, state) -> float: - return float(len(state["child_answers"])) - - -@vf.reward(weight=1.0) -async def exact_answer(task, state) -> float: - return float(state["answer"] == task["answer"]) - - -CHILD_PROMPT_GROUPS = [ - ["hello"], - ["open", "source"], - ["taskset", "harness"], - ["runtime", "boundary"], - ["sandbox", "lease"], - ["toolset", "scope"], - ["group", "reward"], - ["endpoint", "proxy"], - ["cleanup", "signals"], - ["harbor", "tasks"], -] - - -def load_tasks(split: vf.TaskSplit = "train"): - _ = split - return [ - { - "prompt": ( - "Ask child harnesses to uppercase: " + ", ".join(child_prompts) + "." - ), - "child_prompts": child_prompts, - "answer": " ".join(prompt.upper() for prompt in child_prompts), - } - for child_prompts in CHILD_PROMPT_GROUPS - ] - - -async def parent_program(task, state): - tools = state.get_tools() - answers = [] - for prompt in task["child_prompts"]: - answer = await tools["call_harness"](prompt=prompt) - answers.append(answer) - state["child_answers"] = answers - state["answer"] = " ".join(answers) - state["completion"] = [{"role": "assistant", "content": state["answer"]}] - return state - - -class NestedHarness(vf.Harness[NestedHarnessConfig]): - def load_toolsets(self, config: NestedHarnessConfig) -> vf.Toolsets: - _ = config - return {"nested": vf.Toolset(tools=[call_harness])} - - -class NestedTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["exact_answer"] - - -class NestedTaskset(vf.Taskset[NestedTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -class NestedEnvConfig(vf.EnvConfig): - taskset: NestedTasksetConfig = NestedTasksetConfig() - harness: NestedHarnessConfig = NestedHarnessConfig() - - -def load_environment(config: NestedEnvConfig) -> vf.Env: - return vf.Env( - taskset=NestedTaskset(config=config.taskset), - harness=NestedHarness(config=config.harness), - ) diff --git a/environments/nested_harness_v1/nested_harness_v1/__init__.py b/environments/nested_harness_v1/nested_harness_v1/__init__.py new file mode 100644 index 0000000000..209e3ac6c5 --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/__init__.py @@ -0,0 +1 @@ +"""nested-harness-v1 environment package.""" diff --git a/environments/nested_harness_v1/nested_harness_v1/harness.py b/environments/nested_harness_v1/nested_harness_v1/harness.py new file mode 100644 index 0000000000..8a5bee9783 --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import NestedHarness as NestedHarness +from .taskset import NestedHarnessConfig as NestedHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/nested_harness_v1/nested_harness_v1/servers/__init__.py b/environments/nested_harness_v1/nested_harness_v1/servers/__init__.py new file mode 100644 index 0000000000..50fd349908 --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for nested-harness-v1.""" diff --git a/environments/nested_harness_v1/nested_harness_v1/servers/nested/__init__.py b/environments/nested_harness_v1/nested_harness_v1/servers/nested/__init__.py new file mode 100644 index 0000000000..b38fef821e --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/servers/nested/__init__.py @@ -0,0 +1,3 @@ +from .config import NestedToolsetConfig + +__all__ = ["NestedToolsetConfig"] diff --git a/environments/nested_harness_v1/nested_harness_v1/servers/nested/config.py b/environments/nested_harness_v1/nested_harness_v1/servers/nested/config.py new file mode 100644 index 0000000000..d9d5abb15b --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/servers/nested/config.py @@ -0,0 +1,5 @@ +import verifiers.v1 as vf + + +class NestedToolsetConfig(vf.ToolsetConfig): + pass diff --git a/environments/nested_harness_v1/nested_harness_v1/servers/nested/toolset.py b/environments/nested_harness_v1/nested_harness_v1/servers/nested/toolset.py new file mode 100644 index 0000000000..5eb4a4cd4c --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/servers/nested/toolset.py @@ -0,0 +1,9 @@ +import verifiers.v1 as vf + +from .config import NestedToolsetConfig + + +class NestedToolset(vf.Toolset[NestedToolsetConfig]): + @vf.tool + def call_harness(self, prompt: str) -> str: + return prompt.upper() diff --git a/environments/nested_harness_v1/nested_harness_v1/taskset.py b/environments/nested_harness_v1/nested_harness_v1/taskset.py new file mode 100644 index 0000000000..6911d1dc03 --- /dev/null +++ b/environments/nested_harness_v1/nested_harness_v1/taskset.py @@ -0,0 +1,95 @@ +import verifiers.v1 as vf + +from .servers.nested import NestedToolsetConfig + +CHILD_PROMPT_GROUPS = [ + ["hello"], + ["open", "source"], + ["taskset", "harness"], + ["runtime", "boundary"], + ["sandbox", "lease"], + ["toolset", "scope"], + ["group", "reward"], + ["endpoint", "proxy"], + ["cleanup", "signals"], + ["harbor", "tasks"], +] + + +class NestedTasksetConfig(vf.TasksetConfig): + toolsets: vf.ToolsetConfigs = {"nested": NestedToolsetConfig()} + + +class NestedHarnessConfig(vf.HarnessConfig): + max_turns: int = 1 + + +class NestedTask(vf.Task): + child_prompts: list[str] + answer: str + + +class NestedTaskset(vf.Taskset[NestedTasksetConfig]): + task_type = NestedTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [ + { + "prompt": [ + { + "role": "user", + "content": "Ask child harnesses to uppercase: " + + ", ".join(child_prompts) + + ".", + } + ], + "child_prompts": child_prompts, + "answer": " ".join(prompt.upper() for prompt in child_prompts), + } + for child_prompts in CHILD_PROMPT_GROUPS + ] + + @vf.metric + async def child_calls(self, state: vf.State) -> float: + answers = state.extras.get("child_answers") + return float(len(answers) if isinstance(answers, list) else 0) + + @vf.reward(weight=1.0) + async def exact_answer(self, task: NestedTask, state: vf.State) -> float: + messages = [ + message for message in state.completion if message.role == "assistant" + ] + answer = str(messages[-1].content or "").strip() if messages else "" + return float(answer == task.answer) + + +class NestedHarness(vf.Harness[NestedHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = NestedTask.model_validate(context.task.model_dump()) + state = context.state + toolsets = context.toolsets + if toolsets is None: + raise ValueError("NestedHarness requires toolsets.") + answers: list[str] = [] + for prompt in task.child_prompts: + result = await toolsets.call("nested_call_harness", {"prompt": str(prompt)}) + response = result.response + answer = str(response.messages[0].content) if response.messages else "" + answers.append(answer) + state.extras["child_answers"] = answers + answer = " ".join(answers) + message = vf.AssistantMessage(content=answer) + state.transcript.append( + vf.Turn(prompt=self.initial_messages(task), completion=[message]) + ) + state.stop("nested_completed") + + +def load_taskset(config: NestedTasksetConfig) -> NestedTaskset: + return NestedTaskset(config=config) + + +def load_harness(config: NestedHarnessConfig) -> NestedHarness: + return NestedHarness(config=config) diff --git a/environments/nested_harness_v1/pyproject.toml b/environments/nested_harness_v1/pyproject.toml index c5153396b3..2ccbe1e81f 100644 --- a/environments/nested_harness_v1/pyproject.toml +++ b/environments/nested_harness_v1/pyproject.toml @@ -14,7 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["nested_harness_v1.py", "pyproject.toml"] +include = ["nested_harness_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/openai_agents_env/openai_agents_env.py b/environments/openai_agents_env/openai_agents_env.py deleted file mode 100644 index 1bac945361..0000000000 --- a/environments/openai_agents_env/openai_agents_env.py +++ /dev/null @@ -1,141 +0,0 @@ -import re - -import verifiers as vf -from verifiers.utils.data_utils import load_example_dataset - -ANSWER_RE = re.compile(r"^\s*ANSWER\s*:?\s*(.+?)\s*$", re.IGNORECASE) - - -class OpenAIAgentsTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["answer_reward"] - taskset_id: str = "gsm8k-openai-agents" - num_train_examples: int = 50 - num_eval_examples: int = 20 - - -def calculate(expression: str) -> str: - """Evaluate a math expression and return the result.""" - try: - result = eval(expression, {"__builtins__": {}}, {}) - except Exception as exc: - return f"Error: {exc}" - return str(result) - - -async def run_openai_agents_program(task: vf.Task, state: vf.State) -> vf.State: - from agents import ( - Agent, - OpenAIChatCompletionsModel, - Runner, - function_tool, - set_tracing_disabled, - ) - - set_tracing_disabled(True) - endpoint_config = state.get_endpoint_config(api="chat") - client = state.get_client(api="chat") - model = OpenAIChatCompletionsModel( - model=endpoint_config.model, - openai_client=client, - ) - agent = Agent( - name="MathSolver", - instructions=( - "You are a math problem solver. Use the calculate tool to evaluate " - "expressions. Give your final numerical answer after the word ANSWER " - "on its own line, e.g.:\nANSWER: 42" - ), - model=model, - tools=[function_tool(calculate)], - ) - - question = task.get("question") - if question is not None: - query = str(question) - else: - query = "" - prompt = task.get("prompt") - if isinstance(prompt, list) and prompt: - query = str(vf.get_messages(prompt)[-1].content or "") - - result = await Runner.run(agent, input=query) - final_output = str(result.final_output) - state["agent_result"] = final_output - state["completion"] = [{"role": "assistant", "content": final_output}] - return state - - -def load_gsm8k_tasks(split: str, num_examples: int): - n = num_examples if num_examples > 0 else None - return load_example_dataset("gsm8k", split=split, n=n) - - -def load_tasks( - split: vf.TaskSplit = "train", - num_train_examples: int = 50, - num_eval_examples: int = 20, -): - dataset_split = "train" if split == "train" else "test" - num_examples = num_train_examples if split == "train" else num_eval_examples - return load_gsm8k_tasks(dataset_split, num_examples) - - -def extract_answer(text: str) -> str: - for line in reversed(text.splitlines()): - match = ANSWER_RE.match(line) - if match: - return match.group(1).strip() - return "" - - -def answers_match(agent_answer: str, answer: str) -> float: - try: - parsed_agent_answer = float(agent_answer.replace(",", "")) - parsed_answer = float(answer.replace(",", "")) - except (ValueError, TypeError): - return 1.0 if agent_answer.strip() == answer.strip() else 0.0 - return 1.0 if abs(parsed_agent_answer - parsed_answer) < 0.01 else 0.0 - - -def answer_reward(task: vf.Task, state: vf.State) -> float: - """Check if the agent's final output contains the correct answer.""" - result = state.get("agent_result") - if result is not None: - text = str(result) - else: - completion = state.get("completion") - messages = [] - if isinstance(completion, list): - messages = vf.get_messages(completion, role="assistant") or vf.get_messages( - completion - ) - text = str(messages[-1].content or "") if messages else "" - agent_answer = extract_answer(text) - if not agent_answer: - return 0.0 - return answers_match(agent_answer, str(task.get("answer", ""))) - - -class OpenAIAgentsTaskset(vf.Taskset[OpenAIAgentsTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks( - split=split, - num_train_examples=self.config.num_train_examples, - num_eval_examples=self.config.num_eval_examples, - ) - - -class OpenAIAgentsHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="run_openai_agents_program") - - -class OpenAIAgentsEnvConfig(vf.EnvConfig): - taskset: OpenAIAgentsTasksetConfig = OpenAIAgentsTasksetConfig() - harness: OpenAIAgentsHarnessConfig = OpenAIAgentsHarnessConfig() - - -def load_environment(config: OpenAIAgentsEnvConfig) -> vf.Env: - return vf.Env( - taskset=OpenAIAgentsTaskset(config=config.taskset), - harness=vf.Harness(config=config.harness), - ) diff --git a/environments/openai_agents_env/README.md b/environments/openai_agents_env_v1/README.md similarity index 80% rename from environments/openai_agents_env/README.md rename to environments/openai_agents_env_v1/README.md index 9a728a1c21..41a4379d26 100644 --- a/environments/openai_agents_env/README.md +++ b/environments/openai_agents_env_v1/README.md @@ -1,11 +1,11 @@ -# openai-agents-env +# openai-agents-env-v1 - + Source Code ### Overview -- **Environment ID**: `openai-agents-env` +- **Environment ID**: `openai-agents-env-v1` - **Short description**: V1 Taskset/Harness example using the OpenAI Agents SDK with a calculator tool on GSM8K math problems. - **Tags**: v1, taskset, harness, agents, tool-use, math, gsm8k @@ -19,19 +19,19 @@ - **Rubric overview**: Exact match on numeric answer extracted from `ANSWER: ` pattern ### How it works -The taskset owns GSM8K train/eval task loading and reward logic. The harness runs an in-process OpenAI Agents SDK program, builds its client from `state.get_endpoint_config(api="chat")`, and routes every model call through the V1 interception endpoint. +The taskset owns GSM8K train/eval task loading and reward logic. The harness runs an in-process OpenAI Agents SDK agent, starts the v1 protocol endpoint for the rollout, and gives the SDK an OpenAI-compatible client pointed at that endpoint. ### Quickstart Run an evaluation with default settings: ```bash -prime eval run openai-agents-env +prime eval run openai-agents-env-v1 ``` Configure model and sampling: ```bash -prime eval run openai-agents-env \ +prime eval run openai-agents-env-v1 \ -m gpt-4.1-mini \ -n 20 -r 3 -t 1024 -T 0.7 ``` diff --git a/environments/openai_agents_env_v1/openai_agents_env_v1/__init__.py b/environments/openai_agents_env_v1/openai_agents_env_v1/__init__.py new file mode 100644 index 0000000000..99c97e3a4f --- /dev/null +++ b/environments/openai_agents_env_v1/openai_agents_env_v1/__init__.py @@ -0,0 +1 @@ +"""openai-agents-env-v1 environment package.""" diff --git a/environments/openai_agents_env_v1/openai_agents_env_v1/harness.py b/environments/openai_agents_env_v1/openai_agents_env_v1/harness.py new file mode 100644 index 0000000000..b9c5190ca3 --- /dev/null +++ b/environments/openai_agents_env_v1/openai_agents_env_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import OpenAIAgentsHarness as OpenAIAgentsHarness +from .taskset import OpenAIAgentsHarnessConfig as OpenAIAgentsHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/openai_agents_env_v1/openai_agents_env_v1/taskset.py b/environments/openai_agents_env_v1/openai_agents_env_v1/taskset.py new file mode 100644 index 0000000000..e2f5b4e3b6 --- /dev/null +++ b/environments/openai_agents_env_v1/openai_agents_env_v1/taskset.py @@ -0,0 +1,167 @@ +import re + +import verifiers.v1 as vf +from verifiers.utils.data_utils import load_example_dataset + +ANSWER_RE = re.compile(r"^\s*ANSWER\s*:?\s*(.+?)\s*$", re.IGNORECASE) + + +class OpenAIAgentsTasksetConfig(vf.TasksetConfig): + id: str = "gsm8k-openai-agents" + num_train_examples: int = 50 + num_eval_examples: int = 20 + + +class OpenAIAgentsHarnessConfig(vf.HarnessConfig): + max_turns: int = 10 + + +class OpenAIAgentsTask(vf.Task): + question: str + answer: str + + +def calculate(expression: str) -> str: + """Evaluate a math expression and return the result.""" + try: + result = eval(expression, {"__builtins__": {}}, {}) + except Exception as exc: + return f"Error: {exc}" + return str(result) + + +def load_gsm8k_tasks(split: str, num_examples: int) -> vf.Tasks: + n = num_examples if num_examples > 0 else None + return [ + { + **row, + "row_id": index, + "prompt": [{"role": "user", "content": str(row["question"])}], + } + for index, row in enumerate(load_example_dataset("gsm8k", split=split, n=n)) + ] + + +def extract_answer(text: str) -> str: + for line in reversed(text.splitlines()): + match = ANSWER_RE.match(line) + if match: + return match.group(1).strip() + return "" + + +def answers_match(agent_answer: str, answer: str) -> float: + try: + parsed_agent_answer = float(agent_answer.replace(",", "")) + parsed_answer = float(answer.replace(",", "")) + except (ValueError, TypeError): + return float(agent_answer.strip() == answer.strip()) + return float(abs(parsed_agent_answer - parsed_answer) < 0.01) + + +def final_text(state: vf.State) -> str: + result = state.artifacts.get("agent_result") + if isinstance(result, str): + return result + messages = [ + message for message in state.completion if message.role == "assistant" + ] or state.completion + return str(messages[-1].content or "") if messages else "" + + +class OpenAIAgentsTaskset(vf.Taskset[OpenAIAgentsTasksetConfig]): + task_type = OpenAIAgentsTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + dataset_split = "train" if split == "train" else "test" + num_examples = ( + self.config.num_train_examples + if split == "train" + else self.config.num_eval_examples + ) + return load_gsm8k_tasks(dataset_split, num_examples) + + @vf.reward + async def answer_reward(self, task: OpenAIAgentsTask, state: vf.State) -> float: + answer = extract_answer(final_text(state)) + return answers_match(answer, task.answer) if answer else 0.0 + + +class OpenAIAgentsHarness(vf.Harness[OpenAIAgentsHarnessConfig]): + async def run_with_context(self, context: vf.Context) -> None: + task = OpenAIAgentsTask.model_validate(context.task.model_dump()) + state = context.state + runtime = context.runtime + if runtime is None: + raise ValueError("OpenAIAgentsHarness requires a runtime.") + prompt = self.initial_messages(task) + + async def stop_check() -> str | None: + if await self.is_completed(context): + return state.stop_condition or "stop" + return None + + async with vf.InterceptionServer( + context, + task, + state, + protocols=self.protocols, + stop_check=stop_check, + ) as endpoint: + endpoint_url = await runtime.expose(endpoint.port) + endpoint_env = endpoint.env(base_url=endpoint_url, model=context.model) + final_output = await run_openai_agents( + query=task.question, + base_url=endpoint_env["OPENAI_BASE_URL"], + api_key=endpoint_env["OPENAI_API_KEY"], + model=endpoint_env["OPENAI_MODEL"], + ) + + state.artifacts["agent_result"] = final_output + message = vf.AssistantMessage(content=final_output) + if not state.transcript: + state.transcript.append(vf.Turn(prompt=prompt, completion=[message])) + state.stop("agent_completed") + + +async def run_openai_agents( + *, + query: str, + base_url: str, + api_key: str, + model: str, +) -> str: + from agents import ( + Agent, + OpenAIChatCompletionsModel, + Runner, + function_tool, + set_tracing_disabled, + ) + from openai import AsyncOpenAI + + set_tracing_disabled(True) + client = AsyncOpenAI(base_url=base_url, api_key=api_key) + try: + agent = Agent( + name="MathSolver", + instructions=( + "You are a math problem solver. Use the calculate tool to evaluate " + "expressions. Give your final numerical answer after the word ANSWER " + "on its own line, e.g.:\nANSWER: 42" + ), + model=OpenAIChatCompletionsModel(model=model, openai_client=client), + tools=[function_tool(calculate)], + ) + result = await Runner.run(agent, input=query) + return str(result.final_output) + finally: + await client.close() + + +def load_taskset(config: OpenAIAgentsTasksetConfig) -> OpenAIAgentsTaskset: + return OpenAIAgentsTaskset(config=config) + + +def load_harness(config: OpenAIAgentsHarnessConfig) -> OpenAIAgentsHarness: + return OpenAIAgentsHarness(config=config) diff --git a/environments/openai_agents_env/pyproject.toml b/environments/openai_agents_env_v1/pyproject.toml similarity index 81% rename from environments/openai_agents_env/pyproject.toml rename to environments/openai_agents_env_v1/pyproject.toml index 631c1de266..79d6352698 100644 --- a/environments/openai_agents_env/pyproject.toml +++ b/environments/openai_agents_env_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "openai-agents-env" +name = "openai-agents-env-v1" description = "V1 Taskset/Harness environment using the OpenAI Agents SDK" tags = ["v1", "taskset", "harness", "tool-use", "openai-agents"] version = "0.1.0" @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["openai_agents_env.py", "pyproject.toml"] +include = ["openai_agents_env_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/opencode_harbor/README.md b/environments/opencode_harbor/README.md deleted file mode 100644 index bafef40316..0000000000 --- a/environments/opencode_harbor/README.md +++ /dev/null @@ -1,92 +0,0 @@ -# opencode-harbor - -### Overview -- **Environment ID**: `opencode-harbor` -- **Short description**: Environment for running an agent with OpenCode on Harbor tasks -- **Tags**: opencode, cli_agent, harbor - -### Datasets -- **Primary dataset(s)**: Harbor tasks -- **Source links**: -- **Split sizes**: 11 bundled tasks - -### Task -- **Type**: multiturn, cli_agent -- **Rubric overview**: Binary, returned by running task tests - -### Quickstart -Run the environment: - -```bash -prime eval run opencode-harbor -``` - -Configure model and sampling: - -```bash -prime eval run opencode-harbor -m openai/gpt-4.1-mini -n 20 -r 3 -t 1024 -T 0.7 -``` - -Notes: -- Use `-a` / `--env-args` for flat environment arguments. -- Use `taskset` and `harness` config sections for v1 object configuration. - -### Environment Arguments - -| Arg | Type | Default | Description | -| --- | ---- | ------- | ----------- | -| `task_names` | list[str] | `null` | Explicit Harbor task names to run. | -| `dataset` | str | `null` | Harbor Hub dataset id. Defaults to bundled `tasks/`. | - -OpenCode settings belong under the v1 harness config: - -```toml -[env.harness] -max_turns = 4 - -[env.harness.program] -agent_workdir = "/app" -``` - -This environment does not set a custom disabled-tool list. It inherits the -packaged `OpenCodeConfig` defaults. - -### Metrics -Summarize key metrics your rubric emits and how they’re interpreted. - -| Metric | Meaning | -| ------ | ------- | -| `reward` | Main scalar reward (weighted sum of criteria) | - - -## How It Works - -1. `HarborTaskset` loads Harbor tasks and contributes sandbox settings, - task uploads, env vars, and the Harbor reward. -2. `OpenCode` contributes the reusable OpenCode CLI program, install/setup, - intercepted endpoint config, MCP tool proxy, and log artifact collection. -3. The v1 runtime resolves both sides into one sandboxed command program at rollout time. -4. Reward is computed by running the Harbor test scripts after the rollout. - -`HarborTaskset` and `OpenCode` are packaged under `tasksets` and `harnesses` and -imported by the environment package. - -## Requirements - -- Harbor tasks directory with `task.toml` and `instruction.md` files -- Docker images specified in task configs - - -## Reward - -Uses Harbor's standard reward mechanism: - -- Runs `tests/test.sh` after agent completion -- Reads reward from `/logs/verifier/reward.txt` or `/logs/verifier/reward.json` -- Returns float reward value (typically 0 or 1) - -## Notes - -- OpenCode is installed at runtime. -- Agent logs are saved to `/logs/agent/opencode.txt` in the sandbox -- Uses `@ai-sdk/openai-compatible` provider for API interception diff --git a/environments/opencode_harbor/opencode_harbor.py b/environments/opencode_harbor/opencode_harbor.py deleted file mode 100644 index 35ae321268..0000000000 --- a/environments/opencode_harbor/opencode_harbor.py +++ /dev/null @@ -1,21 +0,0 @@ -import verifiers as vf -from harnesses import OpenCode, OpenCodeConfig -from tasksets import HarborTaskset, HarborTasksetConfig - - -def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: - taskset_config = config - if taskset_config.dataset is None and taskset_config.bundle_package is None: - taskset_config = taskset_config.model_copy(update={"bundle_package": __name__}) - return HarborTaskset(config=taskset_config) - - -def load_harness(config: OpenCodeConfig) -> OpenCode: - return OpenCode(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/opencode_harbor_v1/README.md b/environments/opencode_harbor_v1/README.md new file mode 100644 index 0000000000..a59185ab53 --- /dev/null +++ b/environments/opencode_harbor_v1/README.md @@ -0,0 +1,100 @@ +# opencode-harbor-v1 + +### Overview +- **Environment ID**: `opencode-harbor-v1` +- **Short description**: Environment for running an agent with OpenCode on Harbor tasks +- **Tags**: opencode, cli_agent, harbor + +### Datasets +- **Primary dataset(s)**: Harbor tasks +- **Source links**: +- **Split sizes**: 11 bundled tasks + +### Task +- **Type**: multiturn, cli_agent +- **Rubric overview**: Binary, returned by running task tests + +### Quickstart +Run the environment: + +```bash +prime eval run opencode-harbor-v1 +``` + +Configure model and sampling: + +```bash +prime eval run opencode-harbor-v1 -m openai/gpt-4.1-mini -n 20 -r 3 -t 1024 -T 0.7 +``` + +Notes: +- v1 task settings belong under `config.taskset` when passed through `-a` / `--env-args`. +- Use `taskset` and `harness` config sections for v1 object configuration in TOML. + +### Taskset Config + +| Arg | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `source` | `"harbor" \| "package"` | `"package"` in this environment | Dataset source resolver. | +| `dataset` | str | `"opencode_harbor_v1"` in this environment | Harbor dataset id or Python package name. | +| `tasks` | list[str] | `null` | Explicit Harbor task names to run. | +| `cache_dir` | str | `null` | Optional Harbor cache root override. | +| `refresh` | bool | `false` | Refresh Harbor cache before loading. | +| `require_image` | bool | `false` | Require every task to declare `[environment].docker_image`. | + +### Harness Config + +OpenCode settings belong under `config.harness`: + +```toml +[env.harness] +max_turns = 4 +version = "PrimeIntellect-ai/opencode@1.1.63-rl2" +cwd = "/app" +``` + +The harness also accepts the packaged `OpenCodeConfig` fields for `system_prompt`, +`log_path`, `disabled_tools`, `allow_git`, `disable_compaction`, +`provider_timeout_ms`, and runtime settings. + +### Metrics + +| Metric | Meaning | +| ------ | ------- | +| `reward` | Harbor verifier reward, usually `0.0` or `1.0` | +| `num_turns` | Number of intercepted assistant turns | + + +## How It Works + +1. `HarborTaskset` resolves Harbor or package task directories, maps + `[environment].docker_image` and resource hints onto generic v1 `Task` + fields, and owns the Harbor reward. +2. `OpenCode` contributes the reusable OpenCode CLI program, install/setup, + intercepted endpoint config, MCP tool proxy, and log artifact collection. +3. The v1 runtime resolves both sides into one sandboxed command program at rollout time. +4. Reward is computed by staging only `tests/` into the live runtime after the + rollout and running `tests/test.sh`. + +`HarborTaskset` and `OpenCode` are packaged under `tasksets` and `harnesses` and +imported by the environment package. + +## Requirements + +- Harbor tasks directory with `task.toml` and `instruction.md` files +- Docker images specified in task configs + + +## Reward + +Uses Harbor's standard reward mechanism: + +- Runs `tests/test.sh` after agent completion +- Reads reward from `/logs/verifier/reward.txt` +- Returns float reward value (typically 0 or 1) + +## Notes + +- OpenCode is installed at runtime. +- Agent logs are saved to `/logs/agent/opencode.txt` in the sandbox +- Uses `@ai-sdk/openai-compatible` provider for API interception diff --git a/environments/opencode_harbor_v1/opencode_harbor_v1/__init__.py b/environments/opencode_harbor_v1/opencode_harbor_v1/__init__.py new file mode 100644 index 0000000000..f77fdec5ab --- /dev/null +++ b/environments/opencode_harbor_v1/opencode_harbor_v1/__init__.py @@ -0,0 +1 @@ +"""opencode-harbor-v1 environment package.""" diff --git a/environments/opencode_harbor_v1/opencode_harbor_v1/harness.py b/environments/opencode_harbor_v1/opencode_harbor_v1/harness.py new file mode 100644 index 0000000000..7bda02f626 --- /dev/null +++ b/environments/opencode_harbor_v1/opencode_harbor_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import OpenCode as OpenCode +from .taskset import OpenCodeConfig as OpenCodeConfig +from .taskset import load_harness as load_harness diff --git a/environments/opencode_harbor/tasks/build-cython-ext/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/build-cython-ext/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/build-cython-ext/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/build-cython-ext/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/instruction.md diff --git a/environments/opencode_harbor/tasks/build-cython-ext/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/build-cython-ext/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/build-cython-ext/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/build-cython-ext/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/task.toml diff --git a/environments/opencode_harbor/tasks/build-cython-ext/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/build-cython-ext/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/tests/test.sh diff --git a/environments/opencode_harbor/tasks/build-cython-ext/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/build-cython-ext/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/build-cython-ext/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/chess-best-move/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/chess-best-move/environment/make.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/environment/make.py similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/environment/make.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/environment/make.py diff --git a/environments/opencode_harbor/tasks/chess-best-move/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/instruction.md diff --git a/environments/opencode_harbor/tasks/chess-best-move/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/chess-best-move/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/task.toml diff --git a/environments/opencode_harbor/tasks/chess-best-move/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/tests/test.sh diff --git a/environments/opencode_harbor/tasks/chess-best-move/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/chess-best-move/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/chess-best-move/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/instruction.md diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/task.toml diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/tests/test.sh diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/configure-git-webserver/tests/verify.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/tests/verify.sh similarity index 100% rename from environments/opencode_harbor/tasks/configure-git-webserver/tests/verify.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/configure-git-webserver/tests/verify.sh diff --git a/environments/opencode_harbor/tasks/fix-code-vulnerability/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/fix-code-vulnerability/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/fix-code-vulnerability/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/fix-code-vulnerability/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/instruction.md diff --git a/environments/opencode_harbor/tasks/fix-code-vulnerability/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/fix-code-vulnerability/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/fix-code-vulnerability/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/fix-code-vulnerability/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/task.toml diff --git a/environments/opencode_harbor/tasks/fix-code-vulnerability/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/fix-code-vulnerability/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/tests/test.sh diff --git a/environments/opencode_harbor/tasks/fix-code-vulnerability/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/fix-code-vulnerability/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/fix-code-vulnerability/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/hello-world/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/hello-world/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/environment/Dockerfile diff --git a/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/instruction.md new file mode 100644 index 0000000000..11b3518817 --- /dev/null +++ b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/instruction.md @@ -0,0 +1 @@ +Create a file called hello.txt with "Hello, world!" as the content. diff --git a/environments/opencode_harbor/tasks/hello-world/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/hello-world/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/hello-world/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/hello-world/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/task.toml diff --git a/environments/opencode_harbor/tasks/hello-world/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/hello-world/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/tests/test.sh diff --git a/environments/opencode_harbor/tasks/hello-world/tests/test_state.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/tests/test_state.py similarity index 100% rename from environments/opencode_harbor/tasks/hello-world/tests/test_state.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/hello-world/tests/test_state.py diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/environment/log_generator_deterministic.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/environment/log_generator_deterministic.py similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/environment/log_generator_deterministic.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/environment/log_generator_deterministic.py diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/instruction.md diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/task.toml diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/tests/test.sh diff --git a/environments/opencode_harbor/tasks/log-summary-date-ranges/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/log-summary-date-ranges/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/log-summary-date-ranges/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/polyglot-c-py/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/polyglot-c-py/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/polyglot-c-py/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/polyglot-c-py/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/instruction.md diff --git a/environments/opencode_harbor/tasks/polyglot-c-py/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/polyglot-c-py/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/polyglot-c-py/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/polyglot-c-py/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/task.toml diff --git a/environments/opencode_harbor/tasks/polyglot-c-py/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/polyglot-c-py/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/tests/test.sh diff --git a/environments/opencode_harbor/tasks/polyglot-c-py/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/polyglot-c-py/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/polyglot-c-py/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/qemu-alpine-ssh/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/qemu-alpine-ssh/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/qemu-alpine-ssh/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/qemu-alpine-ssh/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/instruction.md diff --git a/environments/opencode_harbor/tasks/qemu-alpine-ssh/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/qemu-alpine-ssh/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/qemu-alpine-ssh/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/qemu-alpine-ssh/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/task.toml diff --git a/environments/opencode_harbor/tasks/qemu-alpine-ssh/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/qemu-alpine-ssh/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/tests/test.sh diff --git a/environments/opencode_harbor/tasks/qemu-alpine-ssh/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/qemu-alpine-ssh/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-alpine-ssh/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/qemu-startup/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/qemu-startup/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/qemu-startup/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/qemu-startup/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/instruction.md diff --git a/environments/opencode_harbor/tasks/qemu-startup/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/qemu-startup/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/qemu-startup/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/qemu-startup/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/task.toml diff --git a/environments/opencode_harbor/tasks/qemu-startup/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/qemu-startup/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/tests/test.sh diff --git a/environments/opencode_harbor/tasks/qemu-startup/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/qemu-startup/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/qemu-startup/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/regex-log/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/regex-log/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/regex-log/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/regex-log/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/instruction.md diff --git a/environments/opencode_harbor/tasks/regex-log/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/regex-log/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/regex-log/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/regex-log/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/task.toml diff --git a/environments/opencode_harbor/tasks/regex-log/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/regex-log/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/tests/test.sh diff --git a/environments/opencode_harbor/tasks/regex-log/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/regex-log/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/regex-log/tests/test_outputs.py diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/environment/Dockerfile b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/environment/Dockerfile similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/environment/Dockerfile rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/environment/Dockerfile diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz.sha256 b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz.sha256 similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz.sha256 rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/environment/vendor/sqlite-fossil-release.tar.gz.sha256 diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/instruction.md b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/instruction.md similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/instruction.md rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/instruction.md diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/solution/solve.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/solution/solve.sh similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/solution/solve.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/solution/solve.sh diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/task.toml b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/task.toml similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/task.toml rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/task.toml diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/tests/test.sh b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/tests/test.sh similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/tests/test.sh rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/tests/test.sh diff --git a/environments/opencode_harbor/tasks/sqlite-with-gcov/tests/test_outputs.py b/environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/tests/test_outputs.py similarity index 100% rename from environments/opencode_harbor/tasks/sqlite-with-gcov/tests/test_outputs.py rename to environments/opencode_harbor_v1/opencode_harbor_v1/tasks/sqlite-with-gcov/tests/test_outputs.py diff --git a/environments/opencode_harbor_v1/opencode_harbor_v1/taskset.py b/environments/opencode_harbor_v1/opencode_harbor_v1/taskset.py new file mode 100644 index 0000000000..fc829cabfc --- /dev/null +++ b/environments/opencode_harbor_v1/opencode_harbor_v1/taskset.py @@ -0,0 +1,18 @@ +from harnesses import OpenCode, OpenCodeConfig +from tasksets import HarborTaskset, HarborTasksetConfig + + +def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: + taskset_config = config + if ( + "source" not in config.model_fields_set + and "dataset" not in config.model_fields_set + ): + taskset_config = taskset_config.model_copy( + update={"source": "package", "dataset": "opencode_harbor_v1"} + ) + return HarborTaskset(config=taskset_config) + + +def load_harness(config: OpenCodeConfig) -> OpenCode: + return OpenCode(config=config) diff --git a/environments/opencode_harbor/pyproject.toml b/environments/opencode_harbor_v1/pyproject.toml similarity index 78% rename from environments/opencode_harbor/pyproject.toml rename to environments/opencode_harbor_v1/pyproject.toml index 4d596996cc..91d32b1306 100644 --- a/environments/opencode_harbor/pyproject.toml +++ b/environments/opencode_harbor_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "opencode-harbor" +name = "opencode-harbor-v1" description = "OpenCode v1 taskset/harness agent on Harbor tasks" license = "MIT" tags = ["eval", "cli_agent", "v1", "taskset", "harness"] @@ -19,11 +19,11 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["opencode_harbor.py", "pyproject.toml", "tasks/**/*"] +include = ["opencode_harbor_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 rollouts_per_example = 3 [tool.ruff] -exclude = ["tasks/**"] +exclude = ["opencode_harbor_v1/tasks/**"] diff --git a/environments/openenv_echo/openenv_echo.py b/environments/openenv_echo/openenv_echo.py deleted file mode 100644 index 85b683d2ad..0000000000 --- a/environments/openenv_echo/openenv_echo.py +++ /dev/null @@ -1,66 +0,0 @@ -from collections.abc import Mapping -from typing import cast - -import verifiers as vf -from tasksets.openenv import OpenEnvTaskset, OpenEnvTasksetConfig -from verifiers.types import Messages, UserMessage -from verifiers.utils.message_utils import MessageInput, normalize_messages - - -class OpenEnvEchoTasksetConfig(OpenEnvTasksetConfig): - prompt_renderer: str = "openenv_echo:render_openenv_prompt" - - -def render_openenv_prompt( - observation: object, - *, - action_schema: vf.ConfigData | None = None, - context: str = "reset", - contract: str = "mcp", - seed: int = 0, -) -> Messages: - del contract, seed - if not isinstance(observation, Mapping): - raise RuntimeError( - f"openenv-echo prompt renderer expected dict observation, got {type(observation).__name__}." - ) - observation_data = cast(vf.ConfigData, observation) - - messages = observation_data.get("messages") - if isinstance(messages, list) and messages: - try: - return normalize_messages( - cast(MessageInput, messages), - field_name="openenv-echo observation messages", - ) - except TypeError as e: - raise RuntimeError(str(e)) from e - - prompt = observation_data.get("prompt") - if isinstance(prompt, str) and prompt.strip(): - return [UserMessage(content=prompt)] - - if context == "reset" and isinstance(action_schema, dict): - return [ - UserMessage( - content=( - "You are connected to an OpenEnv MCP environment. " - "Call at least one tool before your final response. " - "Action contract: call_tool(tool_name: str, arguments: object)." - ) - ) - ] - - raise RuntimeError("openenv-echo observation did not include a renderable prompt.") - - -def load_taskset(config: OpenEnvEchoTasksetConfig) -> OpenEnvTaskset: - return OpenEnvTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/openenv_echo/README.md b/environments/openenv_echo_v1/README.md similarity index 80% rename from environments/openenv_echo/README.md rename to environments/openenv_echo_v1/README.md index 9ab5f51ada..a75c7d9a73 100644 --- a/environments/openenv_echo/README.md +++ b/environments/openenv_echo_v1/README.md @@ -1,12 +1,12 @@ -# openenv-echo +# openenv-echo-v1 - + Source Code ### Overview -- **Environment ID**: `openenv-echo` +- **Environment ID**: `openenv-echo-v1` - **Short description**: OpenEnv Echo environment via `OpenEnvTaskset`, demonstrating MCP tool-calling in Prime Sandboxes. - **Tags**: openenv, mcp, tools, example @@ -27,10 +27,10 @@ Build and register the bundled OpenEnv Docker image in the Prime registry: ```bash -uv run vf-build openenv-echo +uv run vf-build openenv-echo-v1 ``` -This writes `environments/openenv_echo/proj/.build.json` with the fully qualified image reference and runtime metadata. +This writes `environments/openenv_echo_v1/openenv_echo_v1/proj/.build.json` with the fully qualified image reference and runtime metadata. Verify the image is ready (status **Ready** or **Completed**): @@ -41,22 +41,22 @@ prime images list Run an evaluation with default settings: ```bash -prime eval run openenv-echo +prime eval run openenv-echo-v1 ``` Configure model and sampling: ```bash -prime eval run openenv-echo \ +prime eval run openenv-echo-v1 \ -m openai/gpt-4.1-mini \ -n 20 -r 3 -t 1024 -T 0.7 ``` Notes: - If your environments directory is not `./environments`, run: -`uv run vf-build openenv-echo -p /path/to/environments` -- If you customize the bundled OpenEnv project, rerun `uv run vf-build openenv-echo` (the `proj/.build.json` manifest is updated). -- `openenv_echo.py` defines `render_openenv_prompt` and passes it via `prompt_renderer` to keep the initial MCP prompt concise. +`uv run vf-build openenv-echo-v1 -p /path/to/environments` +- If you customize the bundled OpenEnv project, rerun `uv run vf-build openenv-echo-v1` (the `proj/.build.json` manifest is updated). +- `openenv_echo_v1/taskset.py` defines `render_openenv_prompt` and passes it via `prompt_renderer` to keep the initial MCP prompt concise. ### Troubleshooting diff --git a/environments/openenv_echo_v1/openenv_echo_v1/__init__.py b/environments/openenv_echo_v1/openenv_echo_v1/__init__.py new file mode 100644 index 0000000000..17102f2c71 --- /dev/null +++ b/environments/openenv_echo_v1/openenv_echo_v1/__init__.py @@ -0,0 +1 @@ +"""openenv-echo-v1 environment package.""" diff --git a/environments/openenv_echo/proj/.build.json b/environments/openenv_echo_v1/openenv_echo_v1/proj/.build.json similarity index 68% rename from environments/openenv_echo/proj/.build.json rename to environments/openenv_echo_v1/openenv_echo_v1/proj/.build.json index 4537d421f9..97fa7d8fa7 100644 --- a/environments/openenv_echo/proj/.build.json +++ b/environments/openenv_echo_v1/openenv_echo_v1/proj/.build.json @@ -1,8 +1,8 @@ { "app": "server.app:app", "contract": "mcp", - "environment_id": "openenv-echo", - "image": "cmaeni8ji0001ql2z5gw8204f/openenv-echo:latest", + "environment_id": "openenv-echo-v1", + "image": "team-cmlr3u2er002zhr01tj8f48ts/openenv-echo-v1:latest", "image_status": "COMPLETED", "port": 8000, "schema_version": 1, diff --git a/environments/openenv_echo/proj/README.md b/environments/openenv_echo_v1/openenv_echo_v1/proj/README.md similarity index 100% rename from environments/openenv_echo/proj/README.md rename to environments/openenv_echo_v1/openenv_echo_v1/proj/README.md diff --git a/environments/openenv_echo/proj/__init__.py b/environments/openenv_echo_v1/openenv_echo_v1/proj/__init__.py similarity index 100% rename from environments/openenv_echo/proj/__init__.py rename to environments/openenv_echo_v1/openenv_echo_v1/proj/__init__.py diff --git a/environments/openenv_echo/proj/client.py b/environments/openenv_echo_v1/openenv_echo_v1/proj/client.py similarity index 100% rename from environments/openenv_echo/proj/client.py rename to environments/openenv_echo_v1/openenv_echo_v1/proj/client.py diff --git a/environments/openenv_echo/proj/openenv.yaml b/environments/openenv_echo_v1/openenv_echo_v1/proj/openenv.yaml similarity index 100% rename from environments/openenv_echo/proj/openenv.yaml rename to environments/openenv_echo_v1/openenv_echo_v1/proj/openenv.yaml diff --git a/environments/openenv_echo/proj/pyproject.toml b/environments/openenv_echo_v1/openenv_echo_v1/proj/pyproject.toml similarity index 100% rename from environments/openenv_echo/proj/pyproject.toml rename to environments/openenv_echo_v1/openenv_echo_v1/proj/pyproject.toml diff --git a/environments/openenv_echo/proj/server/Dockerfile b/environments/openenv_echo_v1/openenv_echo_v1/proj/server/Dockerfile similarity index 100% rename from environments/openenv_echo/proj/server/Dockerfile rename to environments/openenv_echo_v1/openenv_echo_v1/proj/server/Dockerfile diff --git a/environments/openenv_echo/proj/server/__init__.py b/environments/openenv_echo_v1/openenv_echo_v1/proj/server/__init__.py similarity index 100% rename from environments/openenv_echo/proj/server/__init__.py rename to environments/openenv_echo_v1/openenv_echo_v1/proj/server/__init__.py diff --git a/environments/openenv_echo/proj/server/app.py b/environments/openenv_echo_v1/openenv_echo_v1/proj/server/app.py similarity index 94% rename from environments/openenv_echo/proj/server/app.py rename to environments/openenv_echo_v1/openenv_echo_v1/proj/server/app.py index cbc772dccb..0c5067fd1b 100644 --- a/environments/openenv_echo/proj/server/app.py +++ b/environments/openenv_echo_v1/openenv_echo_v1/proj/server/app.py @@ -37,7 +37,11 @@ # Pass the class (factory) instead of an instance for WebSocket session support # Use MCP types for action/observation since this is a pure MCP environment app = create_app( - EchoEnvironment, CallToolAction, CallToolObservation, env_name="echo_env" + EchoEnvironment, + CallToolAction, + CallToolObservation, + env_name="echo_env", + max_concurrent_envs=64, ) diff --git a/environments/openenv_echo/proj/server/echo_environment.py b/environments/openenv_echo_v1/openenv_echo_v1/proj/server/echo_environment.py similarity index 77% rename from environments/openenv_echo/proj/server/echo_environment.py rename to environments/openenv_echo_v1/openenv_echo_v1/proj/server/echo_environment.py index cbb7c5a165..d458085afd 100644 --- a/environments/openenv_echo/proj/server/echo_environment.py +++ b/environments/openenv_echo_v1/openenv_echo_v1/proj/server/echo_environment.py @@ -35,11 +35,13 @@ try: # In-repo imports (when running from OpenEnv repository) from openenv.core.env_server.mcp_environment import MCPEnvironment - from openenv.core.env_server.types import Action, Observation, State + from openenv.core.env_server.mcp_types import CallToolObservation + from openenv.core.env_server.types import Action, State except ImportError: # Standalone imports (when environment is standalone with openenv from pip) from openenv.core.env_server.mcp_environment import MCPEnvironment - from openenv.core.env_server.types import Action, Observation, State + from openenv.core.env_server.mcp_types import CallToolObservation + from openenv.core.env_server.types import Action, State from fastmcp import FastMCP @@ -66,6 +68,8 @@ class EchoEnvironment(MCPEnvironment): ... print(result) """ + SUPPORTS_CONCURRENT_SESSIONS = True + def __init__(self): """Initialize the echo environment with MCP server and tools.""" # Create MCP server and define tools inline @@ -107,7 +111,7 @@ def reset( seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs: Any, - ) -> Observation: + ) -> CallToolObservation: """ Reset the environment. @@ -125,7 +129,9 @@ def reset( ) self._reset_count += 1 - return Observation( + return CallToolObservation( + tool_name="reset", + result={"status": "ready", "message": "Echo environment ready!"}, done=False, reward=0.0, metadata={"status": "ready", "message": "Echo environment ready!"}, @@ -136,7 +142,7 @@ def _step_impl( action: Action, timeout_s: Optional[float] = None, **kwargs: Any, - ) -> Observation: + ) -> CallToolObservation: """ Handle non-MCP actions. @@ -151,7 +157,9 @@ def _step_impl( Returns: Observation with error for unknown action types """ - return Observation( + return CallToolObservation( + tool_name=type(action).__name__, + result=None, done=False, reward=0.0, metadata={ @@ -165,7 +173,7 @@ def step( action: Action, timeout_s: Optional[float] = None, **kwargs: Any, - ) -> Observation: + ) -> CallToolObservation: """ Execute a step in the environment. @@ -183,7 +191,32 @@ def step( self._state.step_count += 1 # Let the base class handle MCP actions and non-MCP routing - return super().step(action, timeout_s=timeout_s, **kwargs) + observation = super().step(action, timeout_s=timeout_s, **kwargs) + return self._with_echo_reward(action, observation) + + async def step_async( + self, + action: Action, + timeout_s: Optional[float] = None, + **kwargs: Any, + ) -> CallToolObservation: + self._state.step_count += 1 + observation = await super().step_async(action, timeout_s=timeout_s, **kwargs) + return self._with_echo_reward(action, observation) + + @staticmethod + def _with_echo_reward( + action: Action, observation: CallToolObservation + ) -> CallToolObservation: + if getattr(action, "tool_name", None) not in { + "echo_message", + "echo_with_length", + }: + return observation + arguments = getattr(action, "arguments", {}) + message = arguments.get("message") if isinstance(arguments, dict) else None + reward = len(message) * 0.1 if isinstance(message, str) else 0.0 + return observation.model_copy(update={"reward": reward}) @property def state(self) -> State: diff --git a/environments/openenv_echo_v1/openenv_echo_v1/taskset.py b/environments/openenv_echo_v1/openenv_echo_v1/taskset.py new file mode 100644 index 0000000000..f5a247f58b --- /dev/null +++ b/environments/openenv_echo_v1/openenv_echo_v1/taskset.py @@ -0,0 +1,66 @@ +from collections.abc import Mapping +from typing import cast + +import verifiers.v1 as vf +from tasksets.openenv import OpenEnvTaskset, OpenEnvTasksetConfig + + +class OpenEnvEchoTasksetConfig(OpenEnvTasksetConfig): + prompt_renderer: str = "openenv_echo_v1.taskset:render_openenv_prompt" + + +def render_openenv_prompt( + observation: object, + *, + action_schema: vf.JsonData | None = None, + context: str = "reset", + contract: str = "mcp", + seed: int = 0, +) -> vf.Messages: + del contract, seed + if not isinstance(observation, Mapping): + raise RuntimeError( + f"openenv-echo prompt renderer expected dict observation, got {type(observation).__name__}." + ) + observation_data = cast(vf.JsonData, observation) + + messages = observation_data.get("messages") + if isinstance(messages, list) and messages: + parsed: vf.Messages = [] + for message in messages: + if not isinstance(message, Mapping): + raise RuntimeError("openenv-echo observation messages must be objects.") + payload = dict(message) + role = payload.get("role") + if role == "user": + parsed.append(vf.UserMessage.model_validate(payload)) + elif role == "assistant": + parsed.append(vf.AssistantMessage.model_validate(payload)) + elif role == "system": + parsed.append(vf.SystemMessage.model_validate(payload)) + elif role == "tool": + parsed.append(vf.ToolMessage.model_validate(payload)) + else: + raise RuntimeError(f"Unsupported openenv-echo message role: {role!r}.") + return parsed + + prompt = observation_data.get("prompt") + if isinstance(prompt, str) and prompt.strip(): + return [vf.UserMessage(content=prompt)] + + if context == "reset" and isinstance(action_schema, dict): + return [ + vf.UserMessage( + content=( + "You are connected to an OpenEnv MCP environment. " + "Call the echo_message tool with message='hello from openenv', " + "then answer with the echoed message." + ) + ) + ] + + raise RuntimeError("openenv-echo observation did not include a renderable prompt.") + + +def load_taskset(config: OpenEnvEchoTasksetConfig) -> OpenEnvTaskset: + return OpenEnvTaskset(config=config) diff --git a/environments/openenv_echo/pyproject.toml b/environments/openenv_echo_v1/pyproject.toml similarity index 75% rename from environments/openenv_echo/pyproject.toml rename to environments/openenv_echo_v1/pyproject.toml index da09519a45..0c03705e22 100644 --- a/environments/openenv_echo/pyproject.toml +++ b/environments/openenv_echo_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "openenv-echo" +name = "openenv-echo-v1" description = "OpenEnv Echo environment via the v1 OpenEnv taskset" tags = ["openenv", "mcp", "tools", "example"] version = "0.1.0" @@ -14,13 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = [ - "openenv_echo.py", - "pyproject.toml", - "README.md", - "proj/**/*", - "proj/.build.json", -] +include = ["openenv_echo_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/openenv_textarena/README.md b/environments/openenv_textarena_v1/README.md similarity index 80% rename from environments/openenv_textarena/README.md rename to environments/openenv_textarena_v1/README.md index fd084ba89a..f2575ac1bf 100644 --- a/environments/openenv_textarena/README.md +++ b/environments/openenv_textarena_v1/README.md @@ -1,12 +1,12 @@ -# openenv-textarena +# openenv-textarena-v1 - + Source Code ### Overview -- **Environment ID**: `openenv-textarena` +- **Environment ID**: `openenv-textarena-v1` - **Short description**: OpenEnv TextArena gym integration (default game: `Wordle-v0`) via `OpenEnvTaskset`. - **Tags**: openenv, gym, textarena, wordle, example @@ -27,13 +27,13 @@ Build and register the bundled OpenEnv Docker image in the Prime registry: ```bash -uv run vf-build openenv-textarena +uv run vf-build openenv-textarena-v1 ``` Run an evaluation with default settings: ```bash -prime eval run openenv-textarena +prime eval run openenv-textarena-v1 ``` ### Taskset Config @@ -48,4 +48,4 @@ prime eval run openenv-textarena - Upstream TextArena app defaults to `TEXTARENA_ENV_ID=Wordle-v0`. - To use another game, set environment variables in the OpenEnv project/server config before building. -- `openenv_textarena.py` defines `render_textarena_prompt` and passes it via `prompt_renderer` so observations are rendered as useful game messages. +- `openenv_textarena_v1/taskset.py` defines `render_textarena_prompt` and passes it via `prompt_renderer` so observations are rendered as useful game messages. diff --git a/environments/openenv_textarena_v1/openenv_textarena_v1/__init__.py b/environments/openenv_textarena_v1/openenv_textarena_v1/__init__.py new file mode 100644 index 0000000000..c49c2f3de4 --- /dev/null +++ b/environments/openenv_textarena_v1/openenv_textarena_v1/__init__.py @@ -0,0 +1 @@ +"""openenv-textarena-v1 environment package.""" diff --git a/environments/openenv_textarena/proj/.build.json b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/.build.json similarity index 66% rename from environments/openenv_textarena/proj/.build.json rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/.build.json index a76f8dca22..f6c7210dfb 100644 --- a/environments/openenv_textarena/proj/.build.json +++ b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/.build.json @@ -1,8 +1,8 @@ { "app": "server.app:app", "contract": "gym", - "environment_id": "openenv-textarena", - "image": "cmaeni8ji0001ql2z5gw8204f/openenv-textarena:latest", + "environment_id": "openenv-textarena-v1", + "image": "team-cmlr3u2er002zhr01tj8f48ts/openenv-textarena-v1:latest", "image_status": "COMPLETED", "port": 8000, "schema_version": 1, diff --git a/environments/openenv_textarena/proj/README.md b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/README.md similarity index 100% rename from environments/openenv_textarena/proj/README.md rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/README.md diff --git a/environments/openenv_textarena/proj/__init__.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/__init__.py similarity index 100% rename from environments/openenv_textarena/proj/__init__.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/__init__.py diff --git a/environments/openenv_textarena/proj/client.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/client.py similarity index 100% rename from environments/openenv_textarena/proj/client.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/client.py diff --git a/environments/openenv_textarena/proj/models.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/models.py similarity index 100% rename from environments/openenv_textarena/proj/models.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/models.py diff --git a/environments/openenv_textarena/proj/openenv.yaml b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/openenv.yaml similarity index 100% rename from environments/openenv_textarena/proj/openenv.yaml rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/openenv.yaml diff --git a/environments/openenv_textarena/proj/pyproject.toml b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/pyproject.toml similarity index 100% rename from environments/openenv_textarena/proj/pyproject.toml rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/pyproject.toml diff --git a/environments/openenv_textarena/proj/rewards.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/rewards.py similarity index 100% rename from environments/openenv_textarena/proj/rewards.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/rewards.py diff --git a/environments/openenv_textarena/proj/server/Dockerfile b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/Dockerfile similarity index 100% rename from environments/openenv_textarena/proj/server/Dockerfile rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/Dockerfile diff --git a/environments/openenv_textarena/proj/server/__init__.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/__init__.py similarity index 100% rename from environments/openenv_textarena/proj/server/__init__.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/__init__.py diff --git a/environments/openenv_textarena/proj/server/app.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/app.py similarity index 98% rename from environments/openenv_textarena/proj/server/app.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/app.py index e4692781b6..968bd760e8 100644 --- a/environments/openenv_textarena/proj/server/app.py +++ b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/app.py @@ -59,6 +59,7 @@ def create_textarena_environment(): TextArenaAction, TextArenaObservation, env_name="textarena_env", + max_concurrent_envs=64, ) diff --git a/environments/openenv_textarena/proj/server/environment.py b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/environment.py similarity index 99% rename from environments/openenv_textarena/proj/server/environment.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/environment.py index a3aa03549f..4cb6c1839f 100644 --- a/environments/openenv_textarena/proj/server/environment.py +++ b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/environment.py @@ -80,6 +80,8 @@ def _import_textarena() -> Any: class TextArenaEnvironment(Environment): """Wrap any TextArena game behind the OpenEnv ``Environment`` API.""" + SUPPORTS_CONCURRENT_SESSIONS = True + def __init__( self, env_id: str = "Wordle-v0", diff --git a/environments/openenv_textarena/proj/server/run_local.sh b/environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/run_local.sh similarity index 100% rename from environments/openenv_textarena/proj/server/run_local.sh rename to environments/openenv_textarena_v1/openenv_textarena_v1/proj/server/run_local.sh diff --git a/environments/openenv_textarena/openenv_textarena.py b/environments/openenv_textarena_v1/openenv_textarena_v1/taskset.py similarity index 69% rename from environments/openenv_textarena/openenv_textarena.py rename to environments/openenv_textarena_v1/openenv_textarena_v1/taskset.py index 8e7b618438..3bdb3ebd6c 100644 --- a/environments/openenv_textarena/openenv_textarena.py +++ b/environments/openenv_textarena_v1/openenv_textarena_v1/taskset.py @@ -2,7 +2,7 @@ from collections.abc import Mapping from typing import cast -import verifiers as vf +import verifiers.v1 as vf from tasksets.openenv import OpenEnvTaskset, OpenEnvTasksetConfig from verifiers.types import Messages, UserMessage @@ -11,14 +11,14 @@ class OpenEnvTextArenaTasksetConfig(OpenEnvTasksetConfig): - prompt_renderer: str = "openenv_textarena:render_textarena_prompt" + prompt_renderer: str = "openenv_textarena_v1.taskset:render_textarena_prompt" def render_textarena_prompt( observation: object, *, context: str = "reset", - action_schema: vf.ConfigData | None = None, + action_schema: vf.JsonData | None = None, contract: str = "gym", seed: int = 0, ) -> Messages: @@ -27,41 +27,45 @@ def render_textarena_prompt( raise RuntimeError( f"openenv-textarena prompt renderer expected dict observation, got {type(observation).__name__}." ) - observation_data = cast(vf.ConfigData, observation) + observation_data = cast(vf.JsonData, observation) message_text = textarena_message_text(observation_data) prompt_text = textarena_prompt_text(observation_data) + action_instruction = ( + '\n\nReturn only a JSON object like {"message": "[guess]"}. ' + "Do not include markdown or any other text." + ) if context == "step": if message_text is not None: - return [UserMessage(content=message_text)] + return [UserMessage(content=message_text + action_instruction)] if prompt_text is not None: - return [UserMessage(content=prompt_text)] + return [UserMessage(content=prompt_text + action_instruction)] else: if prompt_text is not None: - return [UserMessage(content=prompt_text)] + return [UserMessage(content=prompt_text + action_instruction)] if message_text is not None: - return [UserMessage(content=message_text)] + return [UserMessage(content=message_text + action_instruction)] raise RuntimeError( "openenv-textarena observation did not include renderable prompt text." ) -def textarena_message_text(observation: vf.ConfigData) -> str | None: +def textarena_message_text(observation: vf.JsonData) -> str | None: raw_messages = observation.get("messages") if not isinstance(raw_messages, list): return None for item in reversed(raw_messages): if isinstance(item, Mapping): - message = cast(vf.ConfigData, item) + message = cast(vf.JsonData, item) content = message.get("content") if isinstance(content, str) and content.strip(): return content.strip() return None -def textarena_prompt_text(observation: vf.ConfigData) -> str | None: +def textarena_prompt_text(observation: vf.JsonData) -> str | None: prompt = observation.get("prompt") if not isinstance(prompt, str): return None @@ -77,11 +81,3 @@ def textarena_prompt_text(observation: vf.ConfigData) -> str | None: def load_taskset(config: OpenEnvTextArenaTasksetConfig) -> OpenEnvTaskset: return OpenEnvTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/openenv_textarena/pyproject.toml b/environments/openenv_textarena_v1/pyproject.toml similarity index 74% rename from environments/openenv_textarena/pyproject.toml rename to environments/openenv_textarena_v1/pyproject.toml index 447fb55bae..493da022a9 100644 --- a/environments/openenv_textarena/pyproject.toml +++ b/environments/openenv_textarena_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "openenv-textarena" +name = "openenv-textarena-v1" description = "OpenEnv TextArena (Wordle-v0) environment via the v1 OpenEnv taskset" tags = ["openenv", "gym", "textarena", "wordle", "example"] version = "0.1.0" @@ -14,13 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = [ - "openenv_textarena.py", - "pyproject.toml", - "README.md", - "proj/**/*", - "proj/.build.json", -] +include = ["openenv_textarena_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/reverse_text/pyproject.toml b/environments/reverse_text/pyproject.toml index a97e8e2a87..09080fe712 100644 --- a/environments/reverse_text/pyproject.toml +++ b/environments/reverse_text/pyproject.toml @@ -14,7 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["reverse_text.py", "reverse_text_v1.py"] +include = ["reverse_text.py"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/reverse_text/reverse_text.py b/environments/reverse_text/reverse_text.py index cde3b10093..926bced553 100644 --- a/environments/reverse_text/reverse_text.py +++ b/environments/reverse_text/reverse_text.py @@ -8,24 +8,7 @@ def load_environment( dataset_split: str = "train", system_prompt: str | None = "Reverse the text character-by-character. Put your answer in tags.", - v1: bool = False, ) -> vf.Environment: - if v1: - from reverse_text_v1 import ( - ReverseTextTasksetConfig, - load_environment as load_v1, - ) - - return load_v1( - config=vf.EnvConfig( - taskset=ReverseTextTasksetConfig( - dataset_name=dataset_name, - dataset_split=dataset_split, - system_prompt=system_prompt, - ) - ) - ) - def build_dataset(): train_dataset = load_dataset(dataset_name, split=dataset_split).map( lambda x: { diff --git a/environments/reverse_text_v1/README.md b/environments/reverse_text_v1/README.md new file mode 100644 index 0000000000..4230fb718f --- /dev/null +++ b/environments/reverse_text_v1/README.md @@ -0,0 +1,49 @@ +# reverse-text-v1 + + +Source Code + + +### Overview +- **Environment ID**: `reverse-text-v1` +- **Short description**: Reverse a given text; evaluated by LCS similarity between the parsed answer and ground-truth reversal. +- **Tags**: text, transformation, single-turn, xml + +### Datasets +- **Primary dataset(s)**: `PrimeIntellect/Reverse-Text-RL` mapped to question/answer pairs +- **Source links**: [PrimeIntellect/Reverse-Text-RL](https://huggingface.co/datasets/PrimeIntellect/Reverse-Text-RL) +- **Split sizes**: Uses `train` split + +### Task +- **Type**: single-turn +- **Rubric overview**: LCS ratio between parsed answer and target reversed text + +### Quickstart +Run an evaluation with default settings: + +```bash +prime eval run reverse-text-v1 +``` + +Configure model and sampling: + +```bash +prime eval run reverse-text-v1 \ + -m openai/gpt-4.1-mini \ + -n 20 -r 3 -t 1024 -T 0.7 +``` + +Notes: +- v1 task settings belong under `config.taskset` when passed through `-a` / `--env-args`. + +### Taskset Config +| Arg | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `dataset_name` | str | `"PrimeIntellect/Reverse-Text-RL"` | Name of the dataset to use | +| `dataset_split` | str | `"train"` | Split of the dataset to use | +| `system_prompt` | str | `"Reverse the text character-by-character. Put your answer in tags."` | System prompt to use | + +### Metrics +| Metric | Meaning | +| ------ | ------- | +| `reward` | LCS similarity between reversed text and parsed answer | diff --git a/environments/reverse_text_v1/pyproject.toml b/environments/reverse_text_v1/pyproject.toml new file mode 100644 index 0000000000..02be3c077a --- /dev/null +++ b/environments/reverse_text_v1/pyproject.toml @@ -0,0 +1,21 @@ +[project] +name = "reverse-text-v1" +version = "0.1.4" +tags = ["text", "transformation", "single-turn", "xml"] +license = "Apache-2.0" +description = "Reverse a given text; evaluated by LCS similarity between the parsed answer and ground-truth reversal." +dependencies = [ + "verifiers>=0.1.5.post0", + "datasets", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build] +include = ["reverse_text_v1/**/*", "README.md", "pyproject.toml"] + +[tool.verifiers.eval] +num_examples = 5 +rollouts_per_example = 3 diff --git a/environments/reverse_text_v1/reverse_text_v1/__init__.py b/environments/reverse_text_v1/reverse_text_v1/__init__.py new file mode 100644 index 0000000000..3e195c049c --- /dev/null +++ b/environments/reverse_text_v1/reverse_text_v1/__init__.py @@ -0,0 +1 @@ +"""reverse-text-v1 environment package.""" diff --git a/environments/reverse_text/reverse_text_v1.py b/environments/reverse_text_v1/reverse_text_v1/taskset.py similarity index 68% rename from environments/reverse_text/reverse_text_v1.py rename to environments/reverse_text_v1/reverse_text_v1/taskset.py index 94874ce178..f0adee6a21 100644 --- a/environments/reverse_text/reverse_text_v1.py +++ b/environments/reverse_text_v1/reverse_text_v1/taskset.py @@ -3,15 +3,15 @@ from datasets import load_dataset -import verifiers as vf +import verifiers.v1 as vf class TagExtractor: def __init__(self, tag: str): self.pattern = re.compile(rf"<{tag}>(.*?)", re.DOTALL) - def __call__(self, completion: list[vf.ConfigData]) -> str: - messages = vf.get_messages(completion, role="assistant") + def __call__(self, completion: vf.Messages) -> str: + messages = [message for message in completion if message.role == "assistant"] if not messages: return "" message = messages[-1] @@ -31,13 +31,19 @@ class ReverseTextTasksetConfig(vf.TasksetConfig): ) +class ReverseTextTask(vf.Task): + question: str + answer: str + + class ReverseTextTaskset(vf.Taskset[ReverseTextTasksetConfig]): + task_type = ReverseTextTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: def map_row(row): return { "question": row["prompt"], "answer": row["prompt"][::-1], - "info": {}, } dataset = load_dataset( @@ -47,27 +53,17 @@ def map_row(row): dataset = dataset.remove_columns(["prompt"]) for index, row in enumerate(dataset): yield { - "example_id": index, + "row_id": index, "prompt": [{"role": "user", "content": row["question"]}], "question": row["question"], "answer": row["answer"], - "info": row.get("info") or {}, } @vf.reward(weight=1.0) - async def lcs_reward(self, task, state) -> float: - response = REVERSED_TEXT_EXTRACTOR(state.get("completion") or []) - answer = str(task["answer"]) - return SequenceMatcher(None, response, answer).ratio() + async def lcs_reward(self, task: ReverseTextTask, state: vf.State) -> float: + response = REVERSED_TEXT_EXTRACTOR(state.completion or []) + return SequenceMatcher(None, response, task.answer).ratio() def load_taskset(config: ReverseTextTasksetConfig) -> ReverseTextTaskset: return ReverseTextTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/rlm_swe_v1/README.md b/environments/rlm_swe_v1/README.md index d7389e4f70..6441cb9b42 100644 --- a/environments/rlm_swe_v1/README.md +++ b/environments/rlm_swe_v1/README.md @@ -1,40 +1,70 @@ # rlm-swe-v1 -v1 RLM coding environment using the R2E-Gym SWE taskset and packaged `RLM` -harness. +### Overview +- **Environment ID**: `rlm-swe-v1` +- **Short description**: v1 RLM coding environment on R2E-Gym SWE tasks +- **Tags**: rlm, swe, cli-agent, v1 -```python -import verifiers as vf +### Datasets +- **Primary dataset(s)**: `R2E-Gym/R2E-Gym-Subset` +- **Source links**: [R2E-Gym/R2E-Gym-Subset](https://huggingface.co/datasets/R2E-Gym/R2E-Gym-Subset) +- **Split sizes**: Uses the dataset `train` split by default -env = vf.load_environment("rlm-swe-v1") +### Task +- **Type**: multiturn, cli_agent +- **Rubric overview**: Runs each instance's hidden test command and parses pytest output for pass/fail reward + +### Quickstart +Run an evaluation with default settings: + +```bash +prime eval run rlm-swe-v1 ``` -Tune the taskset and harness through typed v1 config objects: - -```python -import verifiers as vf -from harnesses import RLMConfig, RLMProgramConfig -from rlm_swe_v1 import RlmSweTasksetConfig, load_environment - -env = load_environment( - config=vf.EnvConfig( - taskset=RlmSweTasksetConfig(timeout_minutes=90), - harness=RLMConfig( - program=RLMProgramConfig( - local_checkout="/path/to/checkout", - tools=["bash", "edit"], - ) - ), - ) -) +Configure model and sampling: + +```bash +prime eval run rlm-swe-v1 \ + -m openai/gpt-4.1-mini \ + -n 5 -r 1 -t 4096 -T 0.2 \ + -a '{"config": {"taskset": {"timeout_minutes": 90}, "harness": {"tools": ["bash", "edit"], "cwd": "/testbed"}}}' ``` -The taskset is fully implemented in this environment package on the v1 stack. -It loads the full `R2E-Gym/R2E-Gym-Subset` train split by default, converts each -row into a v1 task, creates the per-instance sandbox config from the dataset -image, stages hidden tests for scoring, runs `run_tests.sh`, and parses pytest -output for reward. +Notes: +- v1 task settings belong under `config.taskset`; reusable RLM agent settings belong under `config.harness`. +- The taskset is discovered from `taskset.py`; the harness is discovered from `harness.py`. + +### Taskset Config +| Arg | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `dataset_name` | str | `"R2E-Gym/R2E-Gym-Subset"` | Dataset to load | +| `repo_path` | str | `"/testbed"` | Repository path inside the sandbox | +| `filter_repos` | list[str] \| null | `null` | Repositories to exclude | +| `ds_num_proc` | int \| null | `null` | Dataset processing parallelism | +| `ds_keep_in_memory` | bool | `true` | Keep processed dataset rows in memory | +| `timeout_minutes` | int \| null | `null` | Override task runtime timeout | +| `env` | object \| null | `null` | Extra task program environment values | + +### Harness Config +| Arg | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `cwd` | str \| null | `"/testbed"` | Agent working directory | +| `tools` | list[str] | `["bash", "edit"]` | RLM tool names exposed to the agent | +| `exec_timeout` | int | `300` | RLM tool execution timeout | +| `max_depth` | int | `0` | RLM recursive agent depth | +| `summarize_at_tokens` | int \| null | `null` | Optional RLM summarization threshold | +| `append_to_system_prompt` | str | `""` | Additional RLM system prompt text | + +### Metrics +| Metric | Meaning | +| ------ | ------- | +| `reward` | Parsed hidden-test reward | +| `rlm_sub_llm_call_count` | RLM sub-model call count | +| `rlm_sub_llm_total_turns` | Total RLM sub-model turns | +| `rlm_sub_llm_total_tool_calls` | Total RLM sub-agent tool calls | -`RLM` owns the CLI program, intercepted endpoint config, RLM installation, and -trajectory filtering. Harbor is not used here because the R2E setup is dataset -and image backed rather than a Harbor task directory corpus. +### How It Works +1. `R2ESWETaskset` loads R2E-Gym rows and converts them into typed v1 tasks with sandbox/runtime metadata. +2. `RLM` owns the reusable CLI command, endpoint interception, and agent configuration. +3. The v1 runtime resolves task and harness runtime settings at rollout time. +4. Reward is computed from hidden test output after the agent finishes. diff --git a/environments/rlm_swe_v1/pyproject.toml b/environments/rlm_swe_v1/pyproject.toml index 8da0ca69c9..6eb83d9650 100644 --- a/environments/rlm_swe_v1/pyproject.toml +++ b/environments/rlm_swe_v1/pyproject.toml @@ -17,10 +17,8 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["rlm_swe_v1.py", "README.md", "pyproject.toml"] +include = ["rlm_swe_v1/**/*", "README.md", "pyproject.toml"] -[project.entry-points."verifiers.environments"] -rlm-swe-v1 = "rlm_swe_v1:load_environment" [tool.verifiers.eval] num_examples = 5 diff --git a/environments/rlm_swe_v1/rlm_swe_v1.py b/environments/rlm_swe_v1/rlm_swe_v1.py deleted file mode 100644 index fa0437344a..0000000000 --- a/environments/rlm_swe_v1/rlm_swe_v1.py +++ /dev/null @@ -1,515 +0,0 @@ -import json -import logging -import re -import shlex -from collections.abc import Mapping -from pathlib import Path -from typing import Protocol, cast - -from datasets import load_dataset -import verifiers as vf -from harnesses import RLM, RLMConfig, RLMProgramConfig - -logger = logging.getLogger(__name__) - -REGISTRY_PREFIX = "us-central1-docker.pkg.dev/prime-intellect-platform/prod-sandbox" -DEFAULT_DATASET_NAME = "R2E-Gym/R2E-Gym-Subset" -DEFAULT_REPO_PATH = "/testbed" -DEFAULT_ALT_PATH = "/root" -DEFAULT_RLM_TOOLS = ("bash", "edit") - - -class RlmSweTasksetConfig(vf.TasksetConfig): - taskset_id: str = "swe/r2e" - dataset_name: str = DEFAULT_DATASET_NAME - repo_path: str = DEFAULT_REPO_PATH - alt_path: str = DEFAULT_ALT_PATH - filter_repos: list[str] | None = None - ds_num_proc: int | None = None - ds_keep_in_memory: bool = True - timeout_minutes: int | None = None - hide_tests_from_agent: bool = True - env: vf.ConfigData | None = None - - -class SandboxCommandResult(Protocol): - exit_code: int - stdout: str | None - stderr: str | None - - -class R2ESandbox(Protocol): - id: str - - async def execute( - self, - command: str, - working_dir: str | None = None, - timeout: int = 90, - ) -> SandboxCommandResult: ... - - async def download_file( - self, remote_path: str, local_path: str, timeout: int = 300 - ) -> None: ... - - async def upload_file( - self, remote_path: str, local_path: str, timeout: int = 300 - ) -> None: ... - - async def upload_bytes(self, remote_path: str, data: bytes, name: str) -> None: ... - - async def run_background_job( - self, command: str, timeout: int, working_dir: str - ) -> SandboxCommandResult: ... - - -def load_tasks( - dataset_name: str = DEFAULT_DATASET_NAME, - repo_path: str = DEFAULT_REPO_PATH, - filter_repos: list[str] | None = None, - ds_num_proc: int | None = None, - ds_keep_in_memory: bool = True, - timeout_minutes: int | None = None, - env: vf.ConfigData | None = None, -) -> list[vf.JsonData]: - dataset_kwargs = dict( - num_proc=ds_num_proc, - keep_in_memory=ds_keep_in_memory, - load_from_cache_file=False, - ) - dataset = load_dataset( - dataset_name, - split="train", - keep_in_memory=ds_keep_in_memory, - num_proc=ds_num_proc, - ) - if filter_repos: - filter_set = frozenset(filter_repos) - dataset = dataset.filter( - lambda row: row.get("repo_name") not in filter_set, - **dataset_kwargs, - ) - dataset = dataset.map( - process_r2e_example, - remove_columns=dataset.column_names, - **dataset_kwargs, - ) - task_env = dict(env or {}) - rows: list[vf.JsonData] = [] - for index, row in enumerate(dataset): - row = dict(row) - info = dict(row["info"]) - instruction = str(info["problem_statement"]) - program_env = env_vars(repo_path=repo_path, env=task_env) - agent_path = program_env.pop("PATH", None) - if agent_path is not None: - program_env.setdefault("AGENT_PATH", agent_path) - program_env.setdefault("AGENT_WORKDIR", repo_path) - task_row: vf.JsonData = { - "example_id": index, - "task_id": info.get("instance_id") or index, - "question": row.get("question", instruction), - "instruction": instruction, - "prompt": [{"role": "user", "content": instruction}], - "answer": row.get("answer", ""), - "info": info, - "sandbox": sandbox_config( - info=info, - repo_path=repo_path, - timeout_minutes=timeout_minutes, - ), - "program": {"env": program_env}, - } - rows.append(task_row) - return rows - - -def sandbox_config( - *, info: vf.JsonData, repo_path: str, timeout_minutes: int | None -) -> vf.JsonData: - config: vf.JsonData = { - "image": f"{REGISTRY_PREFIX}/{info['docker_image']}", - "cpu_cores": 4, - "memory_gb": 4, - "disk_size_gb": 10, - "gpu_count": 0, - "workdir": repo_path, - "scope": "rollout", - } - if timeout_minutes is not None: - config["timeout_minutes"] = timeout_minutes - return config - - -def env_vars(*, repo_path: str, env: vf.ConfigData) -> dict[str, str]: - return { - "PATH": ( - f"/opt/miniconda3/bin:{repo_path}/.venv/bin:/root/.local/bin:" - "/root/.cargo/bin:/go/bin:/usr/local/go/bin:/usr/local/cargo:" - "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin" - ), - "PAGER": "cat", - "MANPAGER": "cat", - "LESS": "-R", - "PIP_PROGRESS_BAR": "off", - "TQDM_DISABLE": "1", - **{str(key): str(value) for key, value in env.items()}, - } - - -class R2ESWETaskset(vf.Taskset[RlmSweTasksetConfig]): - def sandbox_config(self, info: vf.JsonData) -> vf.JsonData: - return sandbox_config( - info=info, - repo_path=self.config.repo_path, - timeout_minutes=self.config.timeout_minutes, - ) - - def get_env_vars(self) -> dict[str, str]: - return env_vars( - repo_path=self.config.repo_path, env=dict(self.config.env or {}) - ) - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks( - dataset_name=self.config.dataset_name, - repo_path=self.config.repo_path, - filter_repos=self.config.filter_repos, - ds_num_proc=self.config.ds_num_proc, - ds_keep_in_memory=self.config.ds_keep_in_memory, - timeout_minutes=self.config.timeout_minutes, - env=dict(self.config.env or {}), - ) - - @vf.setup(priority=250) - async def setup_r2e_sandbox(self, task, state, sandbox=None) -> None: - if sandbox is None: - raise RuntimeError("R2E SWE setup requires the active program sandbox.") - state["_rlm_swe_sandbox"] = sandbox - state["sandbox_id"] = getattr(sandbox, "id", state.get("sandbox_id")) - sandbox_config = task.get("sandbox") - if isinstance(sandbox_config, Mapping): - timeout_minutes = int(sandbox_config.get("timeout_minutes") or 60) - state.setdefault("test_timeout", timeout_minutes * 60) - await self.setup_sandbox(sandbox, state) - - async def setup_sandbox(self, sandbox: R2ESandbox, state: vf.State) -> None: - async def exec_checked( - command: str, working_dir: str | None = None, timeout: int = 90 - ): - result = await sandbox.execute( - command, working_dir=working_dir, timeout=timeout - ) - if result.exit_code != 0: - raise RuntimeError( - f"Setup command failed: {command} exit_code={result.exit_code}" - ) - return result - - link_commands = [ - f"ln -s {self.config.repo_path}/.venv {self.config.alt_path}/.venv", - f"ln -s {self.config.repo_path}/.venv/bin/python {self.config.alt_path}/.local/bin/python", - f"ln -s {self.config.repo_path}/.venv/bin/python {self.config.alt_path}/.local/bin/python3", - f"find {self.config.repo_path}/.venv/bin -type f -executable -exec ln -sfn {{}} {self.config.alt_path}/.local/bin/ \\;", - ] - for command in link_commands: - await exec_checked(command) - - try: - cleanup_commands = [ - ( - "timeout 30 bash -c 'shopt -s globstar; rm -rf **/*.pyc **/__pycache__' 2>/dev/null || timeout 30 find . -name '*.pyc' -delete || true", - self.config.repo_path, - ), - ( - "timeout 30 bash -c 'shopt -s globstar; rm -rf **/__pycache__' 2>/dev/null || timeout 30 find . -name '__pycache__' -exec rm -rf {} + || true", - self.config.repo_path, - ), - ( - "timeout 30 bash -c 'shopt -s globstar; rm -rf /r2e_tests/**/*.pyc /r2e_tests/**/__pycache__' 2>/dev/null || timeout 30 find /r2e_tests -name '*.pyc' -delete || true", - None, - ), - ( - "timeout 30 bash -c 'shopt -s globstar; rm -rf /r2e_tests/**/__pycache__' 2>/dev/null || timeout 30 find /r2e_tests -name '__pycache__' -exec rm -rf {} + || true", - None, - ), - ] - for command, working_dir in cleanup_commands: - await exec_checked(command, working_dir=working_dir) - except Exception as exc: - logger.warning("Continuing without deleting pycache: %r", exc) - - if not self.config.hide_tests_from_agent: - await exec_checked( - f"mv /r2e_tests {self.config.repo_path}/r2e_tests", timeout=60 - ) - return - - remote_archive = "/tmp/r2e_tests.tar.gz" - local_archive_path = str(Path("/tmp") / f"r2e_tests_{sandbox.id}.tar.gz") - await exec_checked(f"tar -C / -czf {remote_archive} r2e_tests", timeout=300) - await sandbox.download_file(remote_archive, local_archive_path, timeout=300) - state["r2e_tests_archive_local_path"] = local_archive_path - await exec_checked("rm -rf /r2e_tests", timeout=300) - await exec_checked(f"rm -f {remote_archive}", timeout=300) - - @vf.reward(weight=1.0) - async def solved(self, task, state) -> float: - if state.get("error") is not None: - return 0.0 - sandbox = state.get("_rlm_swe_sandbox") - if sandbox is None: - return 0.0 - try: - test_output = await self.run_tests( - sandbox, - state, - int(state.get("test_timeout", 900)), - ) - state["test_output"] = test_output - except Exception as exc: - logger.warning("Test execution failed: %r", exc) - state["test_output"] = f"ERROR: {exc}" - return 0.0 - return float(self.calculate_reward(test_output, task["info"])) - - async def run_tests( - self, - sandbox: R2ESandbox, - state: vf.State, - test_timeout: int, - ) -> str: - local_archive_path = state.get("r2e_tests_archive_local_path") - if local_archive_path and Path(str(local_archive_path)).exists(): - remote_archive = "/tmp/r2e_tests_roundtrip.tar.gz" - await sandbox.upload_file( - remote_archive, str(local_archive_path), timeout=300 - ) - result = await sandbox.execute( - f"tar -C {self.config.repo_path} -xzf {remote_archive}", - timeout=300, - ) - if result.exit_code != 0: - raise RuntimeError( - f"Failed to extract r2e_tests: exit_code={result.exit_code}" - ) - Path(str(local_archive_path)).unlink(missing_ok=True) - del state["r2e_tests_archive_local_path"] - elif self.config.hide_tests_from_agent: - raise RuntimeError( - f"Missing cached r2e_tests archive: {local_archive_path}" - ) - - env_str = " ".join( - f"{shlex.quote(key)}={shlex.quote(value)}" - for key, value in self.get_env_vars().items() - ) - command = f"export {env_str}; /bin/bash run_tests.sh > test_output.txt 2>&1" - result = await sandbox.run_background_job( - command, timeout=test_timeout, working_dir=self.config.repo_path - ) - if result.exit_code > 1: - raise RuntimeError(f"Error running tests: exit_code={result.exit_code}") - result = await sandbox.execute( - f"cat {self.config.repo_path}/test_output.txt", timeout=300 - ) - return result.stdout or "" - - def calculate_reward(self, test_output: str, info: vf.JsonData) -> float: - parsed = parse_log_pytest(test_output) - parsed = decolor_dict_keys(parsed) - expected_raw = info["expected_output_json"] - expected: dict[str, str] = json.loads(str(expected_raw)) - expected = decolor_dict_keys(expected) - parsed = {key.split(" - ")[0]: parsed[key] for key in sorted(parsed.keys())} - expected = { - key.split(" - ")[0]: expected[key] for key in sorted(expected.keys()) - } - if any(not key for key in parsed): - return 0.0 - if len(parsed) != len(expected): - return 0.0 - for key in parsed: - if key not in expected or parsed[key] != expected[key]: - return 0.0 - return 1.0 - - async def apply_gold_patch(self, sandbox: R2ESandbox, state: vf.State) -> None: - info = cast(vf.JsonData, state["info"]) - assert isinstance(info, Mapping) - patch = extract_gold_patch( - str(info["parsed_commit_content"]), - test_file=False, - only_python=True, - ) - if not patch.strip(): - raise RuntimeError( - "Empty gold patch reconstructed from parsed_commit_content" - ) - - await sandbox.upload_bytes("/tmp/gold.patch", patch.encode(), "gold.patch") - result = await sandbox.execute( - "git apply --whitespace=fix /tmp/gold.patch", - working_dir=self.config.repo_path, - timeout=30, - ) - if result.exit_code != 0: - stderr = (result.stderr or "")[:500] - raise RuntimeError( - f"git apply failed: exit_code={result.exit_code} stderr={stderr}" - ) - - async def validate_instance(self, state: vf.State) -> bool: - sandbox = cast(R2ESandbox, state["_rlm_swe_sandbox"]) - await self.apply_gold_patch(sandbox, state) - test_timeout = state.get("test_timeout", 900) - assert isinstance(test_timeout, int) - test_output = await self.run_tests( - sandbox, - state, - test_timeout, - ) - state["test_output"] = test_output - info = cast(vf.JsonData, state["info"]) - assert isinstance(info, Mapping) - return self.calculate_reward(test_output, info) > 0 - - @vf.cleanup(priority=100) - async def cleanup_r2e_state(self, task, state) -> None: - archive = state.pop("r2e_tests_archive_local_path", None) - if isinstance(archive, str): - Path(archive).unlink(missing_ok=True) - state.pop("_rlm_swe_sandbox", None) - - -def process_r2e_example(row: vf.JsonData) -> vf.JsonData: - info = cast(vf.JsonData, dict(row)) - info.setdefault("instance_id", row.get("commit_hash")) - info.setdefault("repo", row.get("repo_name")) - return { - "question": row["problem_statement"], - "info": info, - "answer": "", - } - - -def parse_log_pytest(log: str | None) -> dict[str, str]: - if log is None or "short test summary info" not in log: - return {} - test_status_map = {} - for line in log.split("short test summary info", 1)[1].strip().splitlines(): - status, _, test_ref = line.strip().partition(" ") - if status in {"PASSED", "FAILED", "ERROR"}: - test_name = ( - ".".join(test_ref.split("::")[1:]) if "::" in test_ref else test_ref - ) - test_status_map[test_name.split(" - ")[0]] = status - return test_status_map - - -def decolor_dict_keys(values: dict[str, str]) -> dict[str, str]: - return {re.sub(r"\u001b\[\d+m", "", key): value for key, value in values.items()} - - -def extract_gold_patch( - parsed_commit_content: str, test_file: bool = False, only_python: bool = True -) -> str: - data = json.loads(parsed_commit_content) - patch = "" - for file_diff in data.get("file_diffs", []): - path = file_diff.get("header", {}).get("file", {}).get("path", "") - if not path: - continue - if only_python and not path.endswith(".py"): - continue - parts = path.split("/") - is_test = ( - path.endswith("_test.py") - or path.split("/")[-1].startswith("test_") - or any(part in {"tests", "Tests", "test", "Test"} for part in parts) - ) - if is_test and not test_file: - continue - - header = file_diff.get("header", {}) - misc_line = header.get("misc_line") - patch += f"diff --git a/{path} b/{path}\n" - if misc_line: - patch += misc_line + "\n" - - index_line = file_diff.get("index_line") - if index_line: - old_hash = index_line.get("old_commit_hash", "") - new_hash = index_line.get("new_commit_hash", "") - mode = index_line.get("mode", "") - patch += f"index {old_hash}..{new_hash}{' ' if mode else ''}{mode}\n" - - if file_diff.get("is_binary_file"): - binary_line = file_diff.get("binary_line", "") - if binary_line: - patch += binary_line + "\n" - - minus_file = file_diff.get("minus_file") - plus_file = file_diff.get("plus_file") - if minus_file and plus_file: - patch += f"--- {minus_file['path']}\n" - patch += f"+++ {plus_file['path']}\n" - - for hunk in file_diff.get("hunks", []): - descriptor = hunk.get("descriptor", {}) - old_range = descriptor.get("old_range", {}) - new_range = descriptor.get("new_range", {}) - old_str = str(old_range.get("start", 0)) - if old_range.get("length") is not None: - old_str += f",{old_range['length']}" - new_str = str(new_range.get("start", 0)) - if new_range.get("length") is not None: - new_str += f",{new_range['length']}" - section = descriptor.get("section", "") - hunk_header = f"@@ -{old_str} +{new_str} @@" - if section: - hunk_header += f" {section}" - patch += hunk_header + "\n" - - for line in hunk.get("line_group", {}).get("all_lines", []): - content = line.get("content", "") - line_type = line.get("type", "") - if line_type == "context": - patch += f" {content}\n" - elif line_type == "added": - patch += f"+{content}\n" - elif line_type == "deleted": - patch += f"-{content}\n" - elif line_type == "note": - patch += f"\\ {content}\n" - return patch - - -def load_taskset( - config: RlmSweTasksetConfig, -) -> R2ESWETaskset: - return R2ESWETaskset(config=config) - - -class RlmSweProgramConfig(RLMProgramConfig): - workdir: str = DEFAULT_REPO_PATH - tools: list[str] = list(DEFAULT_RLM_TOOLS) - - -class RlmSweHarnessConfig(RLMConfig): - program: RlmSweProgramConfig = RlmSweProgramConfig() - - -def load_harness(config: RlmSweHarnessConfig) -> RLM: - return RLM(config=config) - - -class RlmSweEnvConfig(vf.EnvConfig): - taskset: RlmSweTasksetConfig = RlmSweTasksetConfig() - harness: RlmSweHarnessConfig = RlmSweHarnessConfig() - - -def load_environment(config: RlmSweEnvConfig) -> vf.Env: - taskset = load_taskset(config=config.taskset) - harness = load_harness(config=config.harness) - return vf.Env(taskset=taskset, harness=harness) diff --git a/environments/rlm_swe_v1/rlm_swe_v1/__init__.py b/environments/rlm_swe_v1/rlm_swe_v1/__init__.py new file mode 100644 index 0000000000..326efc591a --- /dev/null +++ b/environments/rlm_swe_v1/rlm_swe_v1/__init__.py @@ -0,0 +1 @@ +"""rlm-swe-v1 environment package.""" diff --git a/environments/rlm_swe_v1/rlm_swe_v1/harness.py b/environments/rlm_swe_v1/rlm_swe_v1/harness.py new file mode 100644 index 0000000000..5b41a3eab8 --- /dev/null +++ b/environments/rlm_swe_v1/rlm_swe_v1/harness.py @@ -0,0 +1,3 @@ +from .taskset import RLM as RLM +from .taskset import RlmSweHarnessConfig as RlmSweHarnessConfig +from .taskset import load_harness as load_harness diff --git a/environments/rlm_swe_v1/rlm_swe_v1/taskset.py b/environments/rlm_swe_v1/rlm_swe_v1/taskset.py new file mode 100644 index 0000000000..eb4313801e --- /dev/null +++ b/environments/rlm_swe_v1/rlm_swe_v1/taskset.py @@ -0,0 +1,250 @@ +import json +import re +from typing import cast + +from datasets import load_dataset +from harnesses import RLM, RLMConfig + +import verifiers.v1 as vf + +REGISTRY_PREFIX = "us-central1-docker.pkg.dev/prime-intellect-platform/prod-sandbox" +DEFAULT_DATASET_NAME = "R2E-Gym/R2E-Gym-Subset" +DEFAULT_REPO_PATH = "/testbed" +DEFAULT_RLM_TOOLS = ("bash", "edit") + + +class RlmSweTasksetConfig(vf.TasksetConfig): + id: str = "swe/r2e" + dataset_name: str = DEFAULT_DATASET_NAME + repo_path: str = DEFAULT_REPO_PATH + filter_repos: list[str] | None = None + ds_num_proc: int | None = None + ds_keep_in_memory: bool = True + timeout_minutes: int | None = None + env: vf.JsonData | None = None + + +class RlmSweHarnessConfig(RLMConfig): + cwd: str | None = DEFAULT_REPO_PATH + tools: list[str] = list(DEFAULT_RLM_TOOLS) + + +class R2ESWETask(vf.Task): + question: str + instruction: str + answer: str + info: vf.JsonData + sandbox: vf.JsonData + program: vf.JsonData + + +class R2ESWETaskset(vf.Taskset[RlmSweTasksetConfig]): + task_type = R2ESWETask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return load_tasks( + dataset_name=self.config.dataset_name, + repo_path=self.config.repo_path, + filter_repos=self.config.filter_repos, + ds_num_proc=self.config.ds_num_proc, + ds_keep_in_memory=self.config.ds_keep_in_memory, + timeout_minutes=self.config.timeout_minutes, + env=dict(self.config.env or {}), + ) + + @vf.reward(weight=1.0) + async def solved(self, task: R2ESWETask, state: vf.State) -> float: + test_output = state.artifacts.get("test_output") + if not isinstance(test_output, str): + command = state.artifacts.get("command") + if isinstance(command, dict): + test_output = str(command.get("stdout") or command.get("stderr") or "") + return float(calculate_reward(test_output or "", task.info)) + + +def load_tasks( + dataset_name: str = DEFAULT_DATASET_NAME, + repo_path: str = DEFAULT_REPO_PATH, + filter_repos: list[str] | None = None, + ds_num_proc: int | None = None, + ds_keep_in_memory: bool = True, + timeout_minutes: int | None = None, + env: dict[str, str] | None = None, +) -> list[vf.JsonData]: + dataset_kwargs = dict( + num_proc=ds_num_proc, + keep_in_memory=ds_keep_in_memory, + load_from_cache_file=False, + ) + dataset = load_dataset( + dataset_name, + split="train", + keep_in_memory=ds_keep_in_memory, + num_proc=ds_num_proc, + ) + if filter_repos: + filter_set = frozenset(filter_repos) + dataset = dataset.filter( + lambda row: row.get("repo_name") not in filter_set, + **dataset_kwargs, + ) + dataset = dataset.map( + process_r2e_example, + remove_columns=dataset.column_names, + **dataset_kwargs, + ) + task_env = dict(env or {}) + rows: list[vf.JsonData] = [] + for index, row in enumerate(dataset): + row = dict(row) + info = dict(row["info"]) + instruction = str(info["problem_statement"]) + program_env = env_vars(repo_path=repo_path, env=task_env) + program_env.setdefault("AGENT_WORKDIR", repo_path) + rows.append( + { + "row_id": index, + "task_id": info.get("instance_id") or index, + "question": row.get("question", instruction), + "instruction": instruction, + "prompt": [{"role": "user", "content": instruction}], + "answer": row.get("answer", ""), + "info": info, + "sandbox": sandbox_config( + info=cast(vf.JsonData, info), + repo_path=repo_path, + timeout_minutes=timeout_minutes, + ), + "program": {"env": program_env}, + } + ) + return rows + + +def sandbox_config( + *, info: vf.JsonData, repo_path: str, timeout_minutes: int | None +) -> vf.JsonData: + config: vf.JsonData = { + "image": f"{REGISTRY_PREFIX}/{info['docker_image']}", + "cpu_cores": 4, + "memory_gb": 4, + "disk_size_gb": 10, + "workdir": repo_path, + } + if timeout_minutes is not None: + config["timeout_minutes"] = timeout_minutes + return config + + +def env_vars(*, repo_path: str, env: dict[str, str]) -> dict[str, str]: + return { + "PATH": ( + f"/opt/miniconda3/bin:{repo_path}/.venv/bin:/root/.local/bin:" + "/root/.cargo/bin:/go/bin:/usr/local/go/bin:/usr/local/cargo:" + "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin" + ), + "PAGER": "cat", + "MANPAGER": "cat", + "LESS": "-R", + "PIP_PROGRESS_BAR": "off", + "TQDM_DISABLE": "1", + **env, + } + + +def process_r2e_example(row: vf.JsonData) -> vf.JsonData: + info = cast(vf.JsonData, dict(row)) + info.setdefault("instance_id", row.get("commit_hash")) + info.setdefault("repo", row.get("repo_name")) + return { + "question": row["problem_statement"], + "info": info, + "answer": "", + } + + +def calculate_reward(test_output: str, info: vf.JsonData) -> float: + parsed = decolor_dict_keys(parse_log_pytest(test_output)) + expected_raw = info["expected_output_json"] + expected: dict[str, str] = json.loads(str(expected_raw)) + expected = decolor_dict_keys(expected) + parsed = {key.split(" - ")[0]: parsed[key] for key in sorted(parsed.keys())} + expected = {key.split(" - ")[0]: expected[key] for key in sorted(expected.keys())} + if any(not key for key in parsed): + return 0.0 + if len(parsed) != len(expected): + return 0.0 + for key in parsed: + if key not in expected or parsed[key] != expected[key]: + return 0.0 + return 1.0 + + +def parse_log_pytest(log: str | None) -> dict[str, str]: + if log is None or "short test summary info" not in log: + return {} + test_status_map = {} + for line in log.split("short test summary info", 1)[1].strip().splitlines(): + status, _, test_ref = line.strip().partition(" ") + if status in {"PASSED", "FAILED", "ERROR"}: + test_name = ( + ".".join(test_ref.split("::")[1:]) if "::" in test_ref else test_ref + ) + test_status_map[test_name.split(" - ")[0]] = status + return test_status_map + + +def decolor_dict_keys(values: dict[str, str]) -> dict[str, str]: + return {re.sub(r"\u001b\[\d+m", "", key): value for key, value in values.items()} + + +def extract_gold_patch( + parsed_commit_content: str, + *, + test_file: bool = False, + only_python: bool = True, +) -> str: + data = json.loads(parsed_commit_content) + patch = "" + for file_diff in data.get("file_diffs", []): + path = file_diff.get("header", {}).get("file", {}).get("path", "") + if not path: + continue + if only_python and not path.endswith(".py"): + continue + parts = path.split("/") + is_test = ( + path.endswith("_test.py") + or path.split("/")[-1].startswith("test_") + or any(part in {"tests", "Tests", "test", "Test"} for part in parts) + ) + if is_test and not test_file: + continue + patch += f"diff --git a/{path} b/{path}\n" + for hunk in file_diff.get("hunks", []): + descriptor = hunk.get("descriptor", {}) + old_range = descriptor.get("old_range", {}) + new_range = descriptor.get("new_range", {}) + patch += ( + f"@@ -{old_range.get('start', 0)} +{new_range.get('start', 0)} @@\n" + ) + for line in hunk.get("line_group", {}).get("all_lines", []): + content = line.get("content", "") + line_type = line.get("type", "") + if line_type == "context": + patch += f" {content}\n" + elif line_type == "added": + patch += f"+{content}\n" + elif line_type == "deleted": + patch += f"-{content}\n" + return patch + + +def load_taskset(config: RlmSweTasksetConfig) -> R2ESWETaskset: + return R2ESWETaskset(config=config) + + +def load_harness(config: RlmSweHarnessConfig) -> RLM: + return RLM(config=config) diff --git a/environments/sft_replay/README.md b/environments/sft_replay_v1/README.md similarity index 83% rename from environments/sft_replay/README.md rename to environments/sft_replay_v1/README.md index a1ba9dfe89..12158c2787 100644 --- a/environments/sft_replay/README.md +++ b/environments/sft_replay_v1/README.md @@ -1,8 +1,8 @@ -# sft-replay +# sft-replay-v1 ### Overview -- **Environment ID**: `sft-replay` -- **Short description**: Replay stored chat transcripts into trajectory steps without making model requests. +- **Environment ID**: `sft-replay-v1` +- **Short description**: Replay stored chat transcripts into v1 transcript turns without making model requests. - **Tags**: replay, sft, v1 ### Datasets @@ -23,19 +23,19 @@ canonical serialized message format. ### Task - **Type**: replay - **Output format expectations (optional)**: Stored assistant messages are replayed exactly in their canonical message shape. -- **Rubric overview**: No default reward; this environment produces replay trajectories for downstream SFT-style processing. +- **Rubric overview**: No default reward; this environment produces replay transcripts for downstream SFT-style processing. ### Quickstart Run an evaluation with default settings: ```bash -prime eval run sft-replay +prime eval run sft-replay-v1 ``` Configure model and sampling: ```bash -prime eval run sft-replay \ +prime eval run sft-replay-v1 \ -m openai/gpt-4.1-mini \ -n 20 -r 3 -t 1024 -T 0.7 ``` @@ -49,7 +49,7 @@ Notes: | Field | Type | Default | Description | | --- | ---- | ------- | ----------- | | `dataset` | str \| null | `null` | Hugging Face dataset ID to load instead of env-local `data/*.jsonl` files. | -| `data_dir` | str \| null | `null` | Local JSONL directory. When unset, `sft-replay` uses its packaged `data/` directory. | +| `data_dir` | str \| null | `null` | Local JSONL directory. When unset, `sft-replay-v1` uses its packaged `data/` directory. | ### Harness Config Uses `vf.HarnessConfig`. By default, every assistant message is replayed. @@ -59,4 +59,4 @@ Set `max_turns` to cap the number of assistant messages replayed per rollout. | Metric | Meaning | | ------ | ------- | -| `num_model_requests` | Number of assistant messages replayed into trajectory steps. | +| `num_model_requests` | Number of assistant messages replayed into transcript turns. | diff --git a/environments/sft_replay/pyproject.toml b/environments/sft_replay_v1/pyproject.toml similarity index 83% rename from environments/sft_replay/pyproject.toml rename to environments/sft_replay_v1/pyproject.toml index fc402d9915..4435b37bc4 100644 --- a/environments/sft_replay/pyproject.toml +++ b/environments/sft_replay_v1/pyproject.toml @@ -1,5 +1,5 @@ [project] -name = "sft-replay" +name = "sft-replay-v1" version = "0.1.0" description = "Thin v1 environment for replaying stored chat transcripts as SFT trajectories." tags = ["sft", "replay", "v1", "taskset", "harness", "train"] @@ -15,7 +15,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["sft_replay.py", "README.md", "pyproject.toml", "data/**/*"] +include = ["sft_replay_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/sft_replay_v1/sft_replay_v1/__init__.py b/environments/sft_replay_v1/sft_replay_v1/__init__.py new file mode 100644 index 0000000000..b75a7b1df4 --- /dev/null +++ b/environments/sft_replay_v1/sft_replay_v1/__init__.py @@ -0,0 +1 @@ +"""sft-replay-v1 environment package.""" diff --git a/environments/sft_replay/data/reverse.jsonl b/environments/sft_replay_v1/sft_replay_v1/data/reverse.jsonl similarity index 100% rename from environments/sft_replay/data/reverse.jsonl rename to environments/sft_replay_v1/sft_replay_v1/data/reverse.jsonl diff --git a/environments/sft_replay_v1/sft_replay_v1/harness.py b/environments/sft_replay_v1/sft_replay_v1/harness.py new file mode 100644 index 0000000000..737318ecb0 --- /dev/null +++ b/environments/sft_replay_v1/sft_replay_v1/harness.py @@ -0,0 +1,2 @@ +from .taskset import ReplayHarness as ReplayHarness +from .taskset import load_harness as load_harness diff --git a/environments/sft_replay/sft_replay.py b/environments/sft_replay_v1/sft_replay_v1/taskset.py similarity index 66% rename from environments/sft_replay/sft_replay.py rename to environments/sft_replay_v1/sft_replay_v1/taskset.py index 699e526a57..df86be5ef7 100644 --- a/environments/sft_replay/sft_replay.py +++ b/environments/sft_replay_v1/sft_replay_v1/taskset.py @@ -1,6 +1,6 @@ from pathlib import Path -import verifiers as vf +import verifiers.v1 as vf from harnesses import ReplayHarness from tasksets import ReplayTaskset, ReplayTasksetConfig @@ -15,10 +15,3 @@ def load_taskset(config: ReplayTasksetConfig) -> SFTReplayTaskset: def load_harness(config: vf.HarnessConfig) -> ReplayHarness: return ReplayHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/environments/tau2_bench_v1/pyproject.toml b/environments/tau2_bench_v1/pyproject.toml index 9a7aa5a3da..d3e4a77a63 100644 --- a/environments/tau2_bench_v1/pyproject.toml +++ b/environments/tau2_bench_v1/pyproject.toml @@ -16,7 +16,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["tau2_bench_v1.py", "pyproject.toml"] +include = ["tau2_bench_v1/**/*", "README.md", "pyproject.toml"] [tool.hatch.metadata] allow-direct-references = true diff --git a/environments/tau2_bench_v1/tau2_bench_v1.py b/environments/tau2_bench_v1/tau2_bench_v1.py deleted file mode 100644 index 987994f5e6..0000000000 --- a/environments/tau2_bench_v1/tau2_bench_v1.py +++ /dev/null @@ -1,687 +0,0 @@ -import asyncio -import json -import os -import shutil -import subprocess -import uuid -from copy import deepcopy -from datetime import datetime, timedelta -from pathlib import Path -from typing import cast - -import verifiers as core_vf -import verifiers as vf -from verifiers.types import Tool - -from tau2.agent.llm_agent import AGENT_INSTRUCTION, SYSTEM_PROMPT, LLMAgent -from tau2.agent.llm_agent import is_valid_agent_history_message -from tau2.config import DEFAULT_LLM_ARGS_AGENT, DEFAULT_LLM_ARGS_USER -from tau2.config import DEFAULT_MAX_ERRORS, DEFAULT_MAX_STEPS -from tau2.data_model.message import AssistantMessage, Message, MultiToolMessage -from tau2.data_model.message import ToolCall, ToolMessage, UserMessage -from tau2.data_model.simulation import SimulationRun, TerminationReason -from tau2.data_model.tasks import Task as TauTask -from tau2.environment.environment import Environment as TauEnvironment -from tau2.evaluator.evaluator import EvaluationType, evaluate_simulation -from tau2.orchestrator.orchestrator import DEFAULT_FIRST_AGENT_MESSAGE, Role -from tau2.registry import registry -from tau2.run import load_tasks as load_tau2_tasks -from tau2.user.user_simulator import UserSimulator, is_valid_user_history_message -from tau2.utils.utils import DATA_DIR, format_time, get_now - -DEFAULT_USER_MODEL = "openai/gpt-4.1-mini" -DEFAULT_USER_BASE_URL = "https://api.pinference.ai/api/v1" -DEFAULT_USER_API_KEY_VAR = "PRIME_API_KEY" - - -def download_tau2_data() -> None: - if os.path.exists(DATA_DIR) and os.path.exists(DATA_DIR / "tau2" / "domains"): - return - os.makedirs(DATA_DIR, exist_ok=True) - temp_dir = Path("/tmp/tau2_bench_v1") - try: - subprocess.run( - [ - "git", - "clone", - "--depth", - "1", - "https://github.com/sierra-research/tau2-bench.git", - temp_dir, - ], - check=True, - capture_output=True, - ) - source_data = temp_dir / "data" - if source_data.exists(): - shutil.copytree(source_data, DATA_DIR, dirs_exist_ok=True) - finally: - if temp_dir.exists(): - shutil.rmtree(temp_dir) - - -def tau_msg_to_vf_dict(message: Message) -> vf.ConfigData: - if isinstance(message, AssistantMessage): - if message.tool_calls: - return core_vf.AssistantMessage( - content=message.content, - tool_calls=[ - core_vf.ToolCall( - id=tool_call.id, - name=tool_call.name, - arguments=json.dumps(tool_call.arguments), - ) - for tool_call in message.tool_calls - ], - ).model_dump(exclude_none=True) - return core_vf.AssistantMessage(content=message.content).model_dump( - exclude_none=True - ) - if isinstance(message, UserMessage): - return core_vf.UserMessage(content=message.content or "").model_dump( - exclude_none=True - ) - if isinstance(message, ToolMessage): - return core_vf.ToolMessage( - tool_call_id=message.id, - content=message.content or "", - ).model_dump(exclude_none=True) - raise ValueError(f"Unknown tau2 message type: {type(message)}") - - -def dump_tau_message(message: Message) -> vf.ConfigData: - return cast(vf.ConfigData, message.model_dump(mode="json", exclude_none=True)) - - -def load_tau_message(payload: vf.JsonData) -> Message: - role = payload.get("role") - if role == "assistant": - return AssistantMessage.model_validate(payload) - if role == "user": - return UserMessage.model_validate(payload) - if role == "tool": - if "tool_messages" in payload: - return MultiToolMessage.model_validate(payload) - return ToolMessage.model_validate(payload) - raise ValueError(f"Unknown tau2 message role: {role!r}") - - -class Tau2Session: - def __init__( - self, - domain: str, - task_payload: vf.JsonData, - user_model: str, - user_args: vf.JsonData, - user_base_url: str, - user_api_key_var: str, - max_steps: int, - max_errors: int, - ): - self.domain = domain - self.task_payload = dict(task_payload) - self.user_model = user_model - self.user_args = dict(user_args) - self.user_base_url = user_base_url - self.user_api_key_var = user_api_key_var - self.max_steps = max_steps - self.max_errors = max_errors - self.ready = False - self.initial_prompt_messages: list[vf.ConfigData] = [] - self.recorded_assistant_messages = 0 - self.pending_agent_tool_calls: list[ToolCall] = [] - self.num_assistant_tool_calls = 0 - self.num_user_tool_calls = 0 - - async def initialize(self, state: vf.State) -> None: - if self.ready: - return - self.task = TauTask.model_validate(self.task_payload) - environment_constructor = registry.get_env_constructor(self.domain) - self.environment = await asyncio.to_thread(environment_constructor) - self.agent = LLMAgent( - tools=self.environment.get_tools(), - domain_policy=self.environment.get_policy(), - llm=str(state.get("runtime", {}).get("model") or ""), - llm_args=dict(state.get("runtime", {}).get("sampling_args") or {}), - ) - user_args = dict(self.user_args) - if self.user_base_url == DEFAULT_USER_BASE_URL: - custom_provider = user_args.get("custom_llm_provider") - assert custom_provider in (None, "custom_openai") - user_args["custom_llm_provider"] = "custom_openai" - if self.user_api_key_var == "PRIME_API_KEY": - team_id = os.getenv("PRIME_TEAM_ID") - if team_id: - extra_headers = user_args.get("extra_headers") or {} - assert isinstance(extra_headers, dict) - user_args["extra_headers"] = { - **extra_headers, - "X-Prime-Team-ID": team_id, - } - user_args["api_base"] = self.user_base_url - user_args["api_key"] = os.getenv(self.user_api_key_var) - self.user = UserSimulator( - tools=self.user_tools(), - instructions=str(self.task.user_scenario), - llm=self.user_model, - llm_args=user_args, - ) - self.init_tau2_state() - self.environment.sync_tools() - self.ready = True - self.initial_prompt_messages = await self.advance_until_agent(state) - self.render_state(state) - - def user_tools(self) -> object: - try: - return self.environment.get_user_tools() - except Exception: - return None - - def init_tau2_state(self) -> None: - initial_state = self.task.initial_state - initialization_data = ( - initial_state.initialization_data if initial_state is not None else None - ) - initialization_actions = ( - initial_state.initialization_actions if initial_state is not None else None - ) - message_history = ( - deepcopy(initial_state.message_history) - if initial_state is not None and initial_state.message_history is not None - else [] - ) - for message in message_history: - message.turn_idx = None - message_history = add_timestamps(message_history) - self.environment.set_state( - initialization_data=initialization_data, - initialization_actions=initialization_actions, - message_history=message_history, - ) - self.done = False - self.termination_reason: TerminationReason | None = None - if message_history: - self.init_from_history(message_history) - else: - self.user_state = self.user.get_init_state() - first_message = deepcopy(DEFAULT_FIRST_AGENT_MESSAGE) - first_message.timestamp = get_now() - self.agent_state = self.agent.get_init_state( - message_history=[first_message] - ) - self.trajectory: list[Message] = [first_message] - self.message = first_message - self.from_role = Role.AGENT - self.to_role = Role.USER - self.step_count = 0 - self.num_errors = 0 - - def init_from_history(self, message_history: list[Message]) -> None: - last_message = message_history[-1] - self.trajectory = message_history - self.message = last_message - if isinstance(last_message, AssistantMessage): - self.from_role = Role.AGENT - self.to_role = Role.ENV if last_message.is_tool_call() else Role.USER - self.agent_state = self.agent.get_init_state( - message_history=[ - msg - for msg in message_history - if is_valid_agent_history_message(msg) - ] - ) - self.user_state = self.user.get_init_state( - message_history=[ - msg - for msg in message_history[:-1] - if is_valid_user_history_message(msg) - ] - ) - if self.agent.is_stop(last_message): - self.done = True - self.termination_reason = TerminationReason.AGENT_STOP - return - if isinstance(last_message, UserMessage): - self.from_role = Role.USER - self.to_role = Role.ENV if last_message.is_tool_call() else Role.AGENT - self.user_state = self.user.get_init_state( - message_history=[ - msg for msg in message_history if is_valid_user_history_message(msg) - ] - ) - self.agent_state = self.agent.get_init_state( - message_history=[ - msg - for msg in message_history[:-1] - if is_valid_agent_history_message(msg) - ] - ) - self.done = UserSimulator.is_stop(last_message) - if self.done: - self.termination_reason = TerminationReason.USER_STOP - return - if isinstance(last_message, ToolMessage): - self.from_role = Role.ENV - self.to_role = ( - Role.AGENT if last_message.requestor == "assistant" else Role.USER - ) - self.agent_state = self.agent.get_init_state( - message_history=[ - msg - for msg in message_history - if is_valid_agent_history_message(msg) - ] - ) - self.user_state = self.user.get_init_state( - message_history=[ - msg for msg in message_history if is_valid_user_history_message(msg) - ] - ) - return - raise ValueError(f"Unsupported tau2 message type: {type(last_message)}") - - async def record_assistant_from_state(self, state: vf.State) -> None: - completion = state.get("completion") or [] - assistant_messages = ( - vf.get_messages(completion, role="assistant") - if isinstance(completion, list) - else [] - ) - for message in assistant_messages[self.recorded_assistant_messages :]: - assistant_message = assistant_from_openai_message(message) - self.agent_state.messages.append(assistant_message) - try: - assistant_message.validate() - except ValueError: - self.done = True - self.termination_reason = TerminationReason.AGENT_ERROR - self.trajectory.append(assistant_message) - continue - if self.agent.is_stop(assistant_message): - self.done = True - self.termination_reason = TerminationReason.AGENT_STOP - self.trajectory.append(assistant_message) - self.pending_agent_tool_calls.extend(assistant_message.tool_calls or []) - self.num_assistant_tool_calls += len(assistant_message.tool_calls or []) - self.message = assistant_message - self.from_role = Role.AGENT - self.to_role = Role.ENV if assistant_message.tool_calls else Role.USER - self.step_count += 1 - self.environment.sync_tools() - self.check_limits() - self.recorded_assistant_messages = len(assistant_messages) - self.render_state(state) - - async def call_agent_tool( - self, name: str, arguments: vf.JsonData, state: vf.State - ) -> str: - await self.record_assistant_from_state(state) - tool_call = self.pop_pending_tool_call(name, arguments) - tool_message = await asyncio.to_thread(self.environment.get_response, tool_call) - if tool_message.error: - self.num_errors += 1 - self.trajectory.append(tool_message) - self.message = tool_message - self.from_role = Role.ENV - self.to_role = Role.AGENT - self.environment.sync_tools() - self.check_limits() - self.render_state(state) - return tool_message.content or "" - - def pop_pending_tool_call(self, name: str, arguments: vf.JsonData) -> ToolCall: - for index, tool_call in enumerate(self.pending_agent_tool_calls): - if tool_call.name == name: - return self.pending_agent_tool_calls.pop(index) - return ToolCall( - id=f"call_{uuid.uuid4().hex[:8]}", - name=name, - arguments=dict(arguments), - requestor="assistant", - ) - - async def user_messages(self, state: vf.State) -> list[vf.ConfigData]: - await self.record_assistant_from_state(state) - if self.done: - self.render_state(state) - return [] - messages = await self.advance_until_agent(state) - self.render_state(state) - return messages - - async def advance_until_agent(self, state: vf.State) -> list[vf.ConfigData]: - messages: list[vf.ConfigData] = [] - while not (self.done or self.to_role == Role.AGENT): - if self.to_role == Role.USER: - user_message, self.user_state = await asyncio.to_thread( - self.user.generate_next_message, - self.message, - self.user_state, - ) - try: - user_message.validate() - except ValueError: - self.done = True - self.termination_reason = TerminationReason.USER_ERROR - self.trajectory.append(user_message) - break - if UserSimulator.is_stop(user_message): - self.done = True - self.termination_reason = TerminationReason.USER_STOP - self.num_user_tool_calls += len(user_message.tool_calls or []) - self.trajectory.append(user_message) - self.message = user_message - self.from_role = Role.USER - self.to_role = Role.ENV if user_message.is_tool_call() else Role.AGENT - self.step_count += 1 - self.check_limits() - if not user_message.is_tool_call(): - messages.append(tau_msg_to_vf_dict(user_message)) - break - continue - if self.to_role == Role.ENV: - await self.execute_user_tools() - self.check_limits() - continue - raise ValueError( - f"Invalid tau2 role transition: {self.from_role} -> {self.to_role}" - ) - return messages - - async def execute_user_tools(self) -> None: - tool_calls = list(getattr(self.message, "tool_calls", []) or []) - tool_messages = [] - for tool_call in tool_calls: - tool_call.requestor = "user" - tool_message = await asyncio.to_thread( - self.environment.get_response, tool_call - ) - if tool_message.error: - self.num_errors += 1 - tool_messages.append(tool_message) - self.trajectory.extend(tool_messages) - if len(tool_messages) > 1: - self.message = MultiToolMessage( - role="tool", - tool_messages=tool_messages, - ) - elif tool_messages: - self.message = tool_messages[0] - self.from_role = Role.ENV - self.to_role = Role.USER - self.step_count += 1 - self.environment.sync_tools() - - def check_limits(self) -> None: - if self.step_count >= self.max_steps and self.to_role != Role.ENV: - self.done = True - self.termination_reason = TerminationReason.MAX_STEPS - if self.num_errors >= self.max_errors: - self.done = True - self.termination_reason = TerminationReason.TOO_MANY_ERRORS - - def render_state(self, state: vf.State) -> None: - state["tau2"] = { - "task_id": self.task.id, - "done": self.done, - "termination_reason": ( - self.termination_reason.value if self.termination_reason else None - ), - "step_count": self.step_count, - "num_errors": self.num_errors, - "num_assistant_tool_calls": self.num_assistant_tool_calls, - "num_user_tool_calls": self.num_user_tool_calls, - "messages": [dump_tau_message(message) for message in self.trajectory], - } - state["num_assistant_tool_calls"] = self.num_assistant_tool_calls - state["num_user_tool_calls"] = self.num_user_tool_calls - if self.done and self.termination_reason is not None: - state.stop(self.termination_reason.value) - - -def assistant_from_openai_message(message: vf.AssistantMessage) -> AssistantMessage: - tool_calls = [] - for raw_tool_call in message.tool_calls or []: - arguments = raw_tool_call.arguments - if isinstance(arguments, str): - parsed_arguments = json.loads(arguments or "{}") - else: - parsed_arguments = arguments - tool_calls.append( - ToolCall( - id=raw_tool_call.id or f"call_{uuid.uuid4().hex[:8]}", - name=raw_tool_call.name, - arguments=cast(vf.ConfigData, parsed_arguments), - requestor="assistant", - ) - ) - content = message.content - return AssistantMessage( - role="assistant", - content=content if isinstance(content, str) and content else None, - tool_calls=tool_calls or None, - raw_data=message.model_dump(exclude_none=True), - ) - - -def add_timestamps(message_history: list[Message]) -> list[Message]: - time_offset = datetime.now() - timedelta(seconds=len(message_history)) - for index, message in enumerate(message_history): - message.timestamp = format_time(time_offset + timedelta(seconds=index)) - return message_history - - -def load_tasks(domain: str, max_turns: int): - download_tau2_data() - environment_constructor = registry.get_env_constructor(domain) - environment = environment_constructor() - system_prompt = SYSTEM_PROMPT.format( - agent_instruction=AGENT_INSTRUCTION, - domain_policy=environment.policy, - ) - for index, task in enumerate( - load_tau2_tasks(task_set_name=domain, task_split_name="base") - ): - yield { - "example_id": index, - "taskset_id": f"tau2_{domain}", - "task_id": task.id, - "domain": domain, - "system_prompt": system_prompt, - "max_turns": max_turns, - "prompt": [], - "info": task.model_dump_json(exclude_none=True), - } - - -def make_tau2_tool(name: str, schema: vf.JsonData) -> vf.Handler: - async def tool(task, state, **arguments) -> str: - _ = task - session = cast( - Tau2Session, - state["tau2_session"], - ) - return await session.call_agent_tool(name, arguments, state) - - function_schema = cast(vf.JsonData, schema["function"]) - tool.__name__ = name - tool.__doc__ = str(function_schema.get("description") or "") - tool.tool_def = Tool( - name=name, - description=str(function_schema.get("description") or ""), - parameters=cast(vf.ConfigData, function_schema.get("parameters") or {}), - strict=False, - ) - return tool - - -def load_toolset( - domain: str = "telecom", -) -> vf.Toolset: - download_tau2_data() - environment_constructor = registry.get_env_constructor(domain) - environment = cast(TauEnvironment, environment_constructor()) - schemas = [tool.openai_schema for tool in environment.get_tools()] - tools = [ - make_tau2_tool( - str(cast(vf.JsonData, schema["function"])["name"]), - schema, - ) - for schema in schemas - ] - return vf.Toolset( - tools=tools, - write=True, - scope="rollout", - ) - - -class Tau2UserConfig(vf.UserConfig): - pass - - -class Tau2User(vf.User[Tau2UserConfig]): - async def get_response(self, task, state) -> list[vf.ConfigData]: - _ = task - session = cast( - Tau2Session, - state["tau2_session"], - ) - return await session.user_messages(state) - - -class Tau2TasksetConfig(vf.TasksetConfig): - taskset_id: str | None = "tau2_telecom" - user: Tau2UserConfig | None = Tau2UserConfig() - domain: str = "telecom" - user_model: str = DEFAULT_USER_MODEL - user_args: vf.ConfigData | None = None - user_base_url: str = DEFAULT_USER_BASE_URL - user_api_key_var: str = DEFAULT_USER_API_KEY_VAR - max_steps: int = DEFAULT_MAX_STEPS - max_errors: int = DEFAULT_MAX_ERRORS - max_turns: int = DEFAULT_MAX_STEPS - - -class Tau2Taskset(vf.Taskset[Tau2TasksetConfig]): - def load_toolsets(self, config: Tau2TasksetConfig) -> vf.Toolsets: - if "toolsets" in self.config.model_fields_set: - return None - return load_toolset(domain=config.domain) - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(domain=self.config.domain, max_turns=self.config.max_turns) - - @vf.setup(priority=100) - async def tau2_setup(self, task: vf.Task, state: vf.State) -> None: - if "setups" in self.config.model_fields_set: - return - runtime = state.runtime_state() - sampling_args = dict(DEFAULT_LLM_ARGS_AGENT) - sampling_args.update(dict(runtime.get("sampling_args") or {})) - runtime["sampling_args"] = sampling_args - task_info = task["info"] - if isinstance(task_info, str): - task_info = json.loads(task_info) - user_args = ( - DEFAULT_LLM_ARGS_USER - if self.config.user_args is None - else self.config.user_args - ) - session = Tau2Session( - domain=self.config.domain, - task_payload=cast(vf.JsonData, task_info), - user_model=self.config.user_model, - user_args=cast(vf.JsonData, user_args), - user_base_url=self.config.user_base_url, - user_api_key_var=self.config.user_api_key_var, - max_steps=self.config.max_steps, - max_errors=self.config.max_errors, - ) - state["tau2_session"] = session - await session.initialize(state) - state.setdefault("prompt", []) - if not state.get("tau2_prompt_initialized"): - state["prompt"].extend(session.initial_prompt_messages) - state["tau2_prompt_initialized"] = True - - @vf.cleanup - async def tau2_cleanup(self, state: vf.State) -> None: - state.pop("tau2_session", None) - - @vf.reward(weight=1.0) - async def tau2_reward(self, task: vf.Task, state: vf.State) -> float: - tau2_state = cast(vf.JsonData, state["tau2"]) - messages = [ - load_tau_message(cast(vf.JsonData, message)) - for message in cast(list[vf.JsonData], tau2_state["messages"]) - ] - termination = tau2_state.get("termination_reason") - if isinstance(termination, str): - termination_reason = TerminationReason(termination) - elif state.get("stop_condition") == "max_turns_reached": - termination_reason = TerminationReason.MAX_STEPS - else: - termination_reason = TerminationReason.AGENT_ERROR - state["tau2"]["termination_reason"] = termination_reason.value - task_info = task["info"] - if isinstance(task_info, str): - task_info = json.loads(task_info) - tau_task = TauTask.model_validate(task_info) - simulation = SimulationRun( - id=f"{task['taskset_id']}_{task['task_id']}_{datetime.now().isoformat()}", - task_id=tau_task.id, - messages=messages, - termination_reason=termination_reason, - timestamp=datetime.now().isoformat(), - start_time=datetime.now().isoformat(), - end_time=datetime.now().isoformat(), - duration=0.0, - agent_cost=0.0, - user_cost=0.0, - ) - reward_info = evaluate_simulation( - simulation=simulation, - task=tau_task, - evaluation_type=EvaluationType.ALL, - solo_mode=False, - domain=str(task["domain"]), - ) - state["tau2"]["evaluation"] = reward_info.model_dump(mode="json") - return float(reward_info.reward) - - @vf.metric - async def tau2_num_errors(self, task: vf.Task, state: vf.State) -> float: - _ = task - return float(state.get("tau2", {}).get("num_errors", 0.0)) - - @vf.metric - async def tau2_num_steps(self, task: vf.Task, state: vf.State) -> float: - _ = task - return float(state.get("tau2", {}).get("step_count", 0.0)) - - @vf.metric - async def tau2_num_assistant_tool_calls( - self, task: vf.Task, state: vf.State - ) -> float: - _ = task - return float(state.get("num_assistant_tool_calls", 0.0)) - - @vf.metric - async def tau2_num_user_tool_calls(self, task: vf.Task, state: vf.State) -> float: - _ = task - return float(state.get("num_user_tool_calls", 0.0)) - - -class Tau2EnvConfig(vf.EnvConfig): - taskset: Tau2TasksetConfig = Tau2TasksetConfig() - harness: vf.HarnessConfig = vf.HarnessConfig() - - -def load_environment(config: Tau2EnvConfig) -> vf.Env: - return vf.Env( - taskset=Tau2Taskset(config=config.taskset), - harness=vf.Harness(config=config.harness), - ) diff --git a/environments/tau2_bench_v1/tau2_bench_v1/__init__.py b/environments/tau2_bench_v1/tau2_bench_v1/__init__.py new file mode 100644 index 0000000000..36dd2e8994 --- /dev/null +++ b/environments/tau2_bench_v1/tau2_bench_v1/__init__.py @@ -0,0 +1 @@ +"""tau2-bench-v1 environment package.""" diff --git a/environments/tau2_bench_v1/tau2_bench_v1/servers/__init__.py b/environments/tau2_bench_v1/tau2_bench_v1/servers/__init__.py new file mode 100644 index 0000000000..597f8dd8b1 --- /dev/null +++ b/environments/tau2_bench_v1/tau2_bench_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for tau2-bench-v1.""" diff --git a/environments/tau2_bench_v1/tau2_bench_v1/servers/user/__init__.py b/environments/tau2_bench_v1/tau2_bench_v1/servers/user/__init__.py new file mode 100644 index 0000000000..be0177256d --- /dev/null +++ b/environments/tau2_bench_v1/tau2_bench_v1/servers/user/__init__.py @@ -0,0 +1,3 @@ +from .config import UserConfig + +__all__ = ["UserConfig"] diff --git a/environments/tau2_bench_v1/tau2_bench_v1/servers/user/config.py b/environments/tau2_bench_v1/tau2_bench_v1/servers/user/config.py new file mode 100644 index 0000000000..57349dd33b --- /dev/null +++ b/environments/tau2_bench_v1/tau2_bench_v1/servers/user/config.py @@ -0,0 +1,5 @@ +import verifiers.v1 as vf + + +class UserConfig(vf.UserConfig): + pass diff --git a/environments/tau2_bench_v1/tau2_bench_v1/servers/user/user.py b/environments/tau2_bench_v1/tau2_bench_v1/servers/user/user.py new file mode 100644 index 0000000000..32c2ee400a --- /dev/null +++ b/environments/tau2_bench_v1/tau2_bench_v1/servers/user/user.py @@ -0,0 +1,354 @@ +from __future__ import annotations + +import os +from copy import deepcopy +from typing import Protocol + +from tau2.data_model.message import ( + AssistantMessage, + Message, + MultiToolMessage, + ToolCall, + ToolMessage, + UserMessage, +) +from tau2.data_model.tasks import Task as TauTask +from tau2.environment.environment import Environment +from tau2.orchestrator.orchestrator import DEFAULT_FIRST_AGENT_MESSAGE +from tau2.registry import registry +from tau2.user.base import UserState, is_valid_user_history_message +from tau2.user.user_simulator import UserSimulator +from tau2.utils.utils import get_now + +import verifiers.v1 as vf +from verifiers.utils.client_utils import load_prime_config + +from .config import UserConfig + +UserInputMessage = AssistantMessage | ToolMessage | MultiToolMessage + + +class TauUserTool(Protocol): + pass + + +class User(vf.User[UserConfig]): + environment: Environment | None + messages: list[vf.JsonData] + user_simulator: UserSimulator | None + user_state: UserState | None + pending_user_input: UserInputMessage | None + bootstrap_user_message: UserMessage | None + num_errors: int + task_id: str + max_errors: int + + def start(self) -> None: + self.environment = None + self.messages = [] + self.user_simulator = None + self.user_state = None + self.pending_user_input = None + self.bootstrap_user_message = None + self.num_errors = 0 + self.task_id = "" + self.max_errors = 0 + + @vf.tool( + hidden=True, + args={ + "domain": "task.domain", + "tau2_task_json": "task.tau2_task_json", + "tau2_user": "task.tau2_user", + }, + sets={ + "tau2": "state.extras.tau2", + }, + ) + def setup(self, domain: str, tau2_task_json: str, tau2_user: dict) -> dict: + tau_task = TauTask.model_validate_json(tau2_task_json) + environment = registry.get_env_constructor(domain)() + self.environment = environment + self.task_id = tau_task.id + self.max_errors = int(tau2_user.get("max_errors") or 0) + initial_messages = self.initialize_environment(environment, tau_task) + user_model, user_args = self.user_model_config(tau2_user) + self.user_simulator = UserSimulator( + tools=self.user_tools(environment), + instructions=tau_task.user_scenario, + llm=user_model, + llm_args=user_args, + ) + self.user_state = self.user_simulator.get_init_state( + message_history=[ + message + for message in initial_messages + if is_valid_user_history_message(message) + ] + ) + self.messages = [self.message_data(message) for message in initial_messages] + self.bootstrap_user_message = self.initial_user_message(initial_messages) + self.pending_user_input = self.initial_user_input(initial_messages) + return { + "tau2": self.tau2_state(), + "tools": self.tau2_tool_defs(environment), + } + + @vf.user( + args={ + "tau2_task_json": "task.tau2_task_json", + "completion": "state.completion", + }, + sets={ + "tau2": "state.extras.tau2", + "stop_condition": "state.stop_condition", + }, + ) + def respond(self, tau2_task_json: str, completion: list[dict]) -> dict: + _ = TauTask.model_validate_json(tau2_task_json) + if self.user_simulator is None or self.user_state is None: + raise RuntimeError("Tau2 user simulator has not started.") + if not completion: + if self.bootstrap_user_message is not None: + user_message = self.bootstrap_user_message + self.bootstrap_user_message = None + return { + "messages": [self.v1_user_message(user_message)], + "tau2": self.tau2_state(), + } + if self.pending_user_input is None: + return {"messages": [], "tau2": self.tau2_state()} + return self.generate_user_response(self.pending_user_input) + text = self.latest_assistant_text(completion) + if text: + assistant_message = AssistantMessage(role="assistant", content=text) + return self.generate_user_response(assistant_message) + return {"messages": [], "tau2": self.tau2_state()} + + @vf.tool( + hidden=True, + sets={ + "tau2": "state.extras.tau2", + "finished": "state.is_completed", + "stop_condition": "state.stop_condition", + }, + ) + def call_tool(self, name: str, input: vf.JsonData) -> dict: + if self.environment is None: + raise RuntimeError("Tau2 environment has not started.") + tool_call = ToolCall( + id=name, + name=name, + arguments=input, + requestor="assistant", + ) + assistant_message = AssistantMessage(role="assistant", tool_calls=[tool_call]) + tool_message = self.environment.get_response(tool_call) + self.record_assistant_tool_exchange(assistant_message, [tool_message]) + payload: dict[str, object] = { + "content": tool_message.content or "", + "tau2": self.tau2_state(), + } + if self.too_many_errors(): + payload["finished"] = True + payload["stop_condition"] = "tau2_too_many_errors" + return payload + + def tau2_state(self) -> vf.JsonData: + return { + "task_id": self.task_id, + "step_count": len(self.messages), + "num_errors": self.num_errors, + "reward": 0.0, + "messages": list(self.messages), + } + + def initialize_environment( + self, environment: Environment, task: TauTask + ) -> list[Message]: + initial_state = task.initial_state + initialization_data = ( + initial_state.initialization_data if initial_state is not None else None + ) + initialization_actions = ( + initial_state.initialization_actions if initial_state is not None else None + ) + message_history = ( + deepcopy(initial_state.message_history or []) + if initial_state is not None and initial_state.message_history is not None + else [] + ) + for message in message_history: + message.turn_idx = None + environment.set_state( + initialization_data=initialization_data, + initialization_actions=initialization_actions, + message_history=message_history, + ) + environment.sync_tools() + return message_history + + def initial_user_message( + self, message_history: list[Message] + ) -> UserMessage | None: + if not message_history: + return None + last_message = message_history[-1] + if isinstance(last_message, UserMessage) and not last_message.is_tool_call(): + return last_message + return None + + def initial_user_input( + self, message_history: list[Message] + ) -> UserInputMessage | None: + if message_history: + last_message = message_history[-1] + if ( + isinstance(last_message, UserMessage) + and not last_message.is_tool_call() + ): + return None + if ( + isinstance(last_message, AssistantMessage) + and not last_message.is_tool_call() + ): + return last_message + if ( + isinstance(last_message, ToolMessage) + and last_message.requestor == "user" + ): + return last_message + raise ValueError( + "Tau2 initial message history must end with a user message, " + "assistant text message, or user tool result." + ) + first_message = deepcopy(DEFAULT_FIRST_AGENT_MESSAGE) + first_message.timestamp = get_now() + self.messages.append(self.message_data(first_message)) + return first_message + + def user_model_config(self, data: dict) -> tuple[str, dict]: + model = data.get("model") + if not isinstance(model, str) or not model: + raise TypeError("Tau2 user model must be a non-empty string.") + args = data.get("args") + llm_args = dict(args) if isinstance(args, dict) else {} + base_url = data.get("base_url") + if isinstance(base_url, str) and base_url: + llm_args.setdefault("api_base", base_url) + api_key_var = data.get("api_key_var") + if isinstance(api_key_var, str) and api_key_var: + api_key = os.environ.get(api_key_var) + if not api_key and api_key_var == "PRIME_API_KEY": + api_key = str(load_prime_config().get("api_key") or "") + if api_key: + llm_args.setdefault("api_key", api_key) + return model, llm_args + + def user_tools(self, environment: Environment) -> list[TauUserTool] | None: + try: + return list(environment.get_user_tools()) + except ValueError: + return None + + def generate_user_response(self, message: UserInputMessage) -> dict: + if self.environment is None: + raise RuntimeError("Tau2 environment has not started.") + if self.user_simulator is None or self.user_state is None: + raise RuntimeError("Tau2 user simulator has not started.") + self.pending_user_input = None + current = message + while True: + user_message, self.user_state = self.user_simulator.generate_next_message( + current, self.user_state + ) + self.messages.append(self.message_data(user_message)) + if self.user_simulator.is_stop(user_message): + return { + "messages": [], + "tau2": self.tau2_state(), + "stop_condition": "tau2_user_done", + } + if not user_message.is_tool_call(): + return { + "messages": [self.v1_user_message(user_message)], + "tau2": self.tau2_state(), + } + tool_messages = [ + self.environment.get_response(tool_call) + for tool_call in user_message.tool_calls or [] + ] + self.record_tool_results(tool_messages) + if self.too_many_errors(): + return { + "messages": [], + "tau2": self.tau2_state(), + "stop_condition": "tau2_too_many_errors", + } + current = ( + MultiToolMessage(role="tool", tool_messages=tool_messages) + if len(tool_messages) > 1 + else tool_messages[0] + ) + + def record_assistant_tool_exchange( + self, assistant_message: AssistantMessage, tool_messages: list[ToolMessage] + ) -> None: + self.messages.append(self.message_data(assistant_message)) + self.record_tool_results(tool_messages) + + def record_tool_results(self, tool_messages: list[ToolMessage]) -> None: + for tool_message in tool_messages: + if tool_message.error: + self.num_errors += 1 + self.messages.append(self.message_data(tool_message)) + + def too_many_errors(self) -> bool: + return self.max_errors > 0 and self.num_errors >= self.max_errors + + @staticmethod + def message_data(message: Message | ToolMessage) -> vf.JsonData: + return message.model_dump(mode="json", exclude_none=True) + + @staticmethod + def v1_user_message(message: UserMessage) -> vf.JsonData: + return {"role": "user", "content": message.content or ""} + + @staticmethod + def tau2_tool_defs(environment: Environment) -> list[vf.JsonData]: + tool_defs: list[vf.JsonData] = [] + for tool in environment.get_tools(): + schema = tool.openai_schema + function = schema.get("function") if isinstance(schema, dict) else None + if not isinstance(function, dict): + raise TypeError( + f"Tau2 tool {tool.name!r} did not expose OpenAI schema." + ) + parameters = function.get("parameters") + tool_defs.append( + { + "name": str(function.get("name") or tool.name), + "description": str(function.get("description") or ""), + "parameters": parameters + if isinstance(parameters, dict) + else {"type": "object", "properties": {}}, + } + ) + return tool_defs + + @staticmethod + def latest_assistant_text(completion: list[dict]) -> str: + for message in reversed(completion): + if not isinstance(message, dict) or message.get("role") != "assistant": + continue + content = message.get("content") + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for item in content: + if isinstance(item, dict) and isinstance(item.get("text"), str): + parts.append(item["text"]) + if parts: + return "\n".join(parts) + return "" diff --git a/environments/tau2_bench_v1/tau2_bench_v1/taskset.py b/environments/tau2_bench_v1/tau2_bench_v1/taskset.py new file mode 100644 index 0000000000..eb7ed7875d --- /dev/null +++ b/environments/tau2_bench_v1/tau2_bench_v1/taskset.py @@ -0,0 +1,188 @@ +import os +import json +import shutil +import subprocess +from pathlib import Path + +from pydantic import TypeAdapter +import verifiers.v1 as vf + +from .servers.user import UserConfig + +DEFAULT_USER_MODEL = "openai/openai/gpt-4.1-mini" +DEFAULT_USER_BASE_URL = "https://api.pinference.ai/api/v1" +DEFAULT_USER_API_KEY_VAR = "PRIME_API_KEY" +DEFAULT_MAX_STEPS = 30 +DEFAULT_MAX_ERRORS = 10 + + +class Tau2TasksetConfig(vf.TasksetConfig): + id: str | None = "tau2_telecom" + user: vf.UserConfig | None = UserConfig() + domain: str = "telecom" + user_model: str = DEFAULT_USER_MODEL + user_args: vf.JsonData | None = None + user_base_url: str = DEFAULT_USER_BASE_URL + user_api_key_var: str = DEFAULT_USER_API_KEY_VAR + max_steps: int = DEFAULT_MAX_STEPS + max_errors: int = DEFAULT_MAX_ERRORS + max_turns: int = DEFAULT_MAX_STEPS + + +class Tau2Task(vf.Task): + domain: str + tau2_task_json: str + tau2_user: vf.JsonData + + +class Tau2Taskset(vf.Taskset[Tau2TasksetConfig]): + task_type = Tau2Task + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return list( + load_tasks( + domain=self.config.domain, + max_turns=self.config.max_turns, + user_model=self.config.user_model, + user_args=self.config.user_args or {}, + user_base_url=self.config.user_base_url, + user_api_key_var=self.config.user_api_key_var, + max_errors=self.config.max_errors, + ) + ) + + @vf.reward(weight=1.0) + async def tau2_reward(self, task: Tau2Task, state: vf.State) -> float: + tau2 = state.extras.get("tau2") + if not isinstance(tau2, dict): + return 0.0 + messages_data = tau2.get("messages") + if not isinstance(messages_data, list): + return 0.0 + from tau2.data_model.message import Message as TauMessage + from tau2.data_model.simulation import SimulationRun, TerminationReason + from tau2.data_model.tasks import Task as TauTask + from tau2.evaluator.evaluator import EvaluationType, evaluate_simulation + from tau2.utils.utils import get_now + + messages = TypeAdapter(list[TauMessage]).validate_python(messages_data) + now = get_now() + simulation = SimulationRun( + id=state.id, + task_id=task.task_id, + start_time=now, + end_time=now, + duration=state.timing.total, + termination_reason=TerminationReason.USER_STOP, + messages=messages, + ) + reward_info = evaluate_simulation( + simulation=simulation, + task=TauTask.model_validate_json(task.tau2_task_json), + evaluation_type=EvaluationType.ALL, + solo_mode=False, + domain=task.domain, + ) + tau2["reward"] = float(reward_info.reward) + tau2["reward_info"] = reward_info.model_dump(mode="json", exclude_none=True) + return float(reward_info.reward) + + @vf.metric + async def tau2_num_steps(self, state: vf.State) -> float: + tau2 = state.extras.get("tau2") + if not isinstance(tau2, dict): + return 0.0 + value = tau2.get("step_count") + if isinstance(value, bool) or not isinstance(value, int | float | str): + return 0.0 + return float(value) + + @vf.metric + async def tau2_num_errors(self, state: vf.State) -> float: + tau2 = state.extras.get("tau2") + if not isinstance(tau2, dict): + return 0.0 + value = tau2.get("num_errors") + if isinstance(value, bool) or not isinstance(value, int | float | str): + return 0.0 + return float(value) + + +def download_tau2_data() -> None: + from tau2.utils.utils import DATA_DIR + + data_dir = Path(DATA_DIR) + if os.path.exists(data_dir) and os.path.exists(data_dir / "tau2" / "domains"): + return + os.makedirs(data_dir, exist_ok=True) + temp_dir = Path("/tmp/tau2_bench_v1") + try: + subprocess.run( + [ + "git", + "clone", + "--depth", + "1", + "https://github.com/sierra-research/tau2-bench.git", + str(temp_dir), + ], + check=True, + capture_output=True, + ) + source_data = temp_dir / "data" + if source_data.exists(): + shutil.copytree(source_data, data_dir, dirs_exist_ok=True) + finally: + if temp_dir.exists(): + shutil.rmtree(temp_dir) + + +def load_tasks( + domain: str, + max_turns: int, + user_model: str, + user_args: vf.JsonData, + user_base_url: str, + user_api_key_var: str, + max_errors: int, +): + from tau2.agent.llm_agent import AGENT_INSTRUCTION, SYSTEM_PROMPT + from tau2.registry import registry + from tau2.run import load_tasks as load_tau2_tasks + + download_tau2_data() + environment_constructor = registry.get_env_constructor(domain) + environment = environment_constructor() + system_prompt = SYSTEM_PROMPT.format( + agent_instruction=AGENT_INSTRUCTION, + domain_policy=environment.policy, + ) + for index, task in enumerate( + load_tau2_tasks(task_set_name=domain, task_split_name="base") + ): + yield { + "row_id": index, + "task_id": task.id, + "domain": domain, + "system_prompt": system_prompt, + "max_turns": max_turns, + "prompt": [], + "tau2_task_json": json.dumps( + task.model_dump(mode="json", exclude_none=True), + sort_keys=True, + separators=(",", ":"), + ), + "tau2_user": { + "model": user_model, + "args": user_args, + "base_url": user_base_url, + "api_key_var": user_api_key_var, + "max_errors": max_errors, + }, + } + + +def load_taskset(config: Tau2TasksetConfig) -> Tau2Taskset: + return Tau2Taskset(config=config) diff --git a/environments/wiki_search/pyproject.toml b/environments/wiki_search/pyproject.toml index 502e100ab2..7759da290d 100644 --- a/environments/wiki_search/pyproject.toml +++ b/environments/wiki_search/pyproject.toml @@ -16,7 +16,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["wiki_search.py", "wiki_search_v1.py", "pyproject.toml"] +include = ["wiki_search.py", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/wiki_search/wiki_search.py b/environments/wiki_search/wiki_search.py index 48a1c5b5f2..bd7dfef853 100644 --- a/environments/wiki_search/wiki_search.py +++ b/environments/wiki_search/wiki_search.py @@ -35,38 +35,7 @@ def load_environment( corpus_dataset: str = "willcb/rare-wiki-pages", corpus_split: str = "train", chroma_db_dir: str = CHROMA_DB_DIR, - v1: bool = False, ) -> vf.Environment: - if v1: - if ( - judge_base_url != DEFAULT_JUDGE_BASE_URL - or judge_api_key_var != DEFAULT_JUDGE_API_KEY_VAR - ): - raise ValueError( - 'v1 wiki_search judges through state.get_endpoint_config(api="chat"); ' - "set the rollout endpoint and only override judge_model." - ) - from wiki_search_v1 import ( - WikiSearchEnvConfig, - WikiSearchTasksetConfig, - load_environment as load_v1, - ) - - return load_v1( - config=WikiSearchEnvConfig( - taskset=WikiSearchTasksetConfig( - max_turns=max_turns, - judge_model=judge_model, - corpus_dataset=corpus_dataset, - corpus_split=corpus_split, - chroma_db_dir=chroma_db_dir, - embed_model=embed_model, - embed_base_url=embed_base_url, - embed_api_key_var=embed_api_key_var, - ) - ) - ) - # lazy corpus loading and chroma initialization _corpus_state: dict = { "loaded": False, diff --git a/environments/wiki_search/wiki_search_v1.py b/environments/wiki_search/wiki_search_v1.py deleted file mode 100644 index d8c0f2f655..0000000000 --- a/environments/wiki_search/wiki_search_v1.py +++ /dev/null @@ -1,319 +0,0 @@ -import asyncio -import os -from typing import cast - -import chromadb -from chromadb.api.types import Embeddable, EmbeddingFunction -from chromadb.utils import embedding_functions -from datasets import load_dataset - -import verifiers as vf - -CHROMA_DB_DIR = ".chroma_db" -_chroma_semaphore: asyncio.Semaphore | None = None - -SYSTEM_PROMPT = "Use the provided Wikipedia search tools to help answer questions." -JUDGE_PROMPT = """Given a ground truth answer \ -and a response, determine if the response is both correct and coherent. - -Question: -``` -{question} -``` - -Ground truth answer: -``` -{answer} -``` - -Response: -``` -{response} -``` - -Respond either "yes" or "no" only. - -If a response contains incoherent text, respond with "no" even if the correct answer is also present. -""" - - -def get_chroma_semaphore() -> asyncio.Semaphore: - global _chroma_semaphore - if _chroma_semaphore is None: - _chroma_semaphore = asyncio.Semaphore(100) - return _chroma_semaphore - - -def load_wiki( - corpus_dataset: str, - corpus_split: str, - chroma_db_dir: str, - embed_model: str, - embed_base_url: str, - embed_api_key_var: str, -) -> vf.ConfigData: - page_id_to_title: dict[str, str] = {} - page_id_to_content: dict[str, str] = {} - corpus = load_dataset(corpus_dataset, split=corpus_split) - for row in corpus: - row = cast(dict, row) - page_id_to_title[row["id"]] = row["title"] - page_id_to_content[row["id"]] = row["content"] - - openai_ef = embedding_functions.OpenAIEmbeddingFunction( - model_name=embed_model, - api_base=embed_base_url, - api_key=os.getenv(embed_api_key_var, "EMPTY"), - ) - client = chromadb.PersistentClient(path=chroma_db_dir) - collection = client.get_or_create_collection( - name="wiki_titles", - embedding_function=cast(EmbeddingFunction[Embeddable], openai_ef), - ) - init_chroma(collection, page_id_to_title) - return { - "collection": collection, - "page_id_to_title": page_id_to_title, - "page_id_to_content": page_id_to_content, - } - - -def init_chroma(collection, page_id_to_title: dict[str, str]) -> None: - all_ids = list(page_id_to_title) - existing: set[str] = set() - for i in range(0, len(all_ids), 500): - batch = all_ids[i : i + 500] - got = collection.get(ids=batch) - existing.update(got.get("ids", [])) - missing = [page_id for page_id in all_ids if page_id not in existing] - if not missing: - return - documents = [] - metadatas = [] - for page_id in missing: - title = str(page_id_to_title[page_id]).strip() - if not title: - raise ValueError(f"Empty title for page_id {page_id}") - documents.append(title) - metadatas.append({"title": title}) - for i in range(0, len(missing), 100): - collection.upsert( - ids=missing[i : i + 100], - documents=documents[i : i + 100], - metadatas=metadatas[i : i + 100], - ) - - -def normalize_id(text: str) -> str: - return text.strip().lower().replace(" ", "_") - - -async def search_pages(query: str, wiki) -> list[dict]: - """Search for top 10 relevant articles using title embedding similarity.""" - async with get_chroma_semaphore(): - results = await asyncio.to_thread( - wiki["collection"].query, query_texts=[query], n_results=10 - ) - if not results or not results["metadatas"]: - raise ValueError(f"No results found for query: {query}") - output = [] - for i in range(len(results["ids"][0])): - output.append( - { - "page_id": results["ids"][0][i], - "title": results["metadatas"][0][i]["title"], - } - ) - return output - - -async def view_sections(page_id: str, wiki) -> list[dict]: - """View the sections of a page.""" - content = wiki["page_id_to_content"][page_id] - sections = [] - lines = content.split("\n") - for i, line in enumerate(lines): - if line.startswith("#"): - section_name = line.lstrip("#").strip() - sections.append( - { - "section_id": f"{page_id}:{normalize_id(section_name)}", - "section_name": section_name, - "start_line": i, - } - ) - if not sections: - sections.append( - { - "section_id": f"{page_id}:full", - "section_name": "Full Page", - "start_line": 0, - } - ) - return [ - {"section_id": section["section_id"], "section_name": section["section_name"]} - for section in sections - ] - - -async def read_section(section_id: str, wiki) -> str: - """Read a section of a page.""" - if ":" not in section_id: - raise ValueError("Invalid section_id format. Expected: page_id:section_name") - page_id, section_name_id = section_id.split(":", 1) - content = wiki["page_id_to_content"][page_id] - if section_name_id == "full": - return content - lines = content.split("\n") - section_start = None - section_end = None - for i, line in enumerate(lines): - if line.startswith("#"): - current_section = normalize_id(line.lstrip("#").strip()) - if current_section == section_name_id and section_start is None: - section_start = i - elif section_start is not None and section_end is None: - section_end = i - break - if section_start is None: - raise ValueError(f"Section not found: {section_id}") - return "\n".join(lines[section_start : section_end or len(lines)]) - - -def load_tasks( - max_turns: int = 10, - judge_model: str | None = None, -): - dataset = load_dataset("willcb/wiki-trivia-questions-v4", split="train") - for index, row in enumerate(dataset): - row = cast(dict, row) - task = { - **row, - "example_id": index, - "max_turns": max_turns, - "prompt": [{"role": "user", "content": row["question"]}], - } - if judge_model is not None: - task["judge_model"] = judge_model - yield task - - -@vf.reward(weight=1.0) -async def judge_reward(task, state) -> float: - completion = state.get("completion") or [] - messages = vf.get_messages(completion, role="assistant") - response = str(messages[-1].content or "") if messages else "" - prompt = JUDGE_PROMPT.format( - question=task["question"], - answer=task["answer"], - response=response, - ) - endpoint_config = state.get_endpoint_config(api="chat") - judge_model = task.get("judge_model") or endpoint_config.model - judge_client = state.get_client(api="chat") - try: - result = await judge_client.chat.completions.create( - model=str(judge_model), - messages=[{"role": "user", "content": prompt}], - ) - finally: - await judge_client.close() - text = result.choices[0].message.content or "" - return 1.0 if "yes" in text.lower() else 0.0 - - -def load_toolset( - corpus_dataset: str = "willcb/rare-wiki-pages", - corpus_split: str = "train", - chroma_db_dir: str = CHROMA_DB_DIR, - embed_model: str = "text-embedding-3-small", - embed_base_url: str = "https://api.openai.com/v1", - embed_api_key_var: str = "OPENAI_API_KEY", - config=None, -): - def load_wiki_index() -> vf.ConfigData: - return load_wiki( - corpus_dataset=corpus_dataset, - corpus_split=corpus_split, - chroma_db_dir=chroma_db_dir, - embed_model=embed_model, - embed_base_url=embed_base_url, - embed_api_key_var=embed_api_key_var, - ) - - wiki_index: vf.ConfigData | None = None - - def wiki() -> vf.ConfigData: - nonlocal wiki_index - if wiki_index is None: - wiki_index = load_wiki_index() - return wiki_index - - async def search_pages_tool(query: str) -> list[dict]: - return await search_pages(query, wiki()) - - async def view_sections_tool(page_id: str) -> list[dict]: - return await view_sections(page_id, wiki()) - - async def read_section_tool(section_id: str) -> str: - return await read_section(section_id, wiki()) - - search_pages_tool.__name__ = "search_pages" - search_pages_tool.__doc__ = search_pages.__doc__ - view_sections_tool.__name__ = "view_sections" - view_sections_tool.__doc__ = view_sections.__doc__ - read_section_tool.__name__ = "read_section" - read_section_tool.__doc__ = read_section.__doc__ - - return vf.Toolset( - tools=[search_pages_tool, view_sections_tool, read_section_tool], - config=config, - ) - - -class WikiSearchTasksetConfig(vf.TasksetConfig): - rewards: list[str] = ["judge_reward"] - max_turns: int = 10 - corpus_dataset: str = "willcb/rare-wiki-pages" - corpus_split: str = "train" - chroma_db_dir: str = CHROMA_DB_DIR - embed_model: str = "text-embedding-3-small" - embed_base_url: str = "https://api.openai.com/v1" - embed_api_key_var: str = "OPENAI_API_KEY" - judge_model: str | None = None - - -class WikiSearchEnvConfig(vf.EnvConfig): - taskset: WikiSearchTasksetConfig = WikiSearchTasksetConfig() - harness: vf.HarnessConfig = vf.HarnessConfig() - - -class WikiSearchTaskset(vf.Taskset[WikiSearchTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks( - max_turns=self.config.max_turns, - judge_model=self.config.judge_model, - ) - - def load_system_prompt(self, config: WikiSearchTasksetConfig) -> vf.SystemPrompt: - _ = config - return SYSTEM_PROMPT - - def load_toolsets(self, config: WikiSearchTasksetConfig) -> vf.Toolsets: - return { - "wiki": load_toolset( - corpus_dataset=config.corpus_dataset, - corpus_split=config.corpus_split, - chroma_db_dir=config.chroma_db_dir, - embed_model=config.embed_model, - embed_base_url=config.embed_base_url, - embed_api_key_var=config.embed_api_key_var, - ) - } - - -def load_environment(config: WikiSearchEnvConfig) -> vf.Env: - return vf.Env( - taskset=WikiSearchTaskset(config=config.taskset), - harness=vf.Harness(config=config.harness), - ) diff --git a/environments/wiki_search_v1/README.md b/environments/wiki_search_v1/README.md new file mode 100644 index 0000000000..5e192e39c1 --- /dev/null +++ b/environments/wiki_search_v1/README.md @@ -0,0 +1,64 @@ +# wiki-search-v1 + + +Source Code + + +### Overview +- **Environment ID**: `wiki-search-v1` +- **Short description**: Multi-turn tool-use QA over a small Wikipedia corpus using a v1 task-owned MCP toolset. +- **Tags**: retrieval, tools, multi-turn, v1 + +### Datasets +- **Primary dataset(s)**: `willcb/wiki-trivia-questions` (HF) and a Wikipedia corpus from `willcb/rare-wiki-pages` +- **Source links**: Hugging Face Datasets +- **Split sizes**: Uses the `train` split for prompts + +### Task +- **Type**: `vf.Env` with a wiki QA `vf.Taskset`, base `vf.Harness`, and env-scope wiki `Toolset`. +- **Rubric overview**: Answer-substring reward against the reference answer. + +### How it works +- **Corpus load**: Reads `willcb/rare-wiki-pages` (HF) into memory: `id → title`, `id → content`. +- **Tools**: + - `search_pages(query)`: Lexical ranking over page titles and content; returns top 10 `{page_id, title}`. + - `view_sections(page_id)`: Parses the page content for Markdown-style headings (`# ...`) and returns section ids/names. Falls back to a single `full` section if no headings. + - `read_section(section_id)`: Returns the content slice for the requested section (or full page). +- **Scoring**: The taskset rewards final assistant responses that contain the reference answer. + +### Quickstart +Run an evaluation with default settings: + +```bash +prime eval run wiki-search-v1 +``` + +Configure model and sampling: + +```bash +prime eval run wiki-search-v1 \ + -m openai/gpt-4.1-mini \ + -n 20 -r 3 -t 1024 -T 0.7 \ + -a '{"config": {"taskset": {"max_turns": 10, "toolsets": {"wiki": {"corpus_dataset": "willcb/rare-wiki-pages", "corpus_split": "train"}}}}}' +``` + +### Required Environment Variables + +No task-specific environment variables are required. + +### Taskset Config +| Field | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `max_turns` | int | `10` | Maximum model turns per rollout | + +### Wiki Toolset Config +| Field | Type | Default | Description | +| --- | ---- | ------- | ----------- | +| `corpus_dataset` | str | `"willcb/rare-wiki-pages"` | HF dataset id containing pages | +| `corpus_split` | str | `"train"` | HF split to load | + +### Metrics +| Metric | Meaning | +| ------ | ------- | +| `reward` | 1.0 if the final assistant response contains the reference answer, else 0.0 | +| `num_turns` | Number of recorded model turns | diff --git a/environments/wiki_search_v1/pyproject.toml b/environments/wiki_search_v1/pyproject.toml new file mode 100644 index 0000000000..d793bb0251 --- /dev/null +++ b/environments/wiki_search_v1/pyproject.toml @@ -0,0 +1,21 @@ +[project] +name = "wiki-search-v1" +description = "Agentic RAG over Wikipedia pages for trivia Q&A" +tags = ["wikipedia", "multi-turn", "agentic-search", "rag", "train", "eval", "llm-judge"] +requires-python = ">=3.11" +version = "0.1.23" +dependencies = [ + "verifiers>=0.1.9", + "datasets", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build] +include = ["wiki_search_v1/**/*", "README.md", "pyproject.toml"] + +[tool.verifiers.eval] +num_examples = 5 +rollouts_per_example = 3 diff --git a/environments/wiki_search_v1/wiki_search_v1/__init__.py b/environments/wiki_search_v1/wiki_search_v1/__init__.py new file mode 100644 index 0000000000..a479f9a213 --- /dev/null +++ b/environments/wiki_search_v1/wiki_search_v1/__init__.py @@ -0,0 +1 @@ +"""wiki-search-v1 environment package.""" diff --git a/environments/wiki_search_v1/wiki_search_v1/servers/__init__.py b/environments/wiki_search_v1/wiki_search_v1/servers/__init__.py new file mode 100644 index 0000000000..00929ef525 --- /dev/null +++ b/environments/wiki_search_v1/wiki_search_v1/servers/__init__.py @@ -0,0 +1 @@ +"""MCP servers for wiki-search-v1.""" diff --git a/environments/wiki_search_v1/wiki_search_v1/servers/wiki/__init__.py b/environments/wiki_search_v1/wiki_search_v1/servers/wiki/__init__.py new file mode 100644 index 0000000000..e68a71b726 --- /dev/null +++ b/environments/wiki_search_v1/wiki_search_v1/servers/wiki/__init__.py @@ -0,0 +1,3 @@ +from .config import WikiToolsetConfig + +__all__ = ["WikiToolsetConfig"] diff --git a/environments/wiki_search_v1/wiki_search_v1/servers/wiki/config.py b/environments/wiki_search_v1/wiki_search_v1/servers/wiki/config.py new file mode 100644 index 0000000000..6776e8c13f --- /dev/null +++ b/environments/wiki_search_v1/wiki_search_v1/servers/wiki/config.py @@ -0,0 +1,8 @@ +import verifiers.v1 as vf + + +class WikiToolsetConfig(vf.ToolsetConfig): + scope: vf.Scope = "env" + startup_timeout_seconds: float = 60.0 + corpus_dataset: str = "willcb/rare-wiki-pages" + corpus_split: str = "train" diff --git a/environments/wiki_search_v1/wiki_search_v1/servers/wiki/toolset.py b/environments/wiki_search_v1/wiki_search_v1/servers/wiki/toolset.py new file mode 100644 index 0000000000..639301ff36 --- /dev/null +++ b/environments/wiki_search_v1/wiki_search_v1/servers/wiki/toolset.py @@ -0,0 +1,33 @@ +import verifiers.v1 as vf + +from wiki_search_v1.taskset import ( + WikiIndex, + load_wiki, + read_section, + search_pages, + view_sections, +) + +from .config import WikiToolsetConfig + + +class WikiToolset(vf.Toolset[WikiToolsetConfig]): + @vf.resource + def wiki(self) -> WikiIndex: + return load_wiki(self.config) + + @vf.tool(args={"wiki": "resources.wiki"}) + async def search_pages_tool( + self, query: str, wiki: WikiIndex + ) -> list[dict[str, str]]: + return await search_pages(query, wiki) + + @vf.tool(args={"wiki": "resources.wiki"}) + async def view_sections_tool( + self, page_id: str, wiki: WikiIndex + ) -> list[dict[str, str]]: + return await view_sections(page_id, wiki) + + @vf.tool(args={"wiki": "resources.wiki"}) + async def read_section_tool(self, section_id: str, wiki: WikiIndex) -> str: + return await read_section(section_id, wiki) diff --git a/environments/wiki_search_v1/wiki_search_v1/taskset.py b/environments/wiki_search_v1/wiki_search_v1/taskset.py new file mode 100644 index 0000000000..fbb05c113c --- /dev/null +++ b/environments/wiki_search_v1/wiki_search_v1/taskset.py @@ -0,0 +1,164 @@ +from __future__ import annotations + +import re +from typing import cast + +from datasets import load_dataset + +import verifiers.v1 as vf + +from .servers.wiki import WikiToolsetConfig + +SYSTEM_PROMPT = "Use the provided Wikipedia search tools to help answer questions." +TOKEN_RE = re.compile(r"[a-z0-9]+") + + +class WikiIndex: + def __init__( + self, + *, + page_id_to_title: dict[str, str], + page_id_to_content: dict[str, str], + ): + self.page_id_to_title = page_id_to_title + self.page_id_to_content = page_id_to_content + + +class WikiSearchTasksetConfig(vf.TasksetConfig): + max_turns: int = 10 + toolsets: vf.ToolsetConfigs = {"wiki": WikiToolsetConfig()} + + +class WikiSearchTask(vf.Task): + question: str + answer: str + + +class WikiSearchTaskset(vf.Taskset[WikiSearchTasksetConfig]): + task_type = WikiSearchTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + _ = split + dataset = load_dataset("willcb/wiki-trivia-questions-v4", split="train") + for index, row in enumerate(dataset): + record = cast(dict[str, object], row) + yield { + "row_id": index, + "question": str(record["question"]), + "answer": str(record["answer"]), + "max_turns": self.config.max_turns, + "prompt": [{"role": "user", "content": str(record["question"])}], + } + + def load_system_prompt(self, config: WikiSearchTasksetConfig) -> vf.SystemPrompt: + _ = config + return SYSTEM_PROMPT + + @vf.reward(weight=1.0) + async def answer_in_response(self, task: WikiSearchTask, state: vf.State) -> float: + messages = [ + message for message in state.completion if message.role == "assistant" + ] + response = str(messages[-1].content or "") if messages else "" + return float(task.answer.lower() in response.lower()) + + +def load_wiki(config: WikiToolsetConfig) -> WikiIndex: + page_id_to_title: dict[str, str] = {} + page_id_to_content: dict[str, str] = {} + corpus = load_dataset(config.corpus_dataset, split=config.corpus_split) + for row in corpus: + record = cast(dict[str, object], row) + page_id = str(record["id"]) + page_id_to_title[page_id] = str(record["title"]) + page_id_to_content[page_id] = str(record["content"]) + return WikiIndex( + page_id_to_title=page_id_to_title, + page_id_to_content=page_id_to_content, + ) + + +def normalize_id(text: str) -> str: + return text.strip().lower().replace(" ", "_") + + +def tokenize(text: str) -> set[str]: + return set(TOKEN_RE.findall(text.lower())) + + +async def search_pages(query: str, wiki: WikiIndex) -> list[dict[str, str]]: + query_tokens = tokenize(query) + if not query_tokens: + raise ValueError("Search query must contain at least one alphanumeric token.") + ranked: list[tuple[int, str, str]] = [] + for page_id, title in wiki.page_id_to_title.items(): + title_score = len(query_tokens & tokenize(title)) + content_score = len(query_tokens & tokenize(wiki.page_id_to_content[page_id])) + score = 5 * title_score + content_score + ranked.append((-score, title.lower(), page_id)) + ranked.sort() + return [ + { + "page_id": page_id, + "title": wiki.page_id_to_title[page_id], + } + for _, _, page_id in ranked[:10] + ] + + +async def view_sections(page_id: str, wiki: WikiIndex) -> list[dict[str, str]]: + content = wiki.page_id_to_content[page_id] + sections: list[dict[str, str | int]] = [] + lines = content.split("\n") + for index, line in enumerate(lines): + if line.startswith("#"): + section_name = line.lstrip("#").strip() + sections.append( + { + "section_id": f"{page_id}:{normalize_id(section_name)}", + "section_name": section_name, + "start_line": index, + } + ) + if not sections: + sections.append( + { + "section_id": f"{page_id}:full", + "section_name": "Full Page", + "start_line": 0, + } + ) + return [ + { + "section_id": str(section["section_id"]), + "section_name": str(section["section_name"]), + } + for section in sections + ] + + +async def read_section(section_id: str, wiki: WikiIndex) -> str: + if ":" not in section_id: + raise ValueError("Invalid section_id format. Expected: page_id:section_name") + page_id, section_name_id = section_id.split(":", 1) + content = wiki.page_id_to_content[page_id] + if section_name_id == "full": + return content + lines = content.split("\n") + section_start = None + section_end = None + for index, line in enumerate(lines): + if line.startswith("#"): + current_section = normalize_id(line.lstrip("#").strip()) + if current_section == section_name_id and section_start is None: + section_start = index + elif section_start is not None and section_end is None: + section_end = index + break + if section_start is None: + raise ValueError(f"Section not found: {section_id}") + return "\n".join(lines[section_start : section_end or len(lines)]) + + +def load_taskset(config: WikiSearchTasksetConfig) -> WikiSearchTaskset: + return WikiSearchTaskset(config=config) diff --git a/environments/wordle_v1/README.md b/environments/wordle_v1/README.md index de44d8574d..7e424c0fed 100644 --- a/environments/wordle_v1/README.md +++ b/environments/wordle_v1/README.md @@ -16,5 +16,5 @@ prime eval run wordle-v1 ### Configuration The environment uses the packaged `TextArenaTaskset` for generic TextArena mechanics. -`wordle_v1.py` owns the Wordle prompt, `WordleUser` response shaping, rewards, +`wordle_v1/taskset.py` owns the Wordle prompt, `WordleUser` response shaping, rewards, and defaults for `Wordle-v0`. diff --git a/environments/wordle_v1/pyproject.toml b/environments/wordle_v1/pyproject.toml index 1ed6bb19a1..5926a092e4 100644 --- a/environments/wordle_v1/pyproject.toml +++ b/environments/wordle_v1/pyproject.toml @@ -14,7 +14,7 @@ requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build] -include = ["wordle_v1.py", "pyproject.toml", "README.md"] +include = ["wordle_v1/**/*", "README.md", "pyproject.toml"] [tool.verifiers.eval] num_examples = 5 diff --git a/environments/wordle_v1/wordle_v1/__init__.py b/environments/wordle_v1/wordle_v1/__init__.py new file mode 100644 index 0000000000..d6786ba847 --- /dev/null +++ b/environments/wordle_v1/wordle_v1/__init__.py @@ -0,0 +1 @@ +"""wordle-v1 environment package.""" diff --git a/environments/wordle_v1/wordle_v1.py b/environments/wordle_v1/wordle_v1/taskset.py similarity index 56% rename from environments/wordle_v1/wordle_v1.py rename to environments/wordle_v1/wordle_v1/taskset.py index 1cf58fc5e9..49f0ac4ab2 100644 --- a/environments/wordle_v1/wordle_v1.py +++ b/environments/wordle_v1/wordle_v1/taskset.py @@ -1,11 +1,10 @@ import re -import verifiers as vf +import verifiers.v1 as vf from tasksets.textarena import ( + TextArenaTask, TextArenaTaskset, TextArenaTasksetConfig, - TextArenaUser, - TextArenaUserConfig, ) WORDLE_SYSTEM_PROMPT = """You are a competitive game player. \ @@ -14,33 +13,10 @@ In each turn, think step-by-step, then give your guess inside ... tags.""" -class WordleUserConfig(TextArenaUserConfig): - pass - - class WordleTasksetConfig(TextArenaTasksetConfig): game: str = "Wordle-v0" answer_state_key: str = "secret_word" - user: WordleUserConfig | None = WordleUserConfig() - system_prompt: vf.PromptInput | vf.SystemPromptConfig | None = WORDLE_SYSTEM_PROMPT - - -class WordleUser(TextArenaUser): - config: WordleUserConfig - - async def get_response( - self, task: vf.Task, state: vf.State, messages: list[vf.Message] - ) -> list[vf.UserMessage]: - response = await super().get_response(task, state, messages) - if state.get("done") is True: - return response - assert len(response) == 1 - content = response[0].content - assert isinstance(content, str) - latest_feedback = content.split("[GAME]")[-1].strip() - if "Feedback:" in latest_feedback: - latest_feedback = latest_feedback.split("Feedback:")[-1] - return [vf.UserMessage(content=latest_feedback)] + system_prompt: vf.SystemPrompt = WORDLE_SYSTEM_PROMPT class WordleTaskset(TextArenaTaskset[WordleTasksetConfig]): @@ -51,12 +27,10 @@ def guesses(self, content: str) -> list[str]: return re.findall(self.guess_pattern, content, re.DOTALL) @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: - answer = task["answer"] - assert isinstance(answer, str) - completion = state.get("completion") or [] - assert isinstance(completion, list) - for message in reversed(vf.get_messages(completion)): + async def correct_answer(self, task: TextArenaTask, state: vf.State) -> float: + answer = task.answer + completion = state.completion + for message in reversed(completion): if not isinstance(message, vf.AssistantMessage): continue content = message.content @@ -67,14 +41,12 @@ async def correct_answer(self, task: vf.Task, state: vf.State) -> float: return 0.0 @vf.reward(weight=1.0) - async def length_bonus(self, task: vf.Task, state: vf.State) -> float: - answer = task["answer"] - assert isinstance(answer, str) - completion = state.get("completion") or [] - assert isinstance(completion, list) + async def length_bonus(self, task: TextArenaTask, state: vf.State) -> float: + answer = task.answer + completion = state.completion guess = "" num_guesses = 0 - for message in vf.get_messages(completion): + for message in completion: if not isinstance(message, vf.AssistantMessage): continue content = message.content @@ -89,12 +61,10 @@ async def length_bonus(self, task: vf.Task, state: vf.State) -> float: return is_correct / (num_guesses or 1) @vf.reward(weight=1.0) - async def partial_answer(self, task: vf.Task, state: vf.State) -> float: - answer = task["answer"] - assert isinstance(answer, str) - completion = state.get("completion") or [] - assert isinstance(completion, list) - for message in reversed(vf.get_messages(completion)): + async def partial_answer(self, task: TextArenaTask, state: vf.State) -> float: + answer = task.answer + completion = state.completion + for message in reversed(completion): if not isinstance(message, vf.AssistantMessage): continue content = message.content @@ -104,7 +74,7 @@ async def partial_answer(self, task: vf.Task, state: vf.State) -> float: if matches[-1].strip() == f"[{answer}]": return 0.0 break - for message in reversed(vf.get_messages(completion)): + for message in reversed(completion): if not isinstance(message, vf.UserMessage): continue content = message.content @@ -118,10 +88,9 @@ async def partial_answer(self, task: vf.Task, state: vf.State) -> float: @vf.reward(weight=0.2) async def format_reward(self, task: vf.Task, state: vf.State) -> float: _ = task - completion = state.get("completion") or [] - assert isinstance(completion, list) + completion = state.completion found = False - for message in vf.get_messages(completion): + for message in completion: if not isinstance(message, vf.AssistantMessage): continue found = True @@ -134,11 +103,3 @@ async def format_reward(self, task: vf.Task, state: vf.State) -> float: def load_taskset(config: WordleTasksetConfig) -> WordleTaskset: return WordleTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) diff --git a/packages/harnesses/README.md b/packages/harnesses/README.md index 191bb71318..758f84110a 100644 --- a/packages/harnesses/README.md +++ b/packages/harnesses/README.md @@ -1,11 +1,10 @@ # harnesses -Reusable v1 `vf.Harness` implementations for Verifiers. +Reusable `verifiers.v1` harness implementations. -Harnesses own rollout execution: programs, command agents, framework adapters, -endpoint interception, primary sandbox placement, execution setup, and execution -artifacts. Task data, task-owned tools, users, rewards, and task-specific config -belong to tasksets. +Harnesses own reusable execution mechanisms: command agents, framework adapters, +runtime calls, and execution artifacts. Task data, task tools, users, rewards, +and task-specific config belong to tasksets. ## Install @@ -26,24 +25,16 @@ Environment packages should expose a typed child loader and let Verifiers coerce the `[env.harness]` config through that annotation: ```python -import verifiers as vf +import verifiers.v1 as vf from harnesses import OpenCode, OpenCodeConfig def load_harness(config: OpenCodeConfig) -> OpenCode: return OpenCode(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` -Use `vf.load_harness(config=config.harness)` when the environment does not -own a reusable execution mechanism. +Omit `harness.py` when the environment does not own a reusable execution +mechanism; the component loader will use the base harness. ## Included Harnesses @@ -54,22 +45,22 @@ own a reusable execution mechanism. | `MiniSWEAgent` | mini-swe-agent. | | `Terminus2` | Harbor Terminus agent. | | `RLM` | Recursive language model command harness. | -| `ReplayHarness` | Replays stored assistant messages into trajectory steps without model calls. | +| `ReplayHarness` | Replays stored assistant messages into transcript turns without model calls. | | `NeMoGymHarness` | NeMo Gym rollout collection. | -Harness implementations resolve to one `ProgramConfig` shape. Command harness -configs may expose task-relevant execution knobs, but the harness owns command -construction, channel wiring, sandbox placement, and artifacts. +Command harness configs may expose task-relevant execution knobs, but the +harness owns command construction and records command output in +`state.artifacts`. ## Replay Stored Transcripts Use `ReplayHarness` when each task row already contains a top-level `messages` -chat transcript and each assistant message should become one trajectory step: +chat transcript and each assistant message should become one transcript turn: ```python from pathlib import Path -import verifiers as vf +import verifiers.v1 as vf from harnesses import ReplayHarness from tasksets import ReplayTaskset, ReplayTasksetConfig @@ -91,7 +82,7 @@ Non-assistant messages may appear before, between, or after assistant messages. `vf.HarnessConfig` defaults to replaying every assistant message; set `max_turns` only when the replay should be capped. -The replayed trajectory keeps `tokens=None`; token IDs and logprobs remain the +The replayed transcript keeps `tokens=None`; token IDs and logprobs remain the responsibility of the trainer or renderer that consumes the final transcript. ## Agent Versions diff --git a/packages/harnesses/harnesses/__init__.py b/packages/harnesses/harnesses/__init__.py index f0f233ccf4..b3ac2ee5ee 100644 --- a/packages/harnesses/harnesses/__init__.py +++ b/packages/harnesses/harnesses/__init__.py @@ -1,11 +1,12 @@ __version__ = "0.1.2" -from .mini_swe_agent import MiniSWEAgent, MiniSWEAgentConfig, MiniSWEAgentProgramConfig -from .opencode import OpenCode, OpenCodeConfig, OpenCodeProgramConfig -from .pi import Pi, PiConfig, PiProgramConfig +from .command import CommandHarness, CommandHarnessConfig +from .mini_swe_agent import MiniSWEAgent, MiniSWEAgentConfig +from .opencode import OpenCode, OpenCodeConfig +from .pi import Pi, PiConfig from .replay import ReplayHarness -from .rlm import RLM, RLMConfig, RLMProgramConfig -from .terminus_2 import Terminus2, Terminus2Config, Terminus2ProgramConfig +from .rlm import RLM, RLMConfig +from .terminus_2 import Terminus2, Terminus2Config LAZY_EXPORTS = { "NeMoGymHarness": (".nemo_gym", "NeMoGymHarness"), @@ -13,23 +14,20 @@ } __all__ = [ + "CommandHarness", + "CommandHarnessConfig", "MiniSWEAgent", "MiniSWEAgentConfig", - "MiniSWEAgentProgramConfig", *LAZY_EXPORTS, "OpenCode", "OpenCodeConfig", - "OpenCodeProgramConfig", "Pi", "PiConfig", - "PiProgramConfig", "ReplayHarness", "RLM", "RLMConfig", - "RLMProgramConfig", "Terminus2", "Terminus2Config", - "Terminus2ProgramConfig", ] diff --git a/packages/harnesses/harnesses/command.py b/packages/harnesses/harnesses/command.py new file mode 100644 index 0000000000..4e9eaa4602 --- /dev/null +++ b/packages/harnesses/harnesses/command.py @@ -0,0 +1,85 @@ +import json +from typing import Generic, TypeVar + +from pydantic import Field + +import verifiers.v1 as vf + + +class CommandHarnessConfig(vf.HarnessConfig): + command: list[str] = Field(default_factory=list) + cwd: str | None = None + env: dict[str, str] = Field(default_factory=dict) + timeout_seconds: float | None = None + + +ConfigT = TypeVar("ConfigT", bound=CommandHarnessConfig) + + +class CommandHarness(vf.Harness[ConfigT], Generic[ConfigT]): + config: ConfigT + + def command(self, task: vf.Task, state: vf.State) -> list[str]: + _ = task, state + if not self.config.command: + raise ValueError(f"{type(self).__name__} requires config.command.") + return list(self.config.command) + + def command_env(self, task: vf.Task, state: vf.State) -> dict[str, str]: + prompt_text = "\n\n".join( + str(getattr(message, "content", "") or "") + for message in self.initial_messages(task) + ) + return { + **self.config.env, + "VF_TASK_JSON": json.dumps( + task.model_dump(mode="json", exclude_none=True, exclude_defaults=True), + ensure_ascii=False, + ), + "VF_STATE_ID": state.id, + "VF_PROMPT": prompt_text, + } + + async def run_with_context(self, context: vf.Context) -> None: + task = context.task + state = context.state + runtime = context.runtime + if runtime is None: + raise ValueError("CommandHarness requires a runtime.") + prompt = self.initial_messages(task) + + async def stop_check() -> str | None: + if await self.is_completed(context): + return state.stop_condition or "stop" + return None + + async with vf.InterceptionServer( + context, + task, + state, + protocols=self.protocols, + stop_check=stop_check, + ) as endpoint: + endpoint_url = await runtime.expose(endpoint.port) + result = await runtime.run( + self.command(task, state), + cwd=self.config.cwd, + env={ + **self.command_env(task, state), + **endpoint.env(base_url=endpoint_url, model=context.model), + }, + timeout=self.config.timeout_seconds, + ) + state.artifacts["command"] = result.model_dump(mode="json") + content = result.stdout.strip() or result.stderr.strip() + if not state.transcript: + message = vf.AssistantMessage(content=content) + state.transcript.append(vf.Turn(prompt=prompt, completion=[message])) + if result.returncode == 0: + state.stop("command_completed") + else: + state.stop("command_failed") + + +def shell_command(command: str) -> list[str]: + return ["bash", "-lc", command] diff --git a/packages/harnesses/harnesses/mini_swe_agent.py b/packages/harnesses/harnesses/mini_swe_agent.py index 29aa084344..a22d8b1cc2 100644 --- a/packages/harnesses/harnesses/mini_swe_agent.py +++ b/packages/harnesses/harnesses/mini_swe_agent.py @@ -1,175 +1,71 @@ import shlex -from pathlib import PurePosixPath -import verifiers as vf -from verifiers.v1.utils.sandbox_python_utils import python_runtime_setup_command +from pydantic import Field -from .utils import split_versioned_agent_spec +import verifiers.v1 as vf -DEFAULT_INSTALL_DIR = "/opt/mini-swe-agent" -DEFAULT_PREFIX_DIR = f"{DEFAULT_INSTALL_DIR}/prefix" -DEFAULT_SITE_PACKAGES_DIR = f"{DEFAULT_PREFIX_DIR}/site-packages" -DEFAULT_MINI_BINARY = f"{DEFAULT_PREFIX_DIR}/bin/mini" -DEFAULT_LOG_DIR = "/logs/agent" -MINI_SWE_AGENT_DEFAULT_AGENT_WORKDIR = "${AGENT_WORKDIR:-/app}" -MINI_SWE_AGENT_DEFAULT_INSTRUCTION_PATH = "/mini-swe-agent/prompt.txt" -MINI_SWE_AGENT_DEFAULT_SYSTEM_PROMPT_PATH = "/mini-swe-agent/system.txt" +from .command import CommandHarness, CommandHarnessConfig, shell_command + +MINI_SWE_AGENT_DEFAULT_WORKDIR = "/app" MINI_SWE_AGENT_DEFAULT_LOG_PATH = "/logs/agent/mini-swe-agent.log" -MINI_SWE_AGENT_DEFAULT_TRAJECTORY_PATH = "/logs/agent/mini-swe-agent.traj.json" +MINI_SWE_AGENT_DEFAULT_OUTPUT_PATH = "/logs/agent/mini-swe-agent.traj.json" MINI_SWE_AGENT_DEFAULT_VERSION = "mini-swe-agent@2.2.8" -MINI_SWE_AGENT_DEFAULT_PACKAGE = MINI_SWE_AGENT_DEFAULT_VERSION MINI_SWE_AGENT_DEFAULT_CONFIG_SPEC = "mini" MINI_SWE_AGENT_DEFAULT_MODEL_CLASS = "litellm" MINI_SWE_AGENT_DEFAULT_ENVIRONMENT_TIMEOUT = 120 -def build_mini_swe_agent_install_script( - version: str = MINI_SWE_AGENT_DEFAULT_VERSION, - prefix_dir: str = DEFAULT_PREFIX_DIR, - package: str | None = None, -) -> str: - root = shlex.quote(str(PurePosixPath(prefix_dir).parent)) - prefix = shlex.quote(prefix_dir) - site_packages = shlex.quote(f"{prefix_dir.rstrip('/')}/site-packages") - if package is not None: - version = package - name, parsed_version = split_versioned_agent_spec(version) - requirement = name - if parsed_version and parsed_version != "latest": - requirement = f"{name}=={parsed_version}" - return f"""\ -set -e -{python_runtime_setup_command()} -rm -rf {prefix} -mkdir -p {root} {prefix}/bin {site_packages} {shlex.quote(DEFAULT_LOG_DIR)} /mini-swe-agent -vf_python_install --target {site_packages} {shlex.quote(requirement)} -echo "$VF_PYTHON" > {prefix}/python -cat > {prefix}/bin/mini <<'EOF' -#!/usr/bin/env sh -export PYTHONPATH={site_packages}:${{PYTHONPATH:-}} -exec "$(cat {prefix}/python)" -m minisweagent.run.mini "$@" -EOF -chmod +x {prefix}/bin/mini -test -x {prefix}/bin/mini -""" - - -class MiniSWEAgentProgramConfig(vf.ProgramConfig): - agent_workdir: str = MINI_SWE_AGENT_DEFAULT_AGENT_WORKDIR - instruction_path: str = MINI_SWE_AGENT_DEFAULT_INSTRUCTION_PATH - system_prompt_path: str = MINI_SWE_AGENT_DEFAULT_SYSTEM_PROMPT_PATH +class MiniSWEAgentConfig(CommandHarnessConfig): + version: str = MINI_SWE_AGENT_DEFAULT_VERSION + cwd: str | None = MINI_SWE_AGENT_DEFAULT_WORKDIR log_path: str = MINI_SWE_AGENT_DEFAULT_LOG_PATH - trajectory_path: str = MINI_SWE_AGENT_DEFAULT_TRAJECTORY_PATH + output_path: str = MINI_SWE_AGENT_DEFAULT_OUTPUT_PATH config_spec: str = MINI_SWE_AGENT_DEFAULT_CONFIG_SPEC model_class: str = MINI_SWE_AGENT_DEFAULT_MODEL_CLASS environment_timeout: int = MINI_SWE_AGENT_DEFAULT_ENVIRONMENT_TIMEOUT parallel_tool_calls: bool = True - extra_config_specs: list[str] | None = None - sandbox: vf.SandboxConfig | None = vf.SandboxConfig() + extra_config_specs: list[str] = Field(default_factory=list) + max_turns: int = 4 - def resolve( - self, version: str = MINI_SWE_AGENT_DEFAULT_VERSION - ) -> vf.ProgramConfig: - files: dict[str, vf.ProgramValue] = { - self.instruction_path: {"fn": "verifiers.v1.utils.prompt_utils:task_text"}, - self.system_prompt_path: { - "fn": "verifiers.v1.utils.prompt_utils:state_system_prompt_text" - }, - } - artifacts = vf.ArtifactsConfig.model_validate( - { - "mini_swe_agent_log": { - "path": self.log_path, - "format": "text", - "optional": True, - }, - "mini_swe_agent_trajectory": { - "path": self.trajectory_path, - "format": "json", - "optional": True, - }, - } + +class MiniSWEAgent(CommandHarness[MiniSWEAgentConfig]): + config: MiniSWEAgentConfig + + def command(self, task: vf.Task, state: vf.State) -> list[str]: + _ = state + instruction = str( + getattr(task, "instruction", None) or getattr(task, "question", None) or "" ) - if self.agent_workdir == MINI_SWE_AGENT_DEFAULT_AGENT_WORKDIR: - workdir_line = ( - f"MINI_SWE_AGENT_WORKDIR={MINI_SWE_AGENT_DEFAULT_AGENT_WORKDIR}" + if not instruction: + instruction = "\n\n".join( + str(getattr(message, "content", "") or "") for message in task.prompt ) - else: - workdir_line = f"MINI_SWE_AGENT_WORKDIR={shlex.quote(self.agent_workdir)}" - config_args = [ "-c", - shlex.quote(self.config_spec), + self.config.config_spec, "-c", "agent.cost_limit=0", "-c", - f"environment.timeout={self.environment_timeout}", + f"environment.timeout={self.config.environment_timeout}", "-c", - f"model.model_class={shlex.quote(self.model_class)}", + f"model.model_class={self.config.model_class}", "-c", "model.cost_tracking=ignore_errors", "-c", "model.model_kwargs.custom_llm_provider=openai", "-c", - f"model.model_kwargs.parallel_tool_calls={str(self.parallel_tool_calls).lower()}", + f"model.model_kwargs.parallel_tool_calls={str(self.config.parallel_tool_calls).lower()}", ] - for spec in self.extra_config_specs or []: - config_args.extend(["-c", shlex.quote(spec)]) - - setup = build_mini_swe_agent_install_script( - version=version, - ) - log_dir = str(PurePosixPath(self.log_path).parent) - trajectory_dir = str(PurePosixPath(self.trajectory_path).parent) - system_prompt_file = shlex.quote(self.system_prompt_path) - script = f"""\ + for spec in self.config.extra_config_specs: + config_args.extend(["-c", spec]) + args = " ".join(shlex.quote(arg) for arg in config_args) + script = f""" set -eo pipefail -export PATH={shlex.quote(DEFAULT_PREFIX_DIR)}/bin:"$PATH" -export PYTHONPATH={shlex.quote(DEFAULT_SITE_PACKAGES_DIR)}:"${{PYTHONPATH:-}}" -export MSWEA_CONFIGURED=true -export MSWEA_SILENT_STARTUP=true -export MSWEA_GLOBAL_CONFIG_DIR=/tmp/mini-swe-agent-config -export OPENAI_API_KEY="${{OPENAI_API_KEY:-intercepted}}" - -{workdir_line} -mkdir -p {shlex.quote(log_dir)} {shlex.quote(trajectory_dir)} "$MINI_SWE_AGENT_WORKDIR" "$MSWEA_GLOBAL_CONFIG_DIR" - -MINI_SWE_AGENT_TASK="$(cat {shlex.quote(self.instruction_path)})" -CONFIG_ARGS=({" ".join(config_args)}) -CONFIG_ARGS+=(-c "environment.cwd=$MINI_SWE_AGENT_WORKDIR") -if [ -s {system_prompt_file} ]; then - CONFIG_ARGS+=(-c "agent.system_template=$(cat {system_prompt_file})") -fi -cd "$MINI_SWE_AGENT_WORKDIR" -timeout --kill-after=30s "${{AGENT_TIMEOUT_SECONDS:-3600}}" {shlex.quote(DEFAULT_MINI_BINARY)} \\ - --model "$OPENAI_MODEL" \\ - --task "$MINI_SWE_AGENT_TASK" \\ - --output {shlex.quote(self.trajectory_path)} \\ - --exit-immediately \\ - --yolo \\ - "${{CONFIG_ARGS[@]}}" 2>&1 | tee -a {shlex.quote(self.log_path)} +mkdir -p "$(dirname {self.config.log_path!r})" "$(dirname {self.config.output_path!r})" +mini --model "$OPENAI_MODEL" --task {instruction!r} --output {self.config.output_path!r} \ + --exit-immediately --yolo {args} 2>&1 | tee -a {self.config.log_path!r} """ - return self.resolve_command( - command=["bash", "-lc", script], - default_sandbox=self.sandbox, - files=files, - setup=setup, - env={"OPENAI_MODEL": "runtime.model"}, - artifacts=artifacts, - ) - - -class MiniSWEAgentConfig(vf.HarnessConfig): - version: str = MINI_SWE_AGENT_DEFAULT_VERSION - program: MiniSWEAgentProgramConfig = MiniSWEAgentProgramConfig() - max_turns: int = 4 - - -class MiniSWEAgent(vf.Harness[MiniSWEAgentConfig]): - config: MiniSWEAgentConfig - - def load_program_config(self, config: MiniSWEAgentConfig) -> vf.ProgramConfig: - return config.program.resolve(version=config.version) + return shell_command(script) def load_harness(config: MiniSWEAgentConfig) -> MiniSWEAgent: diff --git a/packages/harnesses/harnesses/nemo_gym.py b/packages/harnesses/harnesses/nemo_gym.py index 39d35332d4..47150d8560 100644 --- a/packages/harnesses/harnesses/nemo_gym.py +++ b/packages/harnesses/harnesses/nemo_gym.py @@ -1,83 +1,67 @@ -import asyncio -import contextlib -import inspect import json -import logging -import os -import secrets -from collections.abc import Awaitable, Iterator, Sequence -from copy import deepcopy from pathlib import Path -from typing import Protocol, TypeAlias, cast -from urllib.parse import urlparse -import verifiers as vf -from aiohttp import ClientSession, web from pydantic import Field -from verifiers.types import AssistantMessage, ToolCall, ToolMessage -from verifiers.utils.serve_utils import get_free_port +import verifiers.v1 as vf -logger = logging.getLogger(__name__) +from .command import CommandHarness, CommandHarnessConfig, shell_command -NEMO_GYM_POLICY_MODEL_SERVER_NAME = "policy_model" -NEMO_GYM_POLICY_MODEL_TYPE_NAME = "verifiers_proxy" -NEMO_GYM_EXTERNAL_POLICY_MODEL_ENTRYPOINT = "__verifiers_external_policy_model__.py" -PROXY_MODEL_NAME = "verifiers-nemo-gym-proxy" -_NEMO_GYM_GLOBALS_LOCK = asyncio.Lock() -_NEMO_GYM_ACTIVE_RUNNERS = 0 -_NEMO_GYM_OWNS_AIOHTTP_CLIENT = False -_RAY_ENABLE_UV_RUN_RUNTIME_ENV = "RAY_ENABLE_UV_RUN_RUNTIME_ENV" -_NEMO_GYM_CONFIG_PATH_ENV_VAR_NAME = "NEMO_GYM_CONFIG_PATH" -EndpointConfig = dict[str, str] -ConfigData: TypeAlias = dict[str, object] -ConfigMap: TypeAlias = dict[str, object] -TaskRow: TypeAlias = dict[str, object] +DEFAULT_NEMO_GYM_COMMAND = "python -m nemo_gym.cli run-one" +DEFAULT_NEMO_GYM_DATA_NAME = "example.jsonl" -class NeMoGymHarnessConfig(vf.HarnessConfig): +class NeMoGymHarnessConfig(CommandHarnessConfig): + command: list[str] = Field(default_factory=list) nemo_env: str | None = None config_name: str | None = None config_paths: list[str] = Field(default_factory=list) server_name: str | None = None agent_name: str | None = None timeout_seconds: float | None = None - global_config: ConfigMap = Field(default_factory=dict) + global_config: vf.JsonData = Field(default_factory=dict) -class NeMoGymRunner(Protocol): - async def run( - self, - row: TaskRow, - *, - config_paths: Sequence[str], - server_name: str | None, - agent_name: str | None, - endpoint_config: EndpointConfig, - timeout_seconds: float | None, - global_config: ConfigMap, - ) -> ConfigMap: ... - - -class NeMoGymServerClient(Protocol): - global_config_dict: ConfigMap - head_server_config: object - - -class NeMoGymRunHelper(Protocol): - _server_client: NeMoGymServerClient - - def start(self, parser_config: object) -> None: ... - - def shutdown(self) -> None: ... - - def poll(self) -> None: ... - +class NeMoGymHarness(CommandHarness[NeMoGymHarnessConfig]): + config: NeMoGymHarnessConfig -class NeMoGymRolloutCollector(Protocol): - def run_examples( - self, rows: list[ConfigData], *, head_server_config: object - ) -> Iterator[Awaitable[tuple[ConfigData, ConfigMap]]]: ... + def command(self, task: vf.Task, state: vf.State) -> list[str]: + _ = state + if self.config.command: + return list(self.config.command) + row = getattr(task, "nemo_gym_row", None) + if not isinstance(row, dict): + raise ValueError("NeMoGymHarness tasks must contain nemo_gym_row.") + payload = json.dumps( + { + "row": row, + "config_paths": self.config_paths(), + "server_name": self.config.server_name, + "agent_name": self.config.agent_name, + "global_config": self.config.global_config, + }, + ensure_ascii=False, + ) + script = f""" +set -eo pipefail +export NEMO_GYM_ROW_JSON={payload!r} +{DEFAULT_NEMO_GYM_COMMAND} +""" + return shell_command(script) + + def config_paths(self) -> list[str]: + if self.config.config_paths: + return list(self.config.config_paths) + if self.config.nemo_env is None: + raise ValueError("NeMoGymHarness requires config_paths or nemo_env.") + return [ + str( + resolve_nemo_gym_config_path( + self.config.nemo_env, + self.config.config_name, + ) + ) + ] def nemo_gym_package_root() -> Path: @@ -118,895 +102,5 @@ def resolve_nemo_gym_config_path( ) -def infer_nemo_gym_agent_from_config(config_path: str | Path) -> tuple[str, str]: - try: - from omegaconf import OmegaConf # ty: ignore[unresolved-import] - except ImportError as exc: - raise ImportError( - "NeMo Gym config inference requires omegaconf. " - "Install as `verifiers[nemogym]`." - ) from exc - - path = Path(config_path) - raw_config = OmegaConf.to_container(OmegaConf.load(path), resolve=False) - if not isinstance(raw_config, dict): - raise ValueError(f"NeMo Gym config must be a mapping: {path}") - for top_level_name, top_level_value in raw_config.items(): - if not isinstance(top_level_value, dict): - continue - agents = top_level_value.get("responses_api_agents") - if not isinstance(agents, dict) or not agents: - continue - agent_name = next(iter(agents)) - return str(top_level_name), str(agent_name) - raise ValueError(f"No responses_api_agents entry found in {path}") - - -def first_nemo_gym_agent( - config_paths: list[str] | tuple[str, ...], -) -> tuple[str, str] | None: - for config_path in config_paths: - try: - return infer_nemo_gym_agent_from_config(config_path) - except FileNotFoundError: - continue - except ValueError: - continue - return None - - -def agent_ref_name(value: object) -> str | None: - if not isinstance(value, dict): - return None - name = cast(ConfigMap, value).get("name") - return name if isinstance(name, str) and name else None - - -class NeMoGymHarness(vf.Harness[NeMoGymHarnessConfig]): - """Run a NeMo Gym row from a Verifiers rollout. - - The default runner keeps one NeMo Gym server stack alive per harness instance. - NeMo Gym agents call a stable local Verifiers proxy registered as their - policy model server, and that proxy routes each model request into the - matching Verifiers rollout endpoint. - """ - - config_type = NeMoGymHarnessConfig - config: NeMoGymHarnessConfig - runner: NeMoGymRunner | None = None - - def compile_program(self, program: vf.ProgramConfig) -> "vf.ProgramRunner": - configure_nemo_gym_harness_config(self.config) - self.runner = PersistentNeMoGymRunner() - return self._run_nemo_gym - - async def teardown(self) -> None: - runner = self.runner - self.runner = None - teardown = getattr(runner, "teardown", None) - if callable(teardown): - result = teardown() - if inspect.isawaitable(result): - await result - await super().teardown() - - async def _run_nemo_gym(self, task: vf.Task, state: vf.State) -> vf.State: - runner = self.runner - if runner is None: - raise RuntimeError("NeMo Gym harness runner has not been compiled.") - endpoint_config = await nemo_gym_rollout_endpoint_config(state) - row = nemo_gym_row_from_task(cast(ConfigMap, task), self.config.agent_name) - result = await runner.run( - row, - config_paths=self._config_paths(), - server_name=self.config.server_name, - agent_name=self.config.agent_name, - endpoint_config=endpoint_config, - timeout_seconds=self.config.timeout_seconds, - global_config=self.config.global_config, - ) - apply_nemo_gym_result(state, result) - return state - - def _config_paths(self) -> list[str]: - paths = list(self.config.config_paths) - if not paths: - raise ValueError("NeMoGymHarness requires at least one config path.") - return paths - - -def configure_nemo_gym_harness_config(config: NeMoGymHarnessConfig) -> None: - if config.nemo_env and not config.config_paths: - config.config_paths = [ - str(resolve_nemo_gym_config_path(config.nemo_env, config.config_name)) - ] - - if config.server_name is not None and config.agent_name is not None: - return - inferred = first_nemo_gym_agent(tuple(config.config_paths)) - if inferred is None: - return - inferred_server_name, inferred_agent_name = inferred - if config.server_name is None: - config.server_name = inferred_server_name - if config.agent_name is None: - config.agent_name = inferred_agent_name - - -async def nemo_gym_rollout_endpoint_config(state: vf.State) -> EndpointConfig: - endpoint_config = state.get_endpoint_config(api="responses") - client = state.get_client(api="responses") - api_key = getattr(client, "api_key", None) - close = getattr(client, "close", None) - if callable(close): - result = close() - if inspect.isawaitable(result): - await result - if api_key is None: - raise ValueError("NeMo Gym rollout endpoint requires a model API key.") - if not isinstance(api_key, str): - api_key = str(api_key) - return { - "base_url": endpoint_config.base_url, - "api_key": api_key, - "model": endpoint_config.model, - } - - -class PersistentNeMoGymRunner: - """Run one NeMo Gym server stack and route model calls per rollout.""" - - def __init__(self) -> None: - self._lifecycle_lock = asyncio.Lock() - self._helper: NeMoGymRunHelper | None = None - self._rollout_collector: NeMoGymRolloutCollector | None = None - self._proxy: NeMoGymModelProxy | None = None - self._config_key: str | None = None - self._head_server_config: object | None = None - - async def run( - self, - row: TaskRow, - *, - config_paths: Sequence[str], - server_name: str | None, - agent_name: str | None, - endpoint_config: EndpointConfig, - timeout_seconds: float | None, - global_config: ConfigMap, - ) -> ConfigMap: - async with self._lifecycle_lock: - await self._ensure_started( - config_paths=config_paths, - global_config=global_config, - ) - - run_once = self._run_once( - row, - server_name=server_name, - agent_name=agent_name, - endpoint_config=endpoint_config, - ) - if timeout_seconds is None: - return await run_once - return await asyncio.wait_for(run_once, timeout=timeout_seconds) - - async def _ensure_started( - self, *, config_paths: Sequence[str], global_config: ConfigMap - ) -> None: - global _NEMO_GYM_ACTIVE_RUNNERS, _NEMO_GYM_OWNS_AIOHTTP_CLIENT - - key = json.dumps( - { - "config_paths": list(config_paths), - "global_config": jsonable(dict(global_config)), - }, - sort_keys=True, - ) - if self._helper is not None and self._config_key == key: - return - if self._helper is not None: - await self.teardown() - - try: - from omegaconf import OmegaConf # ty: ignore[unresolved-import] - - from nemo_gym import cli as nemo_cli # ty: ignore[unresolved-import] - from nemo_gym.global_config import ( # ty: ignore[unresolved-import] - GlobalConfigDictParserConfig, - ) - from nemo_gym import ( # ty: ignore[unresolved-import] - global_config as nemo_global_config, - ) - from nemo_gym import ( # ty: ignore[unresolved-import] - server_utils as nemo_server_utils, - ) - from nemo_gym.rollout_collection import ( # ty: ignore[unresolved-import] - RolloutCollectionHelper, - ) - from nemo_gym.server_utils import ( # ty: ignore[unresolved-import] - GlobalAIOHTTPAsyncClientConfig, - is_global_aiohttp_client_setup, - set_global_aiohttp_client, - ) - except ImportError as exc: - raise ImportError( - "NeMoGymHarness requires nemo-gym. Install as `verifiers[nemogym]`." - ) from exc - - proxy = await self._ensure_proxy() - config = build_nemo_gym_global_config( - config_paths=config_paths, - endpoint_config=proxy.endpoint_config(), - global_config=global_config, - ) - parser_config = GlobalConfigDictParserConfig( - initial_global_config_dict=OmegaConf.create(config), - skip_load_from_cli=True, - skip_load_from_dotenv=True, - ) - - helper = cast(NeMoGymRunHelper, nemo_cli.RunHelper()) - async with _NEMO_GYM_GLOBALS_LOCK: - reset_nemo_gym_global_config(nemo_global_config) - try: - with disable_ray_uv_run_runtime_env(): - with skip_nemo_gym_policy_model_process(nemo_cli): - await asyncio.to_thread(helper.start, parser_config) - if not is_global_aiohttp_client_setup(): - set_global_aiohttp_client( - GlobalAIOHTTPAsyncClientConfig.model_validate( - helper._server_client.global_config_dict - ) - ) - _NEMO_GYM_OWNS_AIOHTTP_CLIENT = True - _NEMO_GYM_ACTIVE_RUNNERS += 1 - except Exception: - with contextlib.suppress(Exception): - await asyncio.to_thread(helper.shutdown) - if _NEMO_GYM_ACTIVE_RUNNERS == 0 and _NEMO_GYM_OWNS_AIOHTTP_CLIENT: - await close_nemo_gym_aiohttp_client(nemo_server_utils) - _NEMO_GYM_OWNS_AIOHTTP_CLIENT = False - reset_nemo_gym_global_config(nemo_global_config) - raise - - self._helper = helper - self._rollout_collector = cast( - NeMoGymRolloutCollector, RolloutCollectionHelper() - ) - self._config_key = key - self._head_server_config = helper._server_client.head_server_config - - async def _ensure_proxy(self) -> "NeMoGymModelProxy": - if self._proxy is None: - self._proxy = NeMoGymModelProxy(require_auth=False) - await self._proxy.start() - return self._proxy - - async def _run_once( - self, - row: TaskRow, - *, - server_name: str | None, - agent_name: str | None, - endpoint_config: EndpointConfig, - ) -> ConfigMap: - if ( - self._helper is None - or self._rollout_collector is None - or self._proxy is None - or self._head_server_config is None - ): - raise RuntimeError("NeMo Gym runner has not been started.") - - self._helper.poll() - rollout_id = secrets.token_urlsafe(16) - routing_model = nemo_gym_proxy_model_name(rollout_id) - async with self._proxy.activate(routing_model, endpoint_config): - request_row = prepare_nemo_gym_rollout_collection_row( - row, - server_name=server_name, - agent_name=agent_name, - ) - set_nemo_gym_proxy_model(request_row, routing_model) - futures = self._rollout_collector.run_examples( - [request_row], - head_server_config=self._head_server_config, - ) - try: - future = next(futures) - except StopIteration as exc: - raise RuntimeError( - "NeMo Gym rollout collector returned no tasks." - ) from exc - _, result = await future - return cast(ConfigMap, result) - - async def teardown(self) -> None: - global _NEMO_GYM_ACTIVE_RUNNERS, _NEMO_GYM_OWNS_AIOHTTP_CLIENT - - helper = self._helper - proxy = self._proxy - self._helper = None - self._rollout_collector = None - self._proxy = None - self._config_key = None - self._head_server_config = None - - if helper is None: - if proxy is not None: - await proxy.stop() - return - - nemo_global_config: object | None - nemo_server_utils: object | None - try: - from nemo_gym import ( # ty: ignore[unresolved-import] - global_config as imported_nemo_global_config, - ) - from nemo_gym import ( # ty: ignore[unresolved-import] - server_utils as imported_nemo_server_utils, - ) - except ImportError: - nemo_global_config = None - nemo_server_utils = None - else: - nemo_global_config = imported_nemo_global_config - nemo_server_utils = imported_nemo_server_utils - - try: - await asyncio.to_thread(helper.shutdown) - except Exception: - logger.exception("Failed to shut down NeMo Gym runner") - if nemo_global_config is not None: - async with _NEMO_GYM_GLOBALS_LOCK: - _NEMO_GYM_ACTIVE_RUNNERS = max(0, _NEMO_GYM_ACTIVE_RUNNERS - 1) - if _NEMO_GYM_ACTIVE_RUNNERS == 0: - if nemo_server_utils is not None and _NEMO_GYM_OWNS_AIOHTTP_CLIENT: - await close_nemo_gym_aiohttp_client(nemo_server_utils) - _NEMO_GYM_OWNS_AIOHTTP_CLIENT = False - reset_nemo_gym_global_config(nemo_global_config) - if proxy is not None: - await proxy.stop() - - -class NeMoGymModelProxy: - """Stable OpenAI-compatible proxy used by a persistent NeMo Gym stack.""" - - def __init__( - self, - *, - host: str = "127.0.0.1", - port: int | None = None, - require_auth: bool = True, - ) -> None: - self.host = host - self.port = get_free_port() if port is None else port - self.require_auth = require_auth - self.secret = secrets.token_urlsafe(32) - self._app: web.Application | None = None - self._runner: web.AppRunner | None = None - self._site: web.TCPSite | None = None - self._session: ClientSession | None = None - self._lock = asyncio.Lock() - self._endpoints_by_model: dict[str, dict[str, str]] = {} - - async def start(self) -> None: - async with self._lock: - if self._runner is not None: - return - app = web.Application() - app.router.add_get("/health", lambda _: web.json_response({"status": "ok"})) - app.router.add_get("/v1/models", self._handle_models) - app.router.add_route("*", "/v1/{tail:.*}", self._handle_openai_request) - runner = web.AppRunner(app) - await runner.setup() - site = web.TCPSite(runner, self.host, self.port) - await site.start() - self._app = app - self._runner = runner - self._site = site - self._session = ClientSession() - - async def stop(self) -> None: - async with self._lock: - session = self._session - runner = self._runner - self._session = None - self._runner = None - self._site = None - self._app = None - self._endpoints_by_model.clear() - if session is not None: - await session.close() - if runner is not None: - await runner.cleanup() - - def endpoint_config(self) -> EndpointConfig: - return { - "base_url": f"http://{self.host}:{self.port}/v1", - "api_base": f"http://{self.host}:{self.port}/v1", - "api_key": self.secret, - "model": PROXY_MODEL_NAME, - "api_client_type": "openai_responses", - } - - @contextlib.asynccontextmanager - async def activate(self, routing_model: str, endpoint_config: EndpointConfig): - async with self._lock: - if routing_model in self._endpoints_by_model: - raise RuntimeError( - f"NeMo Gym model proxy already has model route {routing_model!r}." - ) - self._endpoints_by_model[routing_model] = { - "base_url": str(endpoint_config["base_url"]), - "api_key": str(endpoint_config["api_key"]), - "model": str(endpoint_config["model"]), - } - try: - yield - finally: - async with self._lock: - self._endpoints_by_model.pop(routing_model, None) - - async def _handle_models(self, request: web.Request) -> web.Response: - if not self._authorized(request): - return web.json_response({"error": "Unauthorized"}, status=401) - async with self._lock: - models = list(self._endpoints_by_model) or [PROXY_MODEL_NAME] - return web.json_response( - { - "object": "list", - "data": [ - { - "id": model, - "object": "model", - "created": 0, - "owned_by": "verifiers", - } - for model in models - ], - } - ) - - async def _handle_openai_request(self, request: web.Request) -> web.Response: - if not self._authorized(request): - return web.json_response({"error": "Unauthorized"}, status=401) - json_payload = await self._request_json(request) - endpoint = await self._endpoint_for_request(json_payload) - if endpoint is None: - return web.json_response( - { - "error": ( - "Missing or unknown model for NeMo Gym model proxy request." - ) - }, - status=409, - ) - session = self._session - if session is None: - return web.json_response({"error": "Proxy not started"}, status=503) - - suffix = request.path.removeprefix("/v1") - upstream_url = f"{endpoint['base_url'].rstrip('/')}{suffix}" - if request.query_string: - upstream_url = f"{upstream_url}?{request.query_string}" - - headers = self._upstream_headers(request, endpoint["api_key"]) - if json_payload is not None: - json_payload["model"] = endpoint["model"] - request_kwargs: ConfigData = {"json": json_payload} - else: - request_kwargs = {"data": await request.read()} - - async with session.request( - request.method, - upstream_url, - headers=headers, - **request_kwargs, - ) as response: - body = await response.read() - return web.Response( - status=response.status, - body=body, - headers=self._response_headers(response.headers.items()), - ) - - async def _endpoint_for_request( - self, json_payload: ConfigMap | None - ) -> EndpointConfig | None: - routing_model = json_payload.get("model") if json_payload is not None else None - async with self._lock: - if isinstance(routing_model, str): - endpoint = self._endpoints_by_model.get(routing_model) - if endpoint: - return dict(endpoint) - if len(self._endpoints_by_model) == 1: - endpoint = next(iter(self._endpoints_by_model.values())) - return dict(endpoint) - return None - - def _authorized(self, request: web.Request) -> bool: - if not self.require_auth: - return True - auth = request.headers.get("Authorization", "") - api_key = request.headers.get("x-api-key", "") - return auth == f"Bearer {self.secret}" or api_key == self.secret - - def _upstream_headers(self, request: web.Request, api_key: str) -> dict[str, str]: - headers = { - key: value - for key, value in request.headers.items() - if key.lower() - not in { - "authorization", - "content-length", - "host", - "x-api-key", - } - } - headers["Authorization"] = f"Bearer {api_key}" - return headers - - async def _request_json(self, request: web.Request) -> ConfigData | None: - content_type = request.headers.get("content-type", "") - if "json" not in content_type: - return None - payload = await request.json() - if not isinstance(payload, dict): - raise TypeError("OpenAI-compatible proxy requests must be JSON objects.") - return cast(ConfigData, payload) - - def _response_headers(self, headers: object) -> dict[str, str]: - header_items = cast(list[tuple[str, object]], headers) - return { - key: str(value) - for key, value in header_items - if key.lower() - not in { - "content-length", - "content-encoding", - "connection", - "keep-alive", - "transfer-encoding", - } - } - - -def build_nemo_gym_global_config( - *, - config_paths: Sequence[str], - endpoint_config: EndpointConfig, - global_config: ConfigMap, -) -> ConfigData: - config = dict(global_config) - config["config_paths"] = list(config_paths) - config["policy_base_url"] = str(endpoint_config["base_url"]) - config["policy_api_key"] = str(endpoint_config["api_key"]) - config["policy_model_name"] = str(endpoint_config["model"]) - config[NEMO_GYM_POLICY_MODEL_SERVER_NAME] = build_nemo_gym_policy_model_config( - endpoint_config - ) - config.setdefault("head_server", {"host": "127.0.0.1", "port": get_free_port()}) - return config - - -def build_nemo_gym_policy_model_config( - endpoint_config: EndpointConfig, -) -> ConfigData: - parsed_url = urlparse(str(endpoint_config["base_url"])) - if not parsed_url.hostname or parsed_url.port is None: - raise ValueError( - "NeMo Gym Verifiers proxy base_url must include an explicit host and port." - ) - return { - "responses_api_models": { - NEMO_GYM_POLICY_MODEL_TYPE_NAME: { - "entrypoint": NEMO_GYM_EXTERNAL_POLICY_MODEL_ENTRYPOINT, - "host": parsed_url.hostname, - "port": parsed_url.port, - } - } - } - - -def nemo_gym_proxy_model_name(rollout_id: str) -> str: - return f"{PROXY_MODEL_NAME}-{rollout_id}" - - -class NoopPolicyModelProcess: - pid = 0 - - def __init__(self) -> None: - self._running = True - - def poll(self) -> int | None: - return None if self._running else 0 - - def communicate(self): - self._running = False - return b"", b"" - - def send_signal(self, signal: int) -> None: - self._running = False - - def wait(self, timeout: float | None = None) -> int: - self._running = False - return 0 - - def kill(self) -> None: - self._running = False - - -@contextlib.contextmanager -def skip_nemo_gym_policy_model_process(nemo_cli_module: object): - original_run_command = getattr(nemo_cli_module, "run_command") - original_setup_env_command = getattr(nemo_cli_module, "setup_env_command") - - def setup_env_command( - dir_path: object, global_config_dict: object, prefix: str - ) -> str: - if prefix == NEMO_GYM_POLICY_MODEL_SERVER_NAME: - return "true" - return original_setup_env_command(dir_path, global_config_dict, prefix) - - def run_command(command: str, working_dir_path: object): - if ( - f"{_NEMO_GYM_CONFIG_PATH_ENV_VAR_NAME}={NEMO_GYM_POLICY_MODEL_SERVER_NAME}" - ) in command: - return NoopPolicyModelProcess() - return original_run_command(command, working_dir_path) - - setattr(nemo_cli_module, "setup_env_command", setup_env_command) - setattr(nemo_cli_module, "run_command", run_command) - try: - yield - finally: - setattr(nemo_cli_module, "run_command", original_run_command) - setattr(nemo_cli_module, "setup_env_command", original_setup_env_command) - - -def reset_nemo_gym_global_config(module: object) -> None: - if hasattr(module, "_GLOBAL_CONFIG_DICT"): - setattr(module, "_GLOBAL_CONFIG_DICT", None) - - -async def close_nemo_gym_aiohttp_client(module: object) -> None: - client = getattr(module, "_GLOBAL_AIOHTTP_CLIENT", None) - if client is None: - return - close = getattr(client, "close", None) - if callable(close): - result = close() - if inspect.isawaitable(result): - await result - setattr(module, "_GLOBAL_AIOHTTP_CLIENT", None) - - -@contextlib.contextmanager -def disable_ray_uv_run_runtime_env(): - original_env = os.environ.get(_RAY_ENABLE_UV_RUN_RUNTIME_ENV) - os.environ[_RAY_ENABLE_UV_RUN_RUNTIME_ENV] = "0" - ray_constants: object | None = None - original_constant: object | None = None - try: - import ray._private.ray_constants as ray_constants # ty: ignore[unresolved-import] - - original_constant = getattr(ray_constants, _RAY_ENABLE_UV_RUN_RUNTIME_ENV) - setattr(ray_constants, _RAY_ENABLE_UV_RUN_RUNTIME_ENV, False) - except (AttributeError, ImportError): - pass - try: - yield - finally: - if original_env is None: - os.environ.pop(_RAY_ENABLE_UV_RUN_RUNTIME_ENV, None) - else: - os.environ[_RAY_ENABLE_UV_RUN_RUNTIME_ENV] = original_env - if ray_constants is not None and original_constant is not None: - setattr(ray_constants, _RAY_ENABLE_UV_RUN_RUNTIME_ENV, original_constant) - - -def prepare_nemo_gym_request_row( - row: TaskRow, - agent_name: str | None, -) -> ConfigData: - prepared: ConfigData = deepcopy(dict(row)) - if agent_name and not agent_ref_name(prepared.get("agent_ref")): - prepared["agent_ref"] = { - "type": "responses_api_agents", - "name": agent_name, - } - if not agent_ref_name(prepared.get("agent_ref")): - raise ValueError( - "NeMo Gym row has no agent_ref.name; pass NeMoGymHarness(agent_name=...) " - "or include agent_ref in each row." - ) - return prepared - - -def prepare_nemo_gym_rollout_collection_row( - row: TaskRow, - *, - server_name: str | None, - agent_name: str | None, -) -> ConfigData: - prepared = prepare_nemo_gym_request_row(row, agent_name) - target_server_name = nemo_gym_server_name(prepared, server_name) - agent_ref = dict(cast(ConfigMap, prepared["agent_ref"])) - agent_ref["name"] = target_server_name - prepared["agent_ref"] = agent_ref - return prepared - - -def set_nemo_gym_proxy_model(row: ConfigData, routing_model: str) -> None: - create_params = row.get("responses_create_params") - if not isinstance(create_params, dict): - raise ValueError("NeMo Gym row requires responses_create_params.") - params = dict(cast(ConfigMap, create_params)) - params["model"] = routing_model - row["responses_create_params"] = params - - -def nemo_gym_server_name(row: TaskRow, server_name: str | None) -> str: - if server_name: - return server_name - name = agent_ref_name(row.get("agent_ref")) - if name is None: - raise ValueError( - "NeMo Gym row has no agent_ref.name; pass " - "NeMoGymHarness(server_name=...) or include agent_ref in each row." - ) - return name - - -def nemo_gym_row_from_task(task: ConfigMap, agent_name: str | None) -> ConfigData: - raw_row = task.get("nemo_gym_row") - if isinstance(raw_row, dict): - return prepare_nemo_gym_request_row(cast(TaskRow, raw_row), agent_name) - if "responses_create_params" not in task: - raise ValueError( - "NeMoGymHarness tasks must contain nemo_gym_row or responses_create_params." - ) - ignored = { - "prompt", - "info", - "example_id", - "task_id", - "taskset_id", - "runtime", - } - row = {key: deepcopy(value) for key, value in task.items() if key not in ignored} - return prepare_nemo_gym_request_row(row, agent_name) - - -def apply_nemo_gym_result(state: vf.State, result: ConfigMap) -> None: - result_dict = jsonable_mapping(result) - state["nemo_gym_result"] = result_dict - response = result_dict.get("response") - if isinstance(response, dict): - completion = messages_from_nemo_gym_response(cast(ConfigMap, response)) - if completion: - state["completion"] = completion - reward = result_dict.get("reward") - if isinstance(reward, bool): - raise TypeError("NeMo Gym reward must be numeric.") - if isinstance(reward, int | float): - state["reward"] = float(reward) - elif isinstance(reward, str): - try: - state["reward"] = float(reward) - except ValueError as exc: - raise TypeError("NeMo Gym reward must be numeric.") from exc - elif reward is not None: - raise TypeError("NeMo Gym reward must be numeric.") - metrics = state.setdefault("metrics", {}) - if isinstance(metrics, dict): - for key, value in result_dict.items(): - if key in {"responses_create_params", "response", "reward"}: - continue - if isinstance(value, bool): - metrics[key] = float(value) - elif isinstance(value, int | float): - metrics[key] = float(value) - state.stop("nemo_gym_completed") - - -def messages_from_nemo_gym_response( - response: ConfigMap, -) -> list[ConfigData]: - output = response.get("output") - if not isinstance(output, list): - return [] - messages: list[ConfigData] = [] - for item in output: - if not isinstance(item, dict): - continue - item = cast(ConfigMap, item) - item_type = item.get("type") - if item_type == "function_call": - call_id = item.get("call_id") or item.get("id") - name = item.get("name") - arguments = item.get("arguments") - if isinstance(call_id, str) and isinstance(name, str): - messages.append( - cast( - ConfigData, - AssistantMessage( - content=None, - tool_calls=[ - ToolCall( - id=call_id, - name=name, - arguments=arguments - if isinstance(arguments, str) - else "{}", - ) - ], - ).model_dump(exclude_none=True), - ) - ) - continue - if item_type == "function_call_output": - call_id = item.get("call_id") - if isinstance(call_id, str): - messages.append( - cast( - ConfigData, - ToolMessage( - tool_call_id=call_id, - content=nemo_output_text(item.get("output")), - ).model_dump(exclude_none=True), - ) - ) - continue - if item_type == "message": - messages.append( - cast( - ConfigData, - AssistantMessage( - content=nemo_content_text(item.get("content")), - ).model_dump(exclude_none=True), - ) - ) - return messages - - -def nemo_content_text(content: object) -> str: - if isinstance(content, str): - return content - if isinstance(content, list): - parts: list[str] = [] - for part in content: - if isinstance(part, dict): - part = cast(ConfigMap, part) - text = part.get("text") - if isinstance(text, str): - parts.append(text) - return "\n".join(parts) - return "" if content is None else str(content) - - -def nemo_output_text(output: object) -> str: - if isinstance(output, str): - return output - return json.dumps(jsonable(output), sort_keys=True) - - -def jsonable_mapping(value: ConfigMap) -> ConfigData: - return cast(ConfigData, jsonable(dict(value))) - - -def jsonable(value: object) -> object: - model_dump = getattr(value, "model_dump", None) - if callable(model_dump): - return jsonable(model_dump(exclude_none=True)) - if isinstance(value, dict): - return { - str(key): jsonable(item) for key, item in cast(ConfigMap, value).items() - } - if isinstance(value, list): - return [jsonable(item) for item in value] - if isinstance(value, tuple): - return [jsonable(item) for item in value] - return json.loads(json.dumps(value)) +def load_harness(config: NeMoGymHarnessConfig) -> NeMoGymHarness: + return NeMoGymHarness(config=config) diff --git a/packages/harnesses/harnesses/opencode.py b/packages/harnesses/harnesses/opencode.py index 40031bec7d..b3e4d98087 100644 --- a/packages/harnesses/harnesses/opencode.py +++ b/packages/harnesses/harnesses/opencode.py @@ -1,16 +1,11 @@ import json -import shlex -from pathlib import PurePosixPath -import verifiers as vf -from verifiers.v1.utils.mcp_proxy_utils import proxy_command +import verifiers.v1 as vf -from .utils import split_versioned_agent_spec +from .command import CommandHarness, CommandHarnessConfig, shell_command OPENCODE_DEFAULT_VERSION = "PrimeIntellect-ai/opencode@1.1.63-rl2" -OPENCODE_DEFAULT_AGENT_WORKDIR = "/app" -OPENCODE_DEFAULT_INSTRUCTION_PATH = "/opencode/instruction.txt" -OPENCODE_DEFAULT_SYSTEM_PROMPT_PATH = "/opencode/system.txt" +OPENCODE_DEFAULT_WORKDIR = "/app" OPENCODE_DEFAULT_LOG_PATH = "/logs/agent/opencode.txt" OPENCODE_DEFAULT_SYSTEM_PROMPT = """\ You are OpenCode, an interactive CLI tool that helps users with tasks. @@ -42,174 +37,64 @@ ] -class OpenCodeProgramConfig(vf.ProgramConfig): - agent_workdir: str = OPENCODE_DEFAULT_AGENT_WORKDIR - instruction_path: str = OPENCODE_DEFAULT_INSTRUCTION_PATH - system_prompt_path: str = OPENCODE_DEFAULT_SYSTEM_PROMPT_PATH +class OpenCodeConfig(CommandHarnessConfig): + system_prompt: vf.SystemPrompt | None = OPENCODE_DEFAULT_SYSTEM_PROMPT + version: str = OPENCODE_DEFAULT_VERSION + cwd: str | None = OPENCODE_DEFAULT_WORKDIR log_path: str = OPENCODE_DEFAULT_LOG_PATH disabled_tools: list[str] = OPENCODE_DEFAULT_DISABLED_TOOLS allow_git: bool = False disable_compaction: bool = True - install_ripgrep: bool = True provider_timeout_ms: int = 3_600_000 + max_turns: int = 4 - def resolve(self, version: str = OPENCODE_DEFAULT_VERSION) -> vf.ProgramConfig: - files: dict[str, vf.ProgramValue] = { - self.instruction_path: {"fn": "verifiers.v1.utils.prompt_utils:task_text"}, - self.system_prompt_path: { - "fn": "verifiers.v1.utils.prompt_utils:state_system_prompt_text" - }, - } - artifacts = vf.ArtifactsConfig.model_validate( - { - "opencode_log": { - "path": self.log_path, - "format": "text", - "optional": True, - } - } - ) - ripgrep_install = ( - "apt-get -o Acquire::Retries=3 install -y -qq ripgrep > /dev/null 2>&1 || true" - if self.install_ripgrep - else "" - ) - repo, parsed_version = split_versioned_agent_spec(version) - path = "releases/latest/download" - if parsed_version and parsed_version != "latest": - tag = ( - parsed_version - if parsed_version.startswith("v") - else f"v{parsed_version}" - ) - path = f"releases/download/{tag}" - # Acquire::Retries=3 mitigates transient archive.ubuntu.com CDN sync - # mismatches that fail fresh-sandbox apt-get calls mid-rollout. - setup = f"""\ -set -e -apt-get -o Acquire::Retries=3 update -qq && apt-get -o Acquire::Retries=3 install -y -qq curl tar ca-certificates > /dev/null 2>&1 -{ripgrep_install} - -OPENCODE_RELEASE_REPO={shlex.quote(repo)} -OPENCODE_RELEASE_PATH={shlex.quote(path)} - -case "$(uname -m)" in - x86_64) OPENCODE_ARCH=x64 ;; - aarch64|arm64) OPENCODE_ARCH=arm64 ;; - *) echo "Unsupported architecture: $(uname -m)"; exit 1 ;; -esac -OPENCODE_ASSET="opencode-linux-$OPENCODE_ARCH.tar.gz" -OPENCODE_RELEASE_URL="https://github.com/$OPENCODE_RELEASE_REPO/$OPENCODE_RELEASE_PATH/$OPENCODE_ASSET" +class OpenCode(CommandHarness[OpenCodeConfig]): + config: OpenCodeConfig -mkdir -p "$HOME/.opencode/bin" -if [ -x "$HOME/.opencode/bin/opencode" ]; then - echo "OpenCode already installed, skipping download" -else - curl -fsSL "$OPENCODE_RELEASE_URL" -o /tmp/opencode.tar.gz - tar -xzf /tmp/opencode.tar.gz -C /tmp - install -m 755 /tmp/opencode "$HOME/.opencode/bin/opencode" - rm -f /tmp/opencode.tar.gz /tmp/opencode -fi -""" - agent_config: vf.ConfigData = { - "title": {"disable": True}, - } - opencode_config: vf.ConfigData = { - "${SCHEMA_DOLLAR}schema": "https://opencode.ai/config.json", + def command(self, task: vf.Task, state: vf.State) -> list[str]: + _ = state + instruction = str(getattr(task, "instruction", "")) + opencode_config: dict[str, object] = { "provider": { - "intercepted": { + "openai": { "npm": "@ai-sdk/openai-compatible", - "name": "Intercepted", + "name": "OpenAI", "options": { "baseURL": "$OPENAI_BASE_URL", "apiKey": "${OPENAI_API_KEY:-intercepted}", - "timeout": self.provider_timeout_ms, + "timeout": self.config.provider_timeout_ms, }, "models": { "model": { - "name": "Intercepted Model", + "name": "Model", "modalities": {"input": ["text"], "output": ["text"]}, } }, } }, - "model": "intercepted/model", - # Keep the small-model pin to avoid falling back to the default small - # model and hitting rate limits; disable title calls below. - "small_model": "intercepted/model", - "agent": agent_config, - "mcp": { - "verifiers-tools": { - "type": "local", - "command": proxy_command(), - "enabled": True, + "model": "openai/model", + "small_model": "openai/model", + "agent": { + "build": { + "prompt": "\n\n".join( + str(message.get("content") or "") + for message in self.system_prompt + ), + "tools": {name: False for name in self.config.disabled_tools}, } }, } - if self.disable_compaction: + if self.config.disable_compaction: opencode_config["compaction"] = {"auto": False, "prune": False} - build_config: vf.ConfigData = { - "prompt": "{file:" + self.system_prompt_path + "}" - } - if self.disabled_tools: - build_config["tools"] = {tool: False for tool in self.disabled_tools} - if build_config: - agent_config["build"] = build_config - config_json = json.dumps(opencode_config, indent=2) - log_dir = str(PurePosixPath(self.log_path).parent) - mcp_setup = f"""\ -set -e -export PATH="$HOME/.opencode/bin:$PATH" - -OPENCODE_WORKDIR="${{AGENT_WORKDIR:-}}" -if [ -z "$OPENCODE_WORKDIR" ]; then - OPENCODE_WORKDIR={shlex.quote(self.agent_workdir)} -fi - -mkdir -p ~/.config/opencode {shlex.quote(log_dir)} "$OPENCODE_WORKDIR" -SCHEMA_DOLLAR='$' -cat > ~/.config/opencode/opencode.json << EOFCONFIG -{config_json} -EOFCONFIG -""" - run_script = f"""\ + config_json = json.dumps(opencode_config) + script = f""" set -eo pipefail -export PATH="$HOME/.opencode/bin:$PATH" -export OPENCODE_DISABLE_FILETIME_CHECK=true -export ALLOW_GIT={"1" if self.allow_git else "0"} - -OPENCODE_WORKDIR="${{AGENT_WORKDIR:-}}" -if [ -z "$OPENCODE_WORKDIR" ]; then - OPENCODE_WORKDIR={shlex.quote(self.agent_workdir)} -fi - -cd "$OPENCODE_WORKDIR" -cat {shlex.quote(self.instruction_path)} | opencode run 2>&1 | tee {shlex.quote(self.log_path)} +mkdir -p "$HOME/.config/opencode" "$(dirname {self.config.log_path!r})" +printf '%s' {config_json!r} > "$HOME/.config/opencode/opencode.json" +printf '%s' {instruction!r} | opencode run 2>&1 | tee {self.config.log_path!r} """ - return self.resolve_command( - command=["bash", "-lc", run_script], - files=files, - setup=setup, - artifacts=artifacts, - channels={"mcp": mcp_setup}, - ) - - -class OpenCodeConfig(vf.HarnessConfig): - system_prompt: vf.PromptInput | vf.SystemPromptConfig | None = ( - OPENCODE_DEFAULT_SYSTEM_PROMPT - ) - version: str = OPENCODE_DEFAULT_VERSION - program: OpenCodeProgramConfig = OpenCodeProgramConfig() - max_turns: int = 4 - - -class OpenCode(vf.Harness[OpenCodeConfig]): - config: OpenCodeConfig - - def load_program_config(self, config: OpenCodeConfig) -> vf.ProgramConfig: - return config.program.resolve(version=config.version) + return shell_command(script) def load_harness(config: OpenCodeConfig) -> OpenCode: diff --git a/packages/harnesses/harnesses/pi.py b/packages/harnesses/harnesses/pi.py index 29a6e6be5b..21e2b67367 100644 --- a/packages/harnesses/harnesses/pi.py +++ b/packages/harnesses/harnesses/pi.py @@ -1,141 +1,46 @@ -import json -import shlex -from pathlib import PurePosixPath +import verifiers.v1 as vf -import verifiers as vf -from verifiers.v1.utils.mcp_proxy_utils import proxy_command +from .command import CommandHarness, CommandHarnessConfig, shell_command PI_DEFAULT_VERSION = "@earendil-works/pi-coding-agent@latest" PI_DEFAULT_WORKDIR = "/app" -PI_DEFAULT_INSTRUCTION_PATH = "/pi/instruction.txt" -PI_DEFAULT_SYSTEM_PROMPT_PATH = "/pi/system.txt" PI_DEFAULT_LOG_PATH = "/logs/agent/pi.txt" PI_DEFAULT_SYSTEM_PROMPT = "Complete the user's task using the available tools." -class PiProgramConfig(vf.ProgramConfig): - agent_workdir: str = PI_DEFAULT_WORKDIR - instruction_path: str = PI_DEFAULT_INSTRUCTION_PATH - system_prompt_path: str = PI_DEFAULT_SYSTEM_PROMPT_PATH +class PiConfig(CommandHarnessConfig): + system_prompt: vf.SystemPrompt | None = PI_DEFAULT_SYSTEM_PROMPT + version: str = PI_DEFAULT_VERSION + cwd: str | None = PI_DEFAULT_WORKDIR log_path: str = PI_DEFAULT_LOG_PATH - install_mcp_adapter: bool = True - sandbox: vf.SandboxConfig | None = vf.SandboxConfig() + max_turns: int = 4 - def resolve(self, version: str = PI_DEFAULT_VERSION) -> vf.ProgramConfig: - files: dict[str, vf.ProgramValue] = { - self.instruction_path: {"fn": "verifiers.v1.utils.prompt_utils:task_text"}, - self.system_prompt_path: { - "fn": "verifiers.v1.utils.prompt_utils:state_system_prompt_text" - }, - } - channels: dict[str, vf.ProgramValue] | None = None - if self.install_mcp_adapter: - command, *args = proxy_command() - mcp_json = json.dumps( - { - "mcpServers": { - "verifiers-tools": { - "command": command, - "args": args, - "lifecycle": "lazy", - } - } - }, - indent=2, - ) - models_json = """\ -{ - "providers": { - "verifiers": { - "baseUrl": "${OPENAI_BASE_URL}", - "api": "openai-completions", - "apiKey": "${OPENAI_API_KEY:-intercepted}", - "models": [{"id": "model", "name": "${OPENAI_MODEL}"}] - } - } -} -""" - mcp_setup = f"""\ -set -e -PI_WORKDIR="${{AGENT_WORKDIR:-}}" -if [ -z "$PI_WORKDIR" ]; then - PI_WORKDIR={shlex.quote(self.agent_workdir)} -fi +class Pi(CommandHarness[PiConfig]): + config: PiConfig -mkdir -p "$HOME/.pi/agent" "$PI_WORKDIR" -cat > "$HOME/.pi/agent/models.json" < "$PI_WORKDIR/.mcp.json" <<'EOFMCP' -{mcp_json} -EOFMCP -cd "$PI_WORKDIR" -pi install npm:pi-mcp-adapter -l -""" - channels = { - "mcp": mcp_setup, - } - setup = f"""\ -set -e -apt-get -o Acquire::Retries=3 update -qq && apt-get -o Acquire::Retries=3 install -y -qq curl ca-certificates nodejs npm xz-utils > /dev/null 2>&1 -npm install -g --ignore-scripts n -n 22.19.0 -hash -r -npm install -g --ignore-scripts {shlex.quote(version)} -""" - artifacts = vf.ArtifactsConfig.model_validate( - { - "pi_log": { - "path": self.log_path, - "format": "text", - "optional": True, - } - } + def command(self, task: vf.Task, state: vf.State) -> list[str]: + _ = state + instruction = str( + getattr(task, "instruction", None) or getattr(task, "question", None) or "" + ) + if not instruction: + instruction = "\n\n".join( + str(getattr(message, "content", "") or "") for message in task.prompt + ) + system_prompt = "\n\n".join( + str(message.get("content") or "") for message in self.system_prompt ) - log_dir = str(PurePosixPath(self.log_path).parent) - system_prompt_path = shlex.quote(self.system_prompt_path) - run_script = f"""\ + system_prompt_arg = ( + f"--system-prompt {system_prompt!r}" if system_prompt else "" + ) + script = f""" set -eo pipefail - -PI_WORKDIR="${{AGENT_WORKDIR:-}}" -if [ -z "$PI_WORKDIR" ]; then - PI_WORKDIR={shlex.quote(self.agent_workdir)} -fi - -mkdir -p {shlex.quote(log_dir)} "$PI_WORKDIR" -cd "$PI_WORKDIR" -SYSTEM_PROMPT_ARGS=() -if [ -s {system_prompt_path} ]; then - SYSTEM_PROMPT_ARGS=(--system-prompt "$(cat {system_prompt_path})") -fi -pi --no-session --no-context-files --provider verifiers --model model \ - "${{SYSTEM_PROMPT_ARGS[@]}}" -p @{shlex.quote(self.instruction_path)} 2>&1 | tee {shlex.quote(self.log_path)} +mkdir -p "$(dirname {self.config.log_path!r})" +pi --no-session --no-context-files --provider openai --model "$OPENAI_MODEL" \ + {system_prompt_arg} -p {instruction!r} 2>&1 | tee {self.config.log_path!r} """ - return self.resolve_command( - command=["bash", "-lc", run_script], - default_sandbox=self.sandbox, - files=files, - setup=setup, - env={"OPENAI_MODEL": "runtime.model"}, - artifacts=artifacts, - channels=channels, - ) - - -class PiConfig(vf.HarnessConfig): - system_prompt: vf.PromptInput | vf.SystemPromptConfig | None = ( - PI_DEFAULT_SYSTEM_PROMPT - ) - version: str = PI_DEFAULT_VERSION - program: PiProgramConfig = PiProgramConfig() - max_turns: int = 4 - - -class Pi(vf.Harness[PiConfig]): - config: PiConfig - - def load_program_config(self, config: PiConfig) -> vf.ProgramConfig: - return config.program.resolve(version=config.version) + return shell_command(script) def load_harness(config: PiConfig) -> Pi: diff --git a/packages/harnesses/harnesses/replay.py b/packages/harnesses/harnesses/replay.py index 7aea26411a..9ee0ea61be 100644 --- a/packages/harnesses/harnesses/replay.py +++ b/packages/harnesses/harnesses/replay.py @@ -1,15 +1,19 @@ import time -from typing import cast -import verifiers as vf -from verifiers.types import TrajectoryStep +from pydantic import TypeAdapter +import verifiers.v1 as vf + +_MESSAGES_ADAPTER = TypeAdapter(vf.Messages) class ReplayHarness(vf.Harness[vf.HarnessConfig]): - async def base_program(self, task: vf.Task, state: vf.State) -> vf.State: - await self.runtime.setup_rollout(task, state) - messages = replay_messages(task) - max_turns = state.get_max_turns(self.config.max_turns) + async def run_with_context(self, context: vf.Context) -> None: + task = context.task + state = context.state + task_messages = getattr(task, "messages", None) + if not isinstance(task_messages, list): + raise TypeError("task.messages must be a list.") + messages = _MESSAGES_ADAPTER.validate_python(task_messages) assistant_indices = [ index for index, message in enumerate(messages) @@ -17,110 +21,37 @@ async def base_program(self, task: vf.Task, state: vf.State) -> vf.State: ] if not assistant_indices: raise ValueError("task.messages has no assistant messages.") + max_turns = self.config.max_turns max_turns_reached = max_turns > 0 and max_turns < len(assistant_indices) if max_turns > 0: assistant_indices = assistant_indices[:max_turns] - - state["trajectory"] = [] - model = state.runtime_state().get("model") - model_name = model if isinstance(model, str) and model else "replay" created = int(time.time()) final_turn = len(assistant_indices) - 1 - for turn, message_index in enumerate(assistant_indices): + for turn_index, message_index in enumerate(assistant_indices): message = messages[message_index] - prompt = replay_messages_data(messages[:message_index]) - completion = [replay_message_data(message)] - is_truncated = (max_turns_reached and turn == final_turn) or bool( - completion[0].get("is_truncated", False) - ) - response = replay_response( - message=message, - model=model_name, - created=created, - turn=turn, - trajectory_id=str(state["trajectory_id"]), - is_truncated=is_truncated, - ) - state["trajectory"].append( - replay_trajectory_step( - prompt=prompt, - completion=completion, - response=response, - trajectory_id=str(state["trajectory_id"]), - message_index=message_index, - is_truncated=is_truncated, + if not isinstance(message, vf.AssistantMessage): + raise TypeError( + "Replay assistant indices must point to assistant messages." ) + prompt = messages[:message_index] + is_truncated = max_turns_reached and turn_index == final_turn + message_finish_reason = getattr(message, "finish_reason", None) + message_is_truncated = bool(getattr(message, "is_truncated", False)) + turn = vf.Turn( + prompt=prompt, + completion=[message], + tool_calls=list(message.tool_calls or []), + response_id=f"replay-{state.id}-{turn_index}", + model=context.model or "replay", + created=created, + finish_reason=message_finish_reason + or ("tool_calls" if message.tool_calls else "stop"), + is_truncated=is_truncated or message_is_truncated, ) - if max_turns_reached: - state._set_stop_condition("max_turns_reached") - else: - state._set_stop_condition("replayed_messages") - return state - - -def replay_messages(task: vf.Task) -> vf.Messages: - value = task.get("messages") - if not isinstance(value, list): - raise TypeError("task.messages must be a list.") - return vf.get_messages(value) - - -def replay_message_data(message: vf.Message) -> vf.JsonData: - return cast(vf.JsonData, message.model_dump(mode="json", exclude_none=True)) - - -def replay_messages_data(messages: list[vf.Message]) -> list[vf.JsonData]: - return [replay_message_data(message) for message in messages] - - -def replay_trajectory_step( - *, - prompt: list[vf.JsonData], - completion: list[vf.JsonData], - response: vf.JsonData, - trajectory_id: str, - message_index: int, - is_truncated: bool, -) -> TrajectoryStep: - return cast( - TrajectoryStep, - { - "prompt": prompt, - "completion": completion, - "response": response, - "tokens": None, - "reward": None, - "advantage": None, - "is_truncated": is_truncated, - "trajectory_id": trajectory_id, - "extras": {"replay": True, "message_index": message_index}, - }, - ) - - -def replay_response( - *, - message: vf.Message, - model: str, - created: int, - turn: int, - trajectory_id: str, - is_truncated: bool, -) -> vf.JsonData: - message_data = replay_message_data(message) - message_data.setdefault( - "finish_reason", "tool_calls" if message_data.get("tool_calls") else "stop" - ) - message_data["is_truncated"] = is_truncated or bool( - message_data.get("is_truncated", False) - ) - return { - "id": f"replay-{trajectory_id}-{turn}", - "created": created, - "model": model, - "usage": None, - "message": cast(vf.JsonData, message_data), - } + state.transcript.append(turn) + if turn.is_truncated: + state.is_truncated = True + state.stop("max_turns" if max_turns_reached else "replayed_messages") def load_harness(config: vf.HarnessConfig) -> ReplayHarness: diff --git a/packages/harnesses/harnesses/rlm.py b/packages/harnesses/harnesses/rlm.py index aa728b8ae0..7d51f10384 100644 --- a/packages/harnesses/harnesses/rlm.py +++ b/packages/harnesses/harnesses/rlm.py @@ -1,221 +1,73 @@ -import shlex - -import verifiers as vf - -from .utils.rlm_utils import ( - DEFAULT_RLM_CHECKOUT_PATH, - DEFAULT_RLM_SKILLS_PATH, - DEFAULT_RLM_TOOL_SKILL_MARKER, - DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH, - DEFAULT_RLM_TOOL_SKILLS_MANIFEST_NAME, -) - -RLM_DEFAULT_REPO_URL = "github.com/PrimeIntellect-ai/rlm-harness.git" -RLM_DEFAULT_REF = "main" -RLM_DEFAULT_EXEC_TIMEOUT = 300 -RLM_DEFAULT_MAX_DEPTH = 0 -RLM_DEFAULT_INSTRUCTION_PATH = "/rlm/instruction.txt" -RLM_DEFAULT_APPEND_TO_SYSTEM_PROMPT_PATH = "/rlm/append_to_system_prompt.txt" +from pydantic import Field + +import verifiers.v1 as vf + +from .command import CommandHarness, CommandHarnessConfig, shell_command + RLM_DEFAULT_WORKDIR = "/workspace" RLM_DEFAULT_TOOLS = ["ipython"] -class RLMProgramConfig(vf.ProgramConfig): - sandbox: vf.SandboxConfig | None = None - workdir: str = RLM_DEFAULT_WORKDIR - instruction_path: str = RLM_DEFAULT_INSTRUCTION_PATH - repo_url: str = RLM_DEFAULT_REPO_URL - ref: str = RLM_DEFAULT_REF - exec_timeout: int = RLM_DEFAULT_EXEC_TIMEOUT - max_depth: int = RLM_DEFAULT_MAX_DEPTH +class RLMConfig(CommandHarnessConfig): + command: list[str] = Field(default_factory=list) + cwd: str | None = RLM_DEFAULT_WORKDIR + tools: list[str] = Field(default_factory=lambda: list(RLM_DEFAULT_TOOLS)) + exec_timeout: int = 300 + max_depth: int = 0 summarize_at_tokens: int | None = None append_to_system_prompt: str = "" - local_checkout: str | None = None - gh_token_var: str | None = "GH_TOKEN" - tools: list[str] = RLM_DEFAULT_TOOLS - env_vars: dict[str, str] = {} - skills: str | None = None - - def resolve(self) -> vf.ProgramConfig: - files: dict[str, vf.ProgramValue] = { - self.instruction_path: { - "fn": "verifiers.v1.utils.prompt_utils:task_text", - "keys": ["instruction", "question"], - }, - RLM_DEFAULT_APPEND_TO_SYSTEM_PROMPT_PATH: self.append_to_system_prompt, - DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH: { - "fn": "harnesses.utils.rlm_utils:rlm_tool_skills_archive" - }, - } - dirs: dict[str, vf.ProgramValue] = { - DEFAULT_RLM_CHECKOUT_PATH: { - "fn": "harnesses.utils.rlm_utils:rlm_checkout_path", - **( - {"local_checkout": self.local_checkout} - if self.local_checkout - else {} - ), - "repo_url": self.repo_url, - "ref": self.ref, - **({"gh_token_var": self.gh_token_var} if self.gh_token_var else {}), - } - } - if self.skills is not None: - dirs[DEFAULT_RLM_SKILLS_PATH] = self.skills - else: - dirs[DEFAULT_RLM_SKILLS_PATH] = { - "fn": "harnesses.utils.rlm_utils:rlm_skills_dir" - } - - env: dict[str, vf.ProgramValue] = { - "PATH": "/root/.local/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin", - "OPENAI_MODEL": "runtime.model", - "RLM_MODEL": "runtime.model", - "RLM_TOOLS": ",".join(self.tools), - "RLM_EXEC_TIMEOUT": str(self.exec_timeout), - "RLM_MAX_DEPTH": str(self.max_depth), - **self.env_vars, - } - if self.summarize_at_tokens is not None: - assert self.summarize_at_tokens > 0 - env["RLM_SUMMARIZE_AT_TOKENS"] = str(self.summarize_at_tokens) - - artifacts = vf.ArtifactsConfig.model_validate( - { - "rlm_metrics": { - "path": f"{self.workdir}/.rlm/sessions/*/meta.json", - "format": "json", - "key": "metrics", - "optional": True, - } - } - ) - command_timeout = max(self.exec_timeout + 120, 600) - setup_timeout = command_timeout - if self.sandbox is not None and "setup_timeout" in self.sandbox.data( - fill_defaults=False - ): - setup_timeout = self.sandbox.setup_timeout - - if self.sandbox is None: - sandbox = vf.SandboxConfig( - image="python:3.11-slim", - workdir=self.workdir, - cpu_cores=1, - memory_gb=2, - disk_size_gb=5, - network_access=True, - timeout_minutes=60, - command_timeout=command_timeout, - setup_timeout=setup_timeout, - ) - else: - sandbox = vf.SandboxConfig.model_validate( - { - "workdir": self.workdir, - "command_timeout": command_timeout, - **self.sandbox.data(), - "setup_timeout": setup_timeout, - } - ) - - skills_install_script = f""" -set -eo pipefail -skills_path={shlex.quote(DEFAULT_RLM_SKILLS_PATH)} -archive_path={shlex.quote(DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH)} -manifest_path="$skills_path/{DEFAULT_RLM_TOOL_SKILLS_MANIFEST_NAME}" -mkdir -p "$skills_path" -if [ -f "$manifest_path" ]; then - while IFS= read -r skill_name; do - case "$skill_name" in ""|.*|*/*|*..*) continue ;; esac - if [ -f "$skills_path/$skill_name/{DEFAULT_RLM_TOOL_SKILL_MARKER}" ]; then - rm -rf "$skills_path/$skill_name" - fi - done < "$manifest_path" - rm -f "$manifest_path" -fi -if [ -s "$archive_path" ]; then - tmp_archive="$(mktemp)" - trap 'rm -f "$tmp_archive"' EXIT - base64 -d "$archive_path" > "$tmp_archive" - tar -tzf "$tmp_archive" \\ - | awk -F/ 'NF > 1 && $1 != "" {{print $1}}' \\ - | sort -u > "$manifest_path" - tar -xzf "$tmp_archive" -C "$skills_path" -fi -""" - checkout_install_script = f""" -set -eo pipefail -export RLM_CHECKOUT_PATH={shlex.quote(DEFAULT_RLM_CHECKOUT_PATH)} -test -f "$RLM_CHECKOUT_PATH/install.sh" -bash "$RLM_CHECKOUT_PATH/install.sh" -""" - run_script = f""" -set -eo pipefail -export PATH="$HOME/.local/bin:${{AGENT_PATH:-$PATH}}" -export RLM_MODEL="${{RLM_MODEL:-$OPENAI_MODEL}}" -export OPENAI_API_KEY="${{OPENAI_API_KEY:-intercepted}}" -export RLM_APPEND_TO_SYSTEM_PROMPT="$(cat {shlex.quote(RLM_DEFAULT_APPEND_TO_SYSTEM_PROMPT_PATH)} 2>/dev/null || true)" -cd "${{AGENT_WORKDIR:-{self.workdir}}}" -rlm "$(cat {shlex.quote(self.instruction_path)})" -""" - return self.resolve_command( - command=["bash", "-lc", run_script], - files=files, - dirs=dirs, - setup=[ - "apt-get -o Acquire::Retries=3 update && " - "apt-get -o Acquire::Retries=3 install -y --no-install-recommends " - "ca-certificates curl git && rm -rf /var/lib/apt/lists/*", - "bash -lc " + shlex.quote(skills_install_script), - "bash -lc " + shlex.quote(checkout_install_script), - ], - env=env, - artifacts=artifacts, - sandbox=sandbox, - setup_timeout=setup_timeout, - ) - - -class RLMConfig(vf.HarnessConfig): - program: RLMProgramConfig = RLMProgramConfig() -class RLMEndpoint(vf.Endpoint): - def trajectory_visibility(self, headers: dict[str, str]) -> vf.TrajectoryVisibility: - if str(headers.get("x-rlm-depth", "0")) != "0": - return "hidden" - return super().trajectory_visibility(headers) - - -class RLM(vf.Harness[RLMConfig]): +class RLM(CommandHarness[RLMConfig]): config: RLMConfig - def load_endpoint(self) -> vf.Endpoint: - return RLMEndpoint( - use_tunnel=self.program_sandbox_config(self.program_config) is not None + def command(self, task: vf.Task, state: vf.State) -> list[str]: + if self.config.command: + return list(self.config.command) + _ = state + instruction = str( + getattr(task, "instruction", None) or getattr(task, "question", None) or "" + ) + if not instruction: + instruction = "\n\n".join( + str(getattr(message, "content", "") or "") for message in task.prompt + ) + return shell_command(f"rlm {instruction!r}") + + def command_env(self, task: vf.Task, state: vf.State) -> dict[str, str]: + env = super().command_env(task, state) + env.update( + { + "RLM_MODEL": state.model.model if state.model is not None else "", + "RLM_TOOLS": ",".join(self.config.tools), + "RLM_EXEC_TIMEOUT": str(self.config.exec_timeout), + "RLM_MAX_DEPTH": str(self.config.max_depth), + "RLM_APPEND_TO_SYSTEM_PROMPT": self.config.append_to_system_prompt, + } ) + if self.config.summarize_at_tokens is not None: + env["RLM_SUMMARIZE_AT_TOKENS"] = str(self.config.summarize_at_tokens) + return env @vf.metric async def rlm_sub_llm_call_count(self, state: vf.State) -> float: - metrics = state["artifacts"].get("rlm_metrics") or {} - assert isinstance(metrics, dict) - value = metrics.get("sub_llm_call_count", 0.0) - return float(value or 0.0) + return rlm_metric(state, "sub_llm_call_count") @vf.metric async def rlm_sub_llm_total_turns(self, state: vf.State) -> float: - metrics = state["artifacts"].get("rlm_metrics") or {} - assert isinstance(metrics, dict) - value = metrics.get("sub_llm_total_turns", 0.0) - return float(value or 0.0) + return rlm_metric(state, "sub_llm_total_turns") @vf.metric async def rlm_sub_llm_total_tool_calls(self, state: vf.State) -> float: - metrics = state["artifacts"].get("rlm_metrics") or {} - assert isinstance(metrics, dict) - value = metrics.get("sub_llm_total_tool_calls", 0.0) - return float(value or 0.0) + return rlm_metric(state, "sub_llm_total_tool_calls") + + +def rlm_metric(state: vf.State, name: str) -> float: + metrics = state.artifacts.get("rlm_metrics") + if not isinstance(metrics, dict): + return 0.0 + value = metrics.get(name) + return float(value) if isinstance(value, int | float) else 0.0 def load_harness(config: RLMConfig) -> RLM: diff --git a/packages/harnesses/harnesses/terminus_2.py b/packages/harnesses/harnesses/terminus_2.py index 4f8bb08e72..7f814913f5 100644 --- a/packages/harnesses/harnesses/terminus_2.py +++ b/packages/harnesses/harnesses/terminus_2.py @@ -1,199 +1,49 @@ -import shlex -from pathlib import PurePosixPath +import verifiers.v1 as vf -import verifiers as vf -from verifiers.v1.utils.sandbox_python_utils import SANDBOX_BIN_DIR, uv_setup_command +from .command import CommandHarness, CommandHarnessConfig, shell_command -TERMINUS_2_DEFAULT_AGENT_WORKDIR = "/app" -TERMINUS_2_DEFAULT_INSTRUCTION_PATH = "/terminus_2/instruction.md" -TERMINUS_2_DEFAULT_SYSTEM_PROMPT_PATH = "/terminus_2/system_prompt.txt" +TERMINUS_2_DEFAULT_WORKDIR = "/app" TERMINUS_2_DEFAULT_LOG_PATH = "/logs/agent/terminus_2.log" TERMINUS_2_DEFAULT_VERSION = "harbor==0.6.6" -TERMINUS_2_DEFAULT_PYTHON_VERSION = "3.12" TERMINUS_2_DEFAULT_MODEL_NAME = "openai/gpt-4.1-mini" TERMINUS_2_DEFAULT_API_BASE_URL = "https://api.pinference.ai/api/v1" -class Terminus2ProgramConfig(vf.ProgramConfig): - agent_workdir: str = TERMINUS_2_DEFAULT_AGENT_WORKDIR - instruction_path: str = TERMINUS_2_DEFAULT_INSTRUCTION_PATH - system_prompt_path: str = TERMINUS_2_DEFAULT_SYSTEM_PROMPT_PATH +class Terminus2Config(CommandHarnessConfig): + version: str = TERMINUS_2_DEFAULT_VERSION + cwd: str | None = TERMINUS_2_DEFAULT_WORKDIR log_path: str = TERMINUS_2_DEFAULT_LOG_PATH - python_version: str = TERMINUS_2_DEFAULT_PYTHON_VERSION model_name: str = TERMINUS_2_DEFAULT_MODEL_NAME api_base_url: str = TERMINUS_2_DEFAULT_API_BASE_URL - sandbox: vf.SandboxConfig | None = vf.SandboxConfig() max_turns: int = 4 - def resolve(self, version: str = TERMINUS_2_DEFAULT_VERSION) -> vf.ProgramConfig: - files: dict[str, vf.ProgramValue] = { - self.instruction_path: {"fn": "verifiers.v1.utils.prompt_utils:task_text"}, - self.system_prompt_path: { - "fn": "verifiers.v1.utils.prompt_utils:state_system_prompt_text" - }, - } - artifacts = vf.ArtifactsConfig.model_validate( - { - "terminus_2_log": { - "path": self.log_path, - "format": "text", - "optional": True, - } - } - ) - log_dir = str(PurePosixPath(self.log_path).parent) - system_prompt_block = f"""\ - system_prompt_path = Path({self.system_prompt_path!r}) - if system_prompt_path.exists() and system_prompt_path.stat().st_size > 0: - instruction = system_prompt_path.read_text() + "\\n\\n" + instruction -""" - agent_script = f"""\ -import asyncio -import logging -import os -import shutil -import subprocess -from pathlib import Path - -from harbor.agents.terminus_2 import Terminus2 -from harbor.environments.base import BaseEnvironment, ExecResult -from harbor.models.agent.context import AgentContext -from harbor.models.environment_type import EnvironmentType -from harbor.models.trial.paths import TrialPaths - - -class LocalEnvironment(BaseEnvironment): - def __init__(self, workdir: Path, logs_dir: Path): - self.workdir = workdir - self.trial_paths = TrialPaths(trial_dir=logs_dir) - self.trial_paths.mkdir() - self.default_user = None - self.session_id = "local" - self.logger = logging.getLogger(__name__) - - def type(self) -> EnvironmentType: - return EnvironmentType.DOCKER - - @property - def is_mounted(self) -> bool: - return True - - @property - def supports_gpus(self) -> bool: - return False - - @property - def can_disable_internet(self) -> bool: - return False - - def _validate_definition(self) -> None: - return None - - async def start(self, force_build: bool) -> None: - return None - - async def stop(self, delete: bool) -> None: - return None - - async def prepare_logs_for_host(self) -> None: - return None - - async def upload_file(self, source_path, target_path): - shutil.copy(source_path, target_path) - - async def upload_dir(self, source_dir, target_dir): - shutil.copytree(source_dir, target_dir, dirs_exist_ok=True) - - async def download_file(self, source_path, target_path): - shutil.copy(source_path, target_path) - async def download_dir(self, source_dir, target_dir): - shutil.copytree(source_dir, target_dir, dirs_exist_ok=True) +class Terminus2(CommandHarness[Terminus2Config]): + config: Terminus2Config - async def exec( - self, - command: str, - cwd: str | None = None, - env: dict | None = None, - timeout_sec: int | None = None, - user: str | int | None = None, - ) -> ExecResult: - _ = user - try: - result = subprocess.run( - command, - shell=True, - cwd=cwd or str(self.workdir), - env={{**os.environ, **(env or {{}})}}, - capture_output=True, - text=True, - timeout=timeout_sec, + def command(self, task: vf.Task, state: vf.State) -> list[str]: + _ = state + instruction = str( + getattr(task, "instruction", None) or getattr(task, "question", None) or "" + ) + if not instruction: + instruction = "\n\n".join( + str(getattr(message, "content", "") or "") for message in task.prompt ) - except subprocess.TimeoutExpired: - return ExecResult(stdout="", stderr="Command timed out", return_code=124) - return ExecResult( - stdout=result.stdout, - stderr=result.stderr, - return_code=result.returncode, + system_prompt = "\n\n".join( + str(message.get("content") or "") for message in self.system_prompt ) - - -async def main() -> None: - workdir = Path(os.environ.get("AGENT_WORKDIR") or {TERMINUS_2_DEFAULT_AGENT_WORKDIR!r}) - logs_dir = Path({log_dir!r}) - instruction = Path({self.instruction_path!r}).read_text() -{system_prompt_block} env = LocalEnvironment(workdir=workdir, logs_dir=logs_dir) - api_base = os.environ.get("OPENAI_BASE_URL") or {self.api_base_url!r} - agent = Terminus2( - logs_dir=logs_dir, - model_name={self.model_name!r}, - api_base=api_base, - max_turns={self.max_turns!r}, - ) - await agent.setup(env) - await agent.run(instruction, env, AgentContext()) - - -asyncio.run(main()) -""" - run_script = f"""\ + if system_prompt: + instruction = f"{system_prompt}\n\n{instruction}" + script = f""" set -eo pipefail -export PATH={shlex.quote(SANDBOX_BIN_DIR)}:"$HOME/.local/bin:$PATH" - -TERMINUS_2_WORKDIR="${{AGENT_WORKDIR:-}}" -if [ -z "$TERMINUS_2_WORKDIR" ]; then - TERMINUS_2_WORKDIR={shlex.quote(self.agent_workdir)} -fi -export AGENT_WORKDIR="$TERMINUS_2_WORKDIR" - -mkdir -p {shlex.quote(log_dir)} "$TERMINUS_2_WORKDIR" -cd "$TERMINUS_2_WORKDIR" -uv --no-config run --no-project --quiet \ - --python {shlex.quote(self.python_version)} \ - --with {shlex.quote(version)} \ - python - <<'PY' 2>&1 | tee -a {shlex.quote(self.log_path)} -{agent_script} -PY +mkdir -p "$(dirname {self.config.log_path!r})" +python -m harbor.agents.terminus_2 \ + --model {self.config.model_name!r} \ + --api-base "${{OPENAI_BASE_URL:-{self.config.api_base_url}}}" \ + --instruction {instruction!r} 2>&1 | tee -a {self.config.log_path!r} """ - return self.resolve_command( - command=["bash", "-lc", run_script], - default_sandbox=self.sandbox, - files=files, - setup=uv_setup_command(), - artifacts=artifacts, - ) - - -class Terminus2Config(vf.HarnessConfig): - version: str = TERMINUS_2_DEFAULT_VERSION - program: Terminus2ProgramConfig = Terminus2ProgramConfig() - - -class Terminus2(vf.Harness[Terminus2Config]): - config: Terminus2Config - - def load_program_config(self, config: Terminus2Config) -> vf.ProgramConfig: - return config.program.resolve(version=config.version) + return shell_command(script) def load_harness(config: Terminus2Config) -> Terminus2: diff --git a/packages/harnesses/harnesses/utils/rlm_utils.py b/packages/harnesses/harnesses/utils/rlm_utils.py deleted file mode 100644 index ab24a8fe07..0000000000 --- a/packages/harnesses/harnesses/utils/rlm_utils.py +++ /dev/null @@ -1,306 +0,0 @@ -import base64 -import io -import json -import keyword -import os -import re -import tarfile -import textwrap -from collections.abc import Callable -from importlib.abc import Traversable -from pathlib import Path -from typing import cast - -from verifiers.envs.experimental.utils.git_checkout_cache import ( - resolve_git_checkout, - validate_git_checkout, -) -from verifiers.v1.runtime import Runtime -from verifiers.v1.state import State -from verifiers.v1.types import ConfigData - -DEFAULT_RLM_CHECKOUT_PATH = "/tmp/rlm-checkout" -DEFAULT_RLM_SKILLS_PATH = "/task/rlm-skills" -DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH = "/tmp/vf-rlm-tool-skills.tar.gz.b64" -DEFAULT_RLM_TOOL_SKILLS_MANIFEST_NAME = ".vf-generated-tool-skills" -DEFAULT_RLM_TOOL_SKILL_MARKER = ".vf-generated-tool-skill" -DEFAULT_RLM_LOCAL_CHECKOUT_CACHE_ROOT = ( - Path.home() / ".cache" / "verifiers" / "rlm-checkouts" -) -REQUIRED_RLM_CHECKOUT_FILES = ("install.sh", "pyproject.toml") - - -def rlm_checkout_path( - repo_url: str, - ref: str, - local_checkout: str | None = None, - gh_token_var: str | None = "GH_TOKEN", -) -> Path: - return rlm_checkout_loader( - local_checkout=local_checkout, - repo_url=repo_url, - ref=ref, - gh_token_var=gh_token_var, - )() - - -def rlm_checkout_loader( - local_checkout: str | Path | None, - repo_url: str, - ref: str, - gh_token_var: str | None, -) -> Callable[[], Path]: - checkout: Path | None = None - - def load() -> Path: - nonlocal checkout - if checkout is not None: - return checkout - if local_checkout is not None: - checkout = validate_git_checkout( - Path(local_checkout), - required_files=REQUIRED_RLM_CHECKOUT_FILES, - ) - else: - checkout = resolve_git_checkout( - repo_url=repo_url, - ref=ref, - cache_root=DEFAULT_RLM_LOCAL_CHECKOUT_CACHE_ROOT, - gh_token=os.environ.get(gh_token_var) if gh_token_var else None, - required_files=REQUIRED_RLM_CHECKOUT_FILES, - ) - return checkout - - return load - - -def rlm_tool_skills_archive(state: State, runtime: Runtime) -> str: - tool_defs = runtime.tool_defs(state) or [] - if not tool_defs: - return "" - used_names: set[str] = set() - skills_dir = rlm_skills_dir(state, runtime) - if skills_dir is not None: - used_names.update( - child.name for child in skills_dir.iterdir() if child.is_dir() - ) - buffer = io.BytesIO() - with tarfile.open(fileobj=buffer, mode="w:gz") as tar: - for raw_tool_def in tool_defs: - tool_def = cast(ConfigData, raw_tool_def.model_dump()) - tool_name = str(tool_def["name"]) - skill_name = re.sub(r"\W", "_", tool_name) - if not skill_name or skill_name[0].isdigit(): - skill_name = f"tool_{skill_name}" - if keyword.iskeyword(skill_name): - skill_name = f"{skill_name}_tool" - base_name = skill_name - index = 2 - while skill_name in used_names: - skill_name = f"{base_name}_{index}" - index += 1 - used_names.add(skill_name) - description = str( - tool_def.get("description") or f"Call the {tool_name} verifier tool." - ) - schema = json.dumps( - tool_def.get("parameters") or {}, indent=2, sort_keys=True - ) - parameters = cast(ConfigData, tool_def.get("parameters") or {}) - properties = cast(ConfigData, parameters.get("properties") or {}) - allowed_arguments = ( - sorted(properties) - if parameters.get("additionalProperties") is False - else None - ) - required = set(cast(list[str], parameters.get("required") or [])) - type_names = { - "array": "list", - "boolean": "bool", - "integer": "int", - "null": "None", - "number": "float", - "object": "dict", - "string": "str", - } - typed_parameters = all( - name.isidentifier() - and not name.startswith("_") - and not keyword.iskeyword(name) - and name - not in { - "arguments", - "kwargs", - } - for name in properties - ) - if typed_parameters: - signature_parts: list[str] = [] - argument_lines = ["arguments = {}"] - for name, raw_schema in sorted( - properties.items(), - key=lambda item: ( - "default" in cast(ConfigData, item[1]) - or item[0] not in required - ), - ): - field_schema = cast(ConfigData, raw_schema) - raw_type = field_schema.get("type") - if isinstance(raw_type, str): - annotation_parts = [type_names.get(raw_type, "object")] - elif isinstance(raw_type, list): - annotation_parts = [ - type_names.get(str(item), "object") for item in raw_type - ] - else: - annotation_parts = ["object"] - annotation = " | ".join(dict.fromkeys(annotation_parts)) - if name not in required and "default" not in field_schema: - if "None" not in annotation_parts: - annotation = f"{annotation} | None" - signature_parts.append(f"{name}: {annotation} = None") - argument_lines.append(f"if {name} is not None:") - argument_lines.append(f" arguments[{name!r}] = {name}") - elif "default" in field_schema: - default = field_schema["default"] - if default is None and "None" not in annotation_parts: - annotation = f"{annotation} | None" - signature_parts.append(f"{name}: {annotation} = {default!r}") - argument_lines.append(f"arguments[{name!r}] = {name}") - else: - signature_parts.append(f"{name}: {annotation}") - argument_lines.append(f"arguments[{name!r}] = {name}") - signature_parts.append("**kwargs") - signature = ", ".join(signature_parts) - argument_lines.append("arguments.update(kwargs)") - argument_source = textwrap.indent("\n".join(argument_lines), " " * 20) - call_example = f'result = await {skill_name}(argument_name="value")' - else: - signature = "arguments: dict | None = None, **kwargs" - argument_source = textwrap.indent( - "arguments = {**(arguments or {}), **kwargs}", " " * 20 - ) - call_example = ( - f"result = await {skill_name}({{'argument_name': 'value'}})" - ) - module_imports = "os, requests" - dependencies = ["requests", "rlm"] - call_source = textwrap.indent( - textwrap.dedent( - f"""\ - base = os.environ.get("ANTHROPIC_BASE_URL") or os.environ.get("OPENAI_BASE_URL") - if not base: - raise RuntimeError("No Verifiers endpoint URL is configured.") - api_key = ( - os.environ.get("OPENAI_API_KEY") - or os.environ.get("ANTHROPIC_API_KEY") - or "intercepted" - ) - response = requests.post( - f"{{base.rsplit('/v1', 1)[0].rstrip('/')}}/vf/tools/" + {tool_name!r}, - json={{"arguments": arguments}}, - headers={{"Authorization": f"Bearer {{api_key}}"}}, - timeout=300, - ) - if not response.content: - response.raise_for_status() - return None - payload = response.json() - if "error" in payload: - raise RuntimeError(str(payload["error"])) - response.raise_for_status() - return payload.get("result") - """ - ), - " " * 20, - ) - module = textwrap.dedent( - f"""\ - import {module_imports} - - - TOOL_ALLOWED_ARGUMENTS = {allowed_arguments!r} - - - async def run({signature}) -> object: - {json.dumps(description)} -{argument_source} - if TOOL_ALLOWED_ARGUMENTS is not None: - arguments = {{ - key: arguments[key] - for key in TOOL_ALLOWED_ARGUMENTS - if key in arguments - }} -{call_source} - """ - ) - distribution_name = skill_name.replace("_", "-") - files = { - f"{skill_name}/SKILL.md": f"""# {skill_name} - -{description} - -This skill calls `{tool_name}`. - -Call it with tool arguments: - -```python -{call_example} -``` - -Tool schema: - -```json -{schema} -``` -""", - f"{skill_name}/{DEFAULT_RLM_TOOL_SKILL_MARKER}": "1\n", - f"{skill_name}/pyproject.toml": textwrap.dedent( - f"""\ - [project] - name = "rlm-skill-{distribution_name}" - version = "0.0.0" - dependencies = {json.dumps(dependencies)} - - [project.scripts] - {skill_name} = "rlm.skill:cli" - - [build-system] - requires = ["hatchling"] - build-backend = "hatchling.build" - - [tool.hatch.build.targets.wheel] - packages = ["src/{skill_name}"] - """ - ), - f"{skill_name}/src/{skill_name}/__init__.py": ( - f"from .{skill_name} import run\n\n__all__ = ['run']\n" - ), - f"{skill_name}/src/{skill_name}/{skill_name}.py": module, - } - for path, content in files.items(): - data = content.encode() - info = tarfile.TarInfo(path) - info.size = len(data) - tar.addfile(info, io.BytesIO(data)) - return base64.b64encode(buffer.getvalue()).decode() - - -def rlm_skills_dir(state: State, runtime: Runtime) -> Path | Traversable | None: - from harnesses.rlm import RLM - - harness = runtime.harness - if not isinstance(harness, RLM): - raise TypeError("rlm_skills_dir requires an RLM harness runtime.") - if harness.config.program.skills is not None: - return Path(harness.config.program.skills) - taskset = runtime.taskset - if taskset is None: - return None - upload_dirs = taskset.get_upload_dirs() - assert isinstance(upload_dirs, dict) - skills_dir = upload_dirs.get("skills") - if skills_dir is None: - return None - assert isinstance(skills_dir, (Path, Traversable)) - return skills_dir diff --git a/packages/tasksets/README.md b/packages/tasksets/README.md index 3feb8808e7..c91c44c4ad 100644 --- a/packages/tasksets/README.md +++ b/packages/tasksets/README.md @@ -1,6 +1,6 @@ # tasksets -Reusable v1 `vf.Taskset` implementations for Verifiers. +Reusable `verifiers.v1` taskset implementations. Tasksets own task data, task controls, task-owned tools, user behavior, rewards, metrics, and task-specific setup/cleanup. They are sibling packages to @@ -35,20 +35,12 @@ Environment packages should expose a typed child loader and let Verifiers coerce the `[env.taskset]` config through that annotation: ```python -import verifiers as vf +import verifiers.v1 as vf from tasksets import HarborTaskset, HarborTasksetConfig def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: return HarborTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` Do not mutate config objects in loaders. Put defaults on the config class or pass @@ -58,11 +50,11 @@ the intended config from Python/TOML. | Taskset | Purpose | | --- | --- | -| `HarborTaskset` | Harbor task directories and Harbor Hub datasets. | -| `OpenEnvTaskset` | Upstream OpenEnv projects with out-of-the-box task/tool use. | -| `OpenRewardTaskset` | Upstream OpenReward environments and rollout-local session tools. | +| `HarborTaskset` | Harbor task directories and Harbor datasets. | +| `OpenEnvTaskset` | Upstream OpenEnv projects with rollout-local user simulation. | +| `OpenRewardTaskset` | Upstream OpenReward environments with rollout-local user simulation. | | `ReplayTaskset` | HF datasets or explicit local JSONL chat transcripts for replay data. | -| `TextArenaTaskset` | Compatible TextArena single-player games with a taskset-owned `vf.User`. | +| `TextArenaTaskset` | Compatible TextArena single-player games with an MCP user server. | | `NeMoGymTaskset` | NeMo Gym JSONL task rows. | Taskset implementations follow the same rules as environment-local tasksets: @@ -72,8 +64,8 @@ serializable, and utilities exist only for shared messy internals. ## Replay Transcript Data Use `ReplayTaskset` with `ReplayHarness` when each training example is already a -chat transcript row and each assistant message should become one trajectory -step. +chat transcript row and each assistant message should become one transcript +turn. For local data, put one JSON object per line in `.jsonl` files under a directory owned by the env package: @@ -85,7 +77,7 @@ owned by the env package: `messages` must be a JSON array of message objects. Each message must have a string `role`; all other message fields are preserved. Assistant messages may appear anywhere in the transcript, and every assistant message is replayed as -one trajectory step. +one transcript turn. Set that local directory explicitly, either through `[env.taskset].data_dir` or on the env-local taskset subclass: @@ -93,7 +85,7 @@ on the env-local taskset subclass: ```python from pathlib import Path -import verifiers as vf +import verifiers.v1 as vf from harnesses import ReplayHarness from tasksets import ReplayTaskset, ReplayTasksetConfig diff --git a/packages/tasksets/tasksets/harbor.py b/packages/tasksets/tasksets/harbor.py index 62b90b6f8e..3d6a7776ed 100644 --- a/packages/tasksets/tasksets/harbor.py +++ b/packages/tasksets/tasksets/harbor.py @@ -1,203 +1,286 @@ +import io +import tarfile from pathlib import Path -from typing import cast +from typing import Literal -import verifiers as vf -from verifiers.utils.import_utils import load_toml -from verifiers.v1.utils.sandbox_utils import SandboxClient +from pydantic import BaseModel, Field +import verifiers.v1 as vf +from verifiers.v1.utils.json_utils import json_data +from verifiers.utils.import_utils import load_toml from tasksets.utils.harbor_utils import ( TASKS_SUBDIR, bundle_tasks_root, download_harbor_dataset, - harbor_sandbox, harbor_task_dirs, parse_gb, parse_number, parse_reward_text, - upload_harbor_tests, ) -HARBOR_DEFAULT_SANDBOX = vf.SandboxConfig( - image="python:3.11-slim", - cpu_cores=2.0, - memory_gb=4.0, - disk_size_gb=10.0, - timeout_minutes=120, - workdir="/app", - command_timeout=900, -) +HarborSource = Literal["harbor", "package"] class HarborTasksetConfig(vf.TasksetConfig): - taskset_id: str | None = "harbor" - dataset: str | None = None - bundle_package: str | None = None - task_names: list[str] | None = None + id: str | None = "harbor" + source: HarborSource = "harbor" + dataset: str = "hello-world" + tasks: list[str] | None = None cache_dir: str | None = None refresh: bool = False - sandbox: vf.SandboxConfig = HARBOR_DEFAULT_SANDBOX - verifier_timeout_seconds: float = 900.0 - task_dir: str = "/task" - env: dict[str, str] = {} + require_image: bool = False + + +class Author(BaseModel, extra="forbid", frozen=True): + name: str | None = None + email: str | None = None + + +class HarborTask(vf.Task, frozen=True): + task_name: str + instruction: str + agent_timeout: float | None = None + scoring_timeout: float | None = None + keywords: list[str] = Field(default_factory=list) + authors: list[Author] = Field(default_factory=list) + difficulty: str | None = None + category: str | None = None + tags: list[str] = Field(default_factory=list) + task_dir: str = Field(default="", exclude=True) + + @classmethod + def from_dir(cls, task_dir: Path, *, require_image: bool) -> "HarborTask": + task_toml_path = task_dir / "task.toml" + instruction_path = task_dir / "instruction.md" + with task_toml_path.open("rb") as f: + task_config = load_toml(f) + sections: dict[str, vf.JsonData] = {} + for name in ("environment", "task", "metadata", "agent", "verifier"): + section = task_config.get(name) + if section is None: + section = {} + if not isinstance(section, dict): + raise TypeError(f"Harbor task [{name}] must be a mapping.") + sections[name] = json_data( + {str(key): value for key, value in section.items()} + ) + + environment = sections["environment"] + task_meta = sections["task"] + metadata = sections["metadata"] + agent_config = sections["agent"] + verifier_config = sections["verifier"] + + raw_authors = task_meta.get("authors", []) + if raw_authors is None: + raw_authors = [] + if not isinstance(raw_authors, list): + raise TypeError("Harbor task [task].authors must be a list.") + authors_data: list[object] = list(raw_authors) + if not authors_data and isinstance(metadata.get("author_name"), str): + authors_data = [ + { + "name": metadata["author_name"], + "email": metadata.get("author_email"), + } + ] + authors = [Author.model_validate(author) for author in authors_data] + + raw_docker_image = environment.get("docker_image") + if raw_docker_image is not None and not isinstance(raw_docker_image, str): + raise TypeError("Harbor task [environment].docker_image must be a string.") + if ( + raw_docker_image is None + and (task_dir / "environment" / "Dockerfile").exists() + ): + raise ValueError( + f"Harbor task {task_dir.name!r} declares environment/Dockerfile " + "but no pullable [environment].docker_image. Building Harbor " + "Dockerfiles is not supported." + ) + if raw_docker_image is None and require_image: + raise ValueError( + f"Harbor task {task_dir.name!r} has no pullable " + "[environment].docker_image." + ) + + raw_agent_timeout = agent_config.get("timeout_sec") + raw_scoring_timeout = verifier_config.get("timeout_sec") + instruction = instruction_path.read_text().strip() + raw_name = task_meta.get("name") + if raw_name is not None and not isinstance(raw_name, str): + raise TypeError("Harbor task [task].name must be a string.") + raw_description = task_meta.get("description") + if raw_description is not None and not isinstance(raw_description, str): + raise TypeError("Harbor task [task].description must be a string.") + raw_difficulty = metadata.get("difficulty") + if raw_difficulty is not None and not isinstance(raw_difficulty, str): + raise TypeError("Harbor task [metadata].difficulty must be a string.") + raw_category = metadata.get("category") + if raw_category is not None and not isinstance(raw_category, str): + raise TypeError("Harbor task [metadata].category must be a string.") + raw_keywords = task_meta.get("keywords", []) + if raw_keywords is None: + raw_keywords = [] + if not isinstance(raw_keywords, list): + raise TypeError("Harbor task [task].keywords must be a list.") + keywords: list[str] = [] + for keyword in raw_keywords: + if not isinstance(keyword, str): + raise TypeError("Harbor task [task].keywords must contain strings.") + keywords.append(keyword) + raw_tags = metadata.get("tags", []) + if raw_tags is None: + raw_tags = [] + if not isinstance(raw_tags, list): + raise TypeError("Harbor task [metadata].tags must be a list.") + tags: list[str] = [] + for tag in raw_tags: + if not isinstance(tag, str): + raise TypeError("Harbor task [metadata].tags must contain strings.") + tags.append(tag) + return cls( + task_name=task_dir.name, + instruction=instruction, + task_dir=str(task_dir), + prompt=[vf.UserMessage(content=instruction)], + name=raw_name or task_dir.name, + description=raw_description, + image=raw_docker_image, + resources=cls.resources_from_environment(environment), + agent_timeout=( + None + if raw_agent_timeout is None + else parse_number(raw_agent_timeout, 0.0) + ), + scoring_timeout=( + None + if raw_scoring_timeout is None + else parse_number(raw_scoring_timeout, 0.0) + ), + keywords=keywords, + authors=authors, + difficulty=raw_difficulty, + category=raw_category, + tags=tags, + ) + + @staticmethod + def resources_from_environment(environment: vf.JsonData) -> vf.Resources: + cpu_cores = ( + None + if environment.get("cpus") is None + else parse_number(environment.get("cpus"), 0.0) + ) + memory_value = environment.get("memory_gb") + if memory_value is None and environment.get("memory_mb") is not None: + memory_gb = parse_number(environment.get("memory_mb"), 0.0) / 1024 + elif memory_value is None and environment.get("memory") is not None: + memory_gb = parse_gb(environment.get("memory"), 0.0) + else: + memory_gb = None if memory_value is None else parse_gb(memory_value, 0.0) + disk_value = environment.get("storage_gb") + if disk_value is None and environment.get("storage_mb") is not None: + disk_gb = parse_number(environment.get("storage_mb"), 0.0) / 1024 + elif disk_value is None and environment.get("storage") is not None: + disk_gb = parse_gb(environment.get("storage"), 0.0) + else: + disk_gb = None if disk_value is None else parse_gb(disk_value, 0.0) + gpu_count = ( + None + if environment.get("gpus") is None + else int(parse_number(environment.get("gpus"), 0.0)) + ) + return vf.Resources( + cpu_cores=cpu_cores, + memory_gb=memory_gb, + gpu_count=gpu_count, + disk_gb=disk_gb, + ) class HarborTaskset(vf.Taskset[HarborTasksetConfig]): + task_type = HarborTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: if split == "eval": return [] + root = self.task_root() + task_dirs = harbor_task_dirs(root, self.config.tasks) + tasks = [ + HarborTask.from_dir(task_dir, require_image=self.config.require_image) + for task_dir in task_dirs + ] + if not tasks: + raise ValueError(f"No valid Harbor tasks found in {root}.") + return tasks + + def to_task(self, task: vf.Task | vf.JsonData) -> vf.Task: + harbor_task = super().to_task(task) + if not isinstance(harbor_task, HarborTask): + raise TypeError("HarborTaskset expected a HarborTask.") + if harbor_task.task_dir: + return harbor_task + return HarborTask.from_dir( + self.task_root() / harbor_task.task_name, + require_image=self.config.require_image, + ) + + def task_root(self) -> Path: config = self.config - if config.dataset is not None: - cache_dir_path = ( - Path(str(config.cache_dir)).expanduser() if config.cache_dir else None + if config.source == "harbor": + cache_dir = ( + Path(config.cache_dir).expanduser() if config.cache_dir else None ) - root = download_harbor_dataset( + return download_harbor_dataset( config.dataset, - cache_dir=cache_dir_path, + cache_dir=cache_dir, refresh=config.refresh, ) - else: - bundle_package = config.bundle_package - if bundle_package is None: - raise RuntimeError( - "HarborTaskset() without a dataset requires bundle_package. " - "Pass dataset='...' to fetch from Harbor Hub, or set " - "bundle_package=__name__ from the package that owns tasks/." - ) - root = bundle_tasks_root(bundle_package) + if config.source == "package": + root = bundle_tasks_root(config.dataset) if not root.exists(): raise FileNotFoundError( - "HarborTaskset() without a dataset requires " - f"{bundle_package}/{TASKS_SUBDIR}/ to contain Harbor task " - f"directories. Not found: {root}" + "HarborTaskset package source must contain " + f"{TASKS_SUBDIR}/. Not found: {root}" ) - task_dirs = harbor_task_dirs(root, list(config.task_names or [])) - tasks: list[vf.JsonData] = [] - for task_dir in task_dirs: - task_toml_path = task_dir / "task.toml" - instruction_path = task_dir / "instruction.md" - with task_toml_path.open("rb") as f: - task_config = load_toml(f) - environment = task_config.get("environment", {}) or {} - assert isinstance(environment, dict) - agent_config = task_config.get("agent", {}) or {} - verifier_config = task_config.get("verifier", {}) or {} - if not isinstance(agent_config, dict): - raise TypeError(f"{task_toml_path} [agent] must be a mapping.") - if not isinstance(verifier_config, dict): - raise TypeError(f"{task_toml_path} [verifier] must be a mapping.") - instruction = instruction_path.read_text().strip() - task_remote_dir = config.task_dir.rstrip("/") or "/task" - sandbox = harbor_sandbox(HARBOR_DEFAULT_SANDBOX, config.sandbox) - sandbox = sandbox.model_copy( - update={ - "image": environment.get("docker_image") or sandbox.image, - "cpu_cores": parse_number( - environment.get("cpus"), sandbox.cpu_cores - ), - "memory_gb": parse_gb(environment.get("memory"), sandbox.memory_gb), - "disk_size_gb": parse_gb( - environment.get("storage"), sandbox.disk_size_gb - ), - "command_timeout": int( - parse_number( - agent_config.get("timeout_sec"), - sandbox.command_timeout or 900, - ) - ), - **( - {"network_access": bool(environment["allow_internet"])} - if "allow_internet" in environment - else {} - ), - } - ) - sandbox_data = sandbox.data(fill_defaults=False) - workdir = sandbox.workdir or "/app" - tasks.append( - { - "task_name": task_dir.name, - "instruction": instruction, - "task_toml": task_toml_path.read_text(), - "task_dir": str(task_dir), - "prompt": [{"role": "user", "content": instruction}], - "sandbox": sandbox_data, - "program": { - "files": { - f"{task_remote_dir}/instruction.md": { - "task": "instruction" - }, - f"{task_remote_dir}/task.toml": {"task": "task_toml"}, - }, - "env": { - "HARBOR_TASK_NAME": task_dir.name, - "HARBOR_TASK_DIR": task_remote_dir, - "HARBOR_INSTRUCTION_PATH": f"{task_remote_dir}/instruction.md", - "AGENT_WORKDIR": workdir, - **config.env, - }, - }, - "harbor": { - "task_dir": str(task_dir), - "task_name": task_dir.name, - "config": task_config, - "docker_image": environment.get("docker_image"), - "test_timeout": parse_number( - verifier_config.get("timeout_sec"), - config.verifier_timeout_seconds, - ), - }, - "info": { - "harbor": { - "task_name": task_dir.name, - "docker_image": environment.get("docker_image"), - } - }, - } - ) - assert tasks, f"No valid Harbor tasks found in {root}." - return tasks + return root + raise ValueError(f"Unknown Harbor source: {config.source!r}") @vf.reward(weight=1.0) - async def harbor_reward(self, task: vf.Task, state: vf.State) -> float: - if state.get("error") is not None: + async def harbor_reward(self, task: HarborTask, runtime: vf.Runtime) -> float: + tests_dir = Path(task.task_dir) / "tests" + try: + tests_archive = self.make_tar(tests_dir) + await runtime.write("/tmp/tests.tgz", tests_archive) + except OSError: return 0.0 - sandbox_id = state["sandbox_id"] - assert isinstance(sandbox_id, str) - harbor = task["harbor"] - assert isinstance(harbor, dict) - task_dir = Path(str(harbor["task_dir"])) - from prime_sandboxes import AsyncSandboxClient - - client = cast(SandboxClient, AsyncSandboxClient()) + extract = await runtime.run( + [ + "sh", + "-c", + "mkdir -p /logs/verifier /tests && tar -xzf /tmp/tests.tgz -C /tests", + ] + ) + if extract.returncode != 0: + return 0.0 + await runtime.run( + ["sh", "-c", "cd /tests && bash test.sh"], + timeout=task.scoring_timeout, + ) try: - await upload_harbor_tests(client, sandbox_id, task_dir) - test_timeout = int(parse_number(harbor.get("test_timeout"), 900)) - result = await client.run_background_job( - sandbox_id=sandbox_id, - command="bash test.sh", - working_dir="/tests", - timeout=test_timeout, - ) - state["harbor_tests"] = { - "returncode": result.exit_code, - "stdout": result.stdout or "", - "stderr": result.stderr or "", - } - reward_result = await client.execute_command( - sandbox_id=sandbox_id, - command=( - "if [ -s /logs/verifier/reward.txt ]; then " - "cat /logs/verifier/reward.txt; " - "elif [ -s /logs/verifier/reward.json ]; then " - "cat /logs/verifier/reward.json; fi" - ), - ) - except Exception as e: - state["harbor_error"] = str(e) + reward = (await runtime.read("/logs/verifier/reward.txt")).decode().strip() + except (OSError, RuntimeError, ValueError): return 0.0 - finally: - await client.aclose() - return parse_reward_text(str(reward_result.stdout or "").strip()) + return parse_reward_text(reward) + + @staticmethod + def make_tar(directory: Path) -> bytes: + buffer = io.BytesIO() + with tarfile.open(fileobj=buffer, mode="w:gz") as tar: + for item in sorted(directory.iterdir()): + tar.add(item, arcname=item.name) + return buffer.getvalue() def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: diff --git a/packages/tasksets/tasksets/nemo_gym.py b/packages/tasksets/tasksets/nemo_gym.py index 8d98e0c2f8..d1700be31c 100644 --- a/packages/tasksets/tasksets/nemo_gym.py +++ b/packages/tasksets/tasksets/nemo_gym.py @@ -1,15 +1,13 @@ import json from copy import deepcopy from pathlib import Path -from typing import TypeAlias, cast -import verifiers as vf -from verifiers.v1.utils.endpoint_utils import normalize_openai_responses_input +from pydantic import Field, TypeAdapter +import verifiers.v1 as vf +from verifiers.v1.utils.json_utils import json_data, json_value DEFAULT_NEMO_GYM_DATA_NAME = "example.jsonl" -ConfigData: TypeAlias = dict[str, object] -ConfigMap: TypeAlias = dict[str, object] -TaskRow: TypeAlias = dict[str, object] +_MESSAGES_ADAPTER = TypeAdapter(vf.Messages) def nemo_gym_package_root() -> Path: @@ -32,15 +30,15 @@ def resolve_nemo_gym_data_path( return path -def agent_ref_name(value: object) -> str | None: +def agent_ref_name(value: vf.JsonValue) -> str | None: if not isinstance(value, dict): return None - name = cast(ConfigMap, value).get("name") + name = value.get("name") return name if isinstance(name, str) and name else None class NeMoGymTasksetConfig(vf.TasksetConfig): - taskset_id: str | None = "nemo_gym" + id: str | None = "nemo_gym" nemo_env: str | None = None jsonl_path: str | None = None data_name: str = DEFAULT_NEMO_GYM_DATA_NAME @@ -48,12 +46,14 @@ class NeMoGymTasksetConfig(vf.TasksetConfig): limit: int | None = None -class NeMoGymTaskset(vf.Taskset[NeMoGymTasksetConfig]): - """Taskset adapter for NeMo Gym JSONL rows. +class NeMoGymTask(vf.Task, frozen=True): + nemo_gym_row: vf.JsonData + info: vf.JsonData = Field(default_factory=dict) + system_prompt: list[vf.JsonData] = Field(default_factory=list) + - Each task keeps the original NeMo Gym row under ``nemo_gym_row`` so the - harness can post it to the configured NeMo Gym agent unchanged. - """ +class NeMoGymTaskset(vf.Taskset[NeMoGymTasksetConfig]): + task_type = NeMoGymTask def jsonl_path(self) -> Path | None: raw_path = self.config.jsonl_path @@ -72,42 +72,45 @@ def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: jsonl_path = self.jsonl_path() if jsonl_path is None: raise ValueError("NeMoGymTaskset requires nemo_env=... or jsonl_path=...") - raw_rows: list[TaskRow] = [] + raw_rows: list[vf.JsonData] = [] with jsonl_path.open(encoding="utf-8") as f: for line in f: stripped = line.strip() if stripped: - raw_rows.append(cast(TaskRow, json.loads(stripped))) + raw_rows.append(json_data(json.loads(stripped))) if self.config.limit is not None: raw_rows = raw_rows[: self.config.limit] tasks = [ normalize_nemo_gym_task_row(row, index, agent_name=self.config.agent_name) for index, row in enumerate(raw_rows) ] - return cast(vf.Tasks, tasks) + return tasks def normalize_nemo_gym_task_row( - row: TaskRow, + row: vf.JsonData, index: int, *, agent_name: str | None = None, -) -> ConfigData: - nemo_row: ConfigData = deepcopy(dict(row)) +) -> vf.JsonData: + nemo_row: vf.JsonData = deepcopy(dict(row)) if agent_name and not agent_ref_name(nemo_row.get("agent_ref")): nemo_row["agent_ref"] = { "type": "responses_api_agents", "name": agent_name, } - task_row: ConfigData = deepcopy(nemo_row) - task_row["nemo_gym_row"] = nemo_row - task_row.setdefault("example_id", index) + task_row: vf.JsonData = {"nemo_gym_row": nemo_row} + for key in NeMoGymTask.model_fields: + if key != "nemo_gym_row" and key in row: + task_row[key] = deepcopy(row[key]) + task_row.setdefault("row_id", index) prompt, system_prompt = prompt_parts_from_nemo_gym_row(nemo_row) - task_row.setdefault("prompt", prompt) - if system_prompt: - task_row.setdefault("system_prompt", system_prompt) + if "prompt" not in task_row: + task_row["prompt"] = json_value(prompt) + if system_prompt and "system_prompt" not in task_row: + task_row["system_prompt"] = json_value(system_prompt) raw_info = task_row.get("info") - info = dict(cast(ConfigMap, raw_info)) if isinstance(raw_info, dict) else {} + info = dict(raw_info) if isinstance(raw_info, dict) else {} info.setdefault( "nemo_gym", { @@ -119,22 +122,34 @@ def normalize_nemo_gym_task_row( def prompt_parts_from_nemo_gym_row( - row: TaskRow, -) -> tuple[list[ConfigData], list[ConfigData]]: + row: vf.JsonData, +) -> tuple[list[vf.JsonData], list[vf.JsonData]]: create_params = row.get("responses_create_params") if not isinstance(create_params, dict): return [], [] - create_params = cast(ConfigMap, create_params) try: - messages = normalize_openai_responses_input(create_params.get("input")) + messages = normalize_responses_input(create_params.get("input")) except Exception: return [], [] - prompt: list[ConfigData] = [] - system_prompt: list[ConfigData] = [] + prompt: list[vf.JsonData] = [] + system_prompt: list[vf.JsonData] = [] for message in messages: - dumped = cast(ConfigData, message.model_dump(exclude_none=True)) + dumped = json_data(message.model_dump(exclude_none=True)) if getattr(message, "role", None) == "system": system_prompt.append(dumped) else: prompt.append(dumped) return prompt, system_prompt + + +def normalize_responses_input(value: vf.JsonValue) -> vf.Messages: + if isinstance(value, str): + return [vf.UserMessage(content=value)] + if isinstance(value, list): + raw_messages: list[vf.JsonData] = [] + for item in value: + if not isinstance(item, dict): + raise TypeError("responses_create_params.input must contain objects.") + raw_messages.append(json_data(item)) + return _MESSAGES_ADAPTER.validate_python(raw_messages) + raise TypeError("responses_create_params.input must be a string or message list.") diff --git a/packages/tasksets/tasksets/openenv.py b/packages/tasksets/tasksets/openenv.py index 9d0040cbd2..1c1e812336 100644 --- a/packages/tasksets/tasksets/openenv.py +++ b/packages/tasksets/tasksets/openenv.py @@ -1,27 +1,26 @@ import asyncio +import contextlib import importlib.util import json -from collections.abc import Awaitable, Callable, Sequence +import weakref +from collections.abc import Awaitable, Callable, Iterable from pathlib import Path from typing import Literal, Protocol, TypeAlias, cast -import verifiers as vf -from openenv.core.env_server.mcp_types import ( - CallToolAction, - CallToolObservation, - Tool as OpenEnvToolSpec, -) from openenv.core.generic_client import GenericEnvClient +from openenv.core.env_server.mcp_types import CallToolAction from openenv.core.mcp_client import MCPToolClient -from verifiers.utils.async_utils import maybe_await -from verifiers.utils.async_utils import maybe_call_with_named_args -from verifiers.utils.message_utils import get_messages, normalize_messages -from verifiers.utils.tool_utils import is_valid_tool_content_parts +from pydantic import Field, TypeAdapter + +import verifiers.v1 as vf +from verifiers.utils.async_utils import maybe_await, maybe_call_with_named_args from verifiers.v1.config import import_config_ref -from verifiers.v1.utils.serialization_utils import serializable +from verifiers.v1.utils.json_utils import json_data, json_value from tasksets.utils.openenv_utils import PrimeSandboxOpenEnvProvider +_MESSAGES_ADAPTER = TypeAdapter(vf.Messages) + OpenEnvPromptRenderer: TypeAlias = Callable[ ..., vf.PromptInput | Awaitable[vf.PromptInput] ] @@ -33,22 +32,29 @@ class OpenEnvResult(Protocol): done: bool -def default_openenv_prompt_renderer( - observation: object, -) -> vf.PromptInput: +class OpenEnvTool(Protocol): + @property + def name(self) -> str: ... + + @property + def description(self) -> str | None: ... + + +def default_openenv_prompt_renderer(observation: object) -> vf.PromptInput: if isinstance(observation, str): return [{"role": "user", "content": observation}] if isinstance(observation, dict): - observation_map = cast(vf.JsonData, observation) - messages = observation_map.get("messages") + observation_data = json_data(observation) + messages = observation_data.get("messages") if messages is not None: - assert isinstance(messages, list) - return cast(vf.PromptInput, messages) + if not isinstance(messages, list): + raise TypeError("OpenEnv observation messages must be a list.") + return _MESSAGES_ADAPTER.validate_python(messages) for key in ("prompt", "question", "instruction", "content", "text"): - value = observation_map.get(key) + value = observation_data.get(key) if isinstance(value, str) and value.strip(): return [{"role": "user", "content": value}] - return [{"role": "user", "content": json.dumps(serializable(observation))}] + return [{"role": "user", "content": json.dumps(observation_data)}] return [{"role": "user", "content": str(observation)}] @@ -59,6 +65,7 @@ class OpenEnvRuntimeConfig(vf.Config): port: int start_command: str contract: Literal["gym", "mcp"] + tools: list[vf.JsonData] = Field(default_factory=list) seed: int startup_timeout_seconds: int startup_poll_interval_seconds: float @@ -73,27 +80,24 @@ class OpenEnvRuntimeConfig(vf.Config): class OpenEnvBuildConfig(vf.Config): + app: str | None = None image: str + environment_id: str | None = None + image_status: str | None = None port: int = 8000 + schema_version: int | None = None start_command: str contract: Literal["gym", "mcp"] + tools: list[vf.JsonData] = Field(default_factory=list) class OpenEnvUserConfig(vf.UserConfig): - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - {"session": "taskset.objects.session"} - ) + scope: vf.Scope = "env" class OpenEnvTasksetConfig(vf.TasksetConfig): - taskset_id: str | None = "openenv" - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - {"session.task": "task"} - ) - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"session": "tasksets.openenv:OpenEnvSession"} - ) - user: OpenEnvUserConfig | None = OpenEnvUserConfig() + id: str | None = "openenv" + user: vf.UserConfig | None = OpenEnvUserConfig() prompt_renderer: str = "tasksets.openenv:default_openenv_prompt_renderer" openenv_project: str = "proj" num_train_examples: int = 100 @@ -104,78 +108,159 @@ class OpenEnvTasksetConfig(vf.TasksetConfig): health_request_timeout_seconds: float = 2.0 schema_request_timeout_seconds: float = 5.0 wait_for_creation_max_attempts: int = 20 - max_retries: int = 5 + max_retries: int = 12 base_delay: float = 0.5 backoff_factor: float = 2.0 - max_backoff_seconds: float = 30.0 + max_backoff_seconds: float = 60.0 jitter: float = 1e-3 +class OpenEnvTask(vf.Task, frozen=True): + openenv: vf.JsonData + info: vf.JsonData + + class OpenEnvSession: - def __init__(self, task: vf.Task): - task_config = task["openenv"] - assert isinstance(task_config, dict) - self.config = OpenEnvRuntimeConfig.model_validate(task_config) - self.provider: PrimeSandboxOpenEnvProvider | None = None + _startup_slots: weakref.WeakKeyDictionary[ + asyncio.AbstractEventLoop, asyncio.Semaphore + ] = weakref.WeakKeyDictionary() + + def __init__(self, config: OpenEnvRuntimeConfig, server: "OpenEnvServer"): + self.config = config + self.server = server self.client: GenericEnvClient | MCPToolClient | None = None - self.action_schema: vf.JsonData = {} + + @property + def action_schema(self) -> vf.JsonData: + return self.server.action_schema + + @classmethod + @contextlib.asynccontextmanager + async def startup_slot(cls): + loop = asyncio.get_running_loop() + semaphore = cls._startup_slots.get(loop) + if semaphore is None: + semaphore = asyncio.Semaphore(1) + cls._startup_slots[loop] = semaphore + async with semaphore: + yield async def start(self) -> GenericEnvClient | MCPToolClient: if self.client is not None: return self.client - config = self.config - provider = PrimeSandboxOpenEnvProvider(config) - self.provider = provider - client_class = MCPToolClient if config.contract == "mcp" else GenericEnvClient - try: - self.client = await client_class.from_docker_image( - config.image, - provider=provider, - port=config.port, - start_command=config.start_command, - env_vars={"ENABLE_WEB_INTERFACE": "false"}, - ) - schema = await asyncio.to_thread(provider.fetch_schema) - except Exception: - provider.stop_container() - raise - action_schema = schema.get("action", {}) - assert isinstance(action_schema, dict) - self.action_schema = cast(vf.JsonData, dict(action_schema)) + self.client = await self.server.open_client() return self.client async def reset(self) -> OpenEnvResult: client = await self.start() return cast(OpenEnvResult, await client.reset(seed=self.config.seed)) - async def list_tools(self) -> Sequence[OpenEnvToolSpec]: + async def step(self, action: vf.JsonData) -> OpenEnvResult: client = await self.start() - assert isinstance(client, MCPToolClient) - return await client.list_tools() + if not isinstance(client, GenericEnvClient): + raise RuntimeError("MCP OpenEnv tasks require an MCP tool server.") + return cast(OpenEnvResult, await client.step(action)) - async def call_tool(self, name: str, arguments: vf.JsonData) -> OpenEnvResult: + async def call_tool(self, name: str, input: vf.JsonData) -> OpenEnvResult: client = await self.start() - assert isinstance(client, MCPToolClient) - result = await client.step( - CallToolAction(tool_name=name, arguments=dict(arguments)) + if not isinstance(client, MCPToolClient): + raise RuntimeError("Gym OpenEnv tasks require assistant JSON actions.") + return cast( + OpenEnvResult, + await client.step(CallToolAction(tool_name=name, arguments=input)), ) - return cast(OpenEnvResult, result) - async def step(self, action: vf.JsonData) -> OpenEnvResult: - client = await self.start() - assert isinstance(client, GenericEnvClient) - return cast(OpenEnvResult, await client.step(action)) + async def tool_defs(self) -> list[vf.JsonData]: + return await self.server.tool_defs() async def close(self) -> None: if self.client is not None: await maybe_await(self.client.close) self.client = None + + +class OpenEnvServer: + def __init__(self, config: OpenEnvRuntimeConfig): + self.config = config + self.provider: PrimeSandboxOpenEnvProvider | None = None + self.base_url: str | None = None + self.action_schema: vf.JsonData = {} + self._lock = asyncio.Lock() + self._tools_lock = asyncio.Lock() + self._tool_defs: list[vf.JsonData] | None = None + + @property + def client_class(self) -> type[GenericEnvClient] | type[MCPToolClient]: + return MCPToolClient if self.config.contract == "mcp" else GenericEnvClient + + async def start(self) -> None: + if self.base_url is not None: + return + async with self._lock: + if self.base_url is not None: + return + config = self.config + provider = PrimeSandboxOpenEnvProvider(config) + schema: vf.JsonData | None = None + try: + async with OpenEnvSession.startup_slot(): + base_url = await asyncio.to_thread( + provider.start_container, + config.image, + port=config.port, + start_command=config.start_command, + env_vars={"ENABLE_WEB_INTERFACE": "false"}, + ) + await asyncio.to_thread(provider.wait_for_ready, base_url) + schema = await asyncio.to_thread(provider.fetch_schema) + action_schema = schema.get("action", {}) + if isinstance(action_schema, dict): + self.action_schema = json_data(action_schema) + self.provider = provider + self.base_url = base_url + except BaseException: + provider.stop_container() + raise + + async def open_client(self) -> GenericEnvClient | MCPToolClient: + await self.start() + if self.base_url is None: + raise RuntimeError("OpenEnv server did not start.") + client = self.client_class(base_url=self.base_url) + await client.connect() + return client + + async def tool_defs(self) -> list[vf.JsonData]: + if self.config.contract != "mcp": + return [] + if self._tool_defs is not None: + return list(self._tool_defs) + async with self._tools_lock: + if self._tool_defs is not None: + return list(self._tool_defs) + client = await self.open_client() + try: + if not isinstance(client, MCPToolClient): + self._tool_defs = [] + else: + self._tool_defs = openenv_tool_defs( + await client.list_tools(use_cache=False) + ) + finally: + await maybe_await(client.close) + return list(self._tool_defs) + + async def close(self) -> None: + if self.provider is not None: + self.provider.stop_container() self.provider = None + self.base_url = None + self.action_schema = {} + self._tool_defs = None class OpenEnvTaskset(vf.Taskset[OpenEnvTasksetConfig]): - def load_toolsets(self, config: OpenEnvTasksetConfig) -> vf.Toolsets: - return {"openenv": vf.Toolset(scope="rollout", handler=self.call_tool)} + task_type = OpenEnvTask def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: if split == "eval": @@ -188,33 +273,46 @@ def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: first_seed=self.config.seed, ) - def openenv_tasks( - self, - *, - num_examples: int, - first_seed: int, - ) -> vf.Tasks: - config = self.config + def openenv_tasks(self, *, num_examples: int, first_seed: int) -> vf.Tasks: if num_examples <= 0: return [] - project = Path(config.openenv_project).expanduser() - if not project.is_absolute(): - spec = importlib.util.find_spec(type(config).__module__) - assert spec is not None - assert spec.origin is not None - project = Path(spec.origin).parent / project - project = project.resolve() + project = self.openenv_project_path() build = OpenEnvBuildConfig.model_validate( json.loads((project / ".build.json").read_text()) ) - runtime_config = OpenEnvRuntimeConfig( + return [ + { + "prompt": [], + "openenv": self.runtime_config( + project=project, build=build, seed=first_seed + index + ).model_dump(mode="json"), + "info": {"seed": first_seed + index, "contract": build.contract}, + } + for index in range(num_examples) + ] + + def openenv_project_path(self) -> Path: + project = Path(self.config.openenv_project).expanduser() + if project.is_absolute(): + return project.resolve() + spec = importlib.util.find_spec(type(self.config).__module__) + if spec is None or spec.origin is None: + return project.resolve() + return (Path(spec.origin).parent / project).resolve() + + def runtime_config( + self, *, project: Path, build: OpenEnvBuildConfig, seed: int + ) -> OpenEnvRuntimeConfig: + config = self.config + return OpenEnvRuntimeConfig( openenv_project=str(project), prompt_renderer=config.prompt_renderer, image=build.image, port=build.port, start_command=build.start_command, contract=build.contract, - seed=config.seed, + tools=list(build.tools), + seed=seed, startup_timeout_seconds=config.startup_timeout_seconds, startup_poll_interval_seconds=config.startup_poll_interval_seconds, health_request_timeout_seconds=config.health_request_timeout_seconds, @@ -226,134 +324,279 @@ def openenv_tasks( max_backoff_seconds=config.max_backoff_seconds, jitter=config.jitter, ) - return [ - { - "prompt": [ - { - "role": "user", - "content": "OpenEnv rollout is initializing.", - } - ], - "openenv": runtime_config.model_copy( - update={"seed": first_seed + index} - ).model_dump(), - "info": {"seed": first_seed + index, "contract": build.contract}, - } - for index in range(num_examples) - ] - @vf.setup - async def setup_openenv(self, task: vf.Task, state: vf.State) -> None: - session = await self.get_object("session", task, state) - assert isinstance(session, OpenEnvSession) - result = await session.reset() - config = session.config - if config.contract == "mcp": - for tool in await session.list_tools(): - schema = tool.input_schema or {"type": "object", "properties": {}} - tool_def = vf.Tool( - name=tool.name, - description=tool.description, - parameters={str(key): value for key, value in schema.items()}, - ) - state.add_tool("openenv", tool_def) - state["openenv_done"] = bool(result.done) - renderer = import_config_ref(config.prompt_renderer) - assert callable(renderer) - rendered = await maybe_call_with_named_args( - cast(OpenEnvPromptRenderer, renderer), - observation=result.observation, - context="reset", - action_schema=dict(session.action_schema), - contract=config.contract, - seed=config.seed, + @vf.reward(weight=1.0) + async def openenv_reward(self, state: vf.State) -> float: + return sum(float(turn.reward or 0.0) for turn in state.transcript) + + +async def render_messages( + config: OpenEnvRuntimeConfig, + session: OpenEnvSession, + observation: object, + context: str, +) -> list[vf.JsonData]: + renderer = import_config_ref(config.prompt_renderer) + if not callable(renderer): + raise TypeError("OpenEnv prompt_renderer must be callable.") + rendered = await maybe_call_with_named_args( + cast(OpenEnvPromptRenderer, renderer), + observation=observation, + context=context, + action_schema=dict(session.action_schema), + contract=config.contract, + seed=config.seed, + ) + messages = ( + [vf.UserMessage(content=rendered)] + if isinstance(rendered, str) + else _MESSAGES_ADAPTER.validate_python(rendered) + ) + return [ + json_data(message.model_dump(mode="json", exclude_none=True)) + for message in messages + ] + + +def latest_assistant_json(raw_completion: list[dict]) -> vf.JsonData: + messages = [message for message in raw_completion if isinstance(message, dict)] + assistant_messages = [ + message for message in messages if message.get("role") == "assistant" + ] + if not assistant_messages: + return {} + content = assistant_messages[-1].get("content") + text = content if isinstance(content, str) else json.dumps(content) + action = json.loads(text.strip()) + return json_data(action, context="OpenEnv assistant action") + + +def mcp_tool_content(observation: object) -> vf.JsonValue: + model_dump = getattr(observation, "model_dump", None) + if callable(model_dump): + observation = model_dump() + if not isinstance(observation, dict): + return json_value(observation) + observation_data = json_data(observation) + if observation_data.get("error") is not None: + return {"error": json_value(observation_data.get("error"))} + value = observation_data.get("result") + data = getattr(value, "data", None) + if data is not None: + return json_value(data) + if isinstance(value, dict) and "data" in value: + return json_value(value["data"]) + return json_value(value) + + +def openenv_tool_defs(tools: Iterable[OpenEnvTool]) -> list[vf.JsonData]: + tool_defs: list[vf.JsonData] = [] + for tool in tools: + name = getattr(tool, "name", None) + if not isinstance(name, str) or not name: + raise TypeError("OpenEnv MCP tools must have non-empty names.") + description = getattr(tool, "description", "") or "" + schema = getattr(tool, "input_schema", None) + if schema is None: + schema = getattr(tool, "inputSchema", None) + parameters = ( + json_data(schema, context=f"OpenEnv tool {name} schema") + if isinstance(schema, dict) + else {"type": "object", "properties": {}} ) - state["prompt"] = normalize_messages( - cast(vf.PromptInput, rendered), field_name="openenv" + tool_defs.append( + json_data( + { + "name": name, + "description": str(description), + "parameters": parameters, + }, + context=f"OpenEnv tool {name}", + ) ) + return tool_defs - @vf.stop - async def openenv_done(self, state: vf.State) -> bool: - return bool(state.get("openenv_done")) - @vf.reward(weight=1.0) - async def openenv_reward(self, state: vf.State) -> float: - return state.total_step_reward() +def result_payload( + *, + messages: list[vf.JsonData], + result: OpenEnvResult, + include_reward: bool, +) -> vf.JsonData: + done = openenv_result_done(result) + payload: dict[str, object] = { + "messages": messages, + "observation": json_value(result.observation), + "openenv_done": done, + } + if include_reward: + payload["reward"] = openenv_result_reward(result) + if done: + payload["stop_condition"] = "openenv_done" + return json_data(payload) - async def call_tool( - self, task: vf.Task, state: vf.State, tool: vf.Tool, arguments: vf.JsonData - ) -> vf.MessageContent: - session = await self.get_object("session", task, state) - assert isinstance(session, OpenEnvSession) - result = await session.call_tool(tool.name, arguments) - state.add_step_reward(result.reward) - state["openenv_done"] = bool(result.done) - if result.done: - state.stop("openenv_done") - observation = result.observation - if isinstance(observation, CallToolObservation): - content: object = ( - {"error": observation.error.message} - if observation.error is not None - else observation.result.data - ) - elif isinstance(observation, dict): - observation_map = cast(vf.JsonData, observation) - if observation_map.get("error") is not None: - content = {"error": observation_map.get("error")} - else: - result_value = observation_map["result"] - assert isinstance(result_value, dict) - result_data = cast(vf.JsonData, result_value) - content = result_data["data"] - else: - content = observation - if is_valid_tool_content_parts(content): - return cast(vf.MessageContent, content) - if isinstance(content, str): - return content - return json.dumps(content, ensure_ascii=True) + +def openenv_result_done(result: OpenEnvResult) -> bool: + observation_done = getattr(result.observation, "done", None) + return bool(result.done or observation_done) + + +def openenv_result_reward(result: OpenEnvResult) -> float: + reward = result.reward + if reward is None: + reward = getattr(result.observation, "reward", None) + return float(reward or 0.0) class OpenEnvUser(vf.User[OpenEnvUserConfig]): - async def get_response( + sessions: dict[str, OpenEnvSession] + servers: dict[str, OpenEnvServer] + + def start(self) -> None: + self.sessions = {} + self.servers = {} + + def stop(self) -> None: + if self.sessions or self.servers: + asyncio.run(self.close_resources()) + self.sessions = {} + self.servers = {} + + async def close_resources(self) -> None: + sessions = list(self.sessions.values()) + servers = list(self.servers.values()) + self.sessions.clear() + self.servers.clear() + await asyncio.gather(*(session.close() for session in sessions)) + await asyncio.gather(*(server.close() for server in servers)) + + def session_for(self, state_id: str) -> OpenEnvSession: + try: + return self.sessions[state_id] + except KeyError as exc: + raise RuntimeError("OpenEnv setup has not started.") from exc + + def server_for(self, config: OpenEnvRuntimeConfig) -> OpenEnvServer: + key = self.server_key(config) + server = self.servers.get(key) + if server is None: + server = OpenEnvServer(config) + self.servers[key] = server + return server + + @staticmethod + def server_key(config: OpenEnvRuntimeConfig) -> str: + data = config.model_dump(mode="json") + data.pop("seed", None) + return json.dumps(data, sort_keys=True, separators=(",", ":")) + + @vf.tool( + hidden=True, + args={"state_id": "state.id", "openenv": "task.openenv"}, + sets={ + "observation": "state.extras.openenv.observation", + "openenv_done": "state.extras.openenv.done", + "finished": "state.is_completed", + "stop_condition": "state.stop_condition", + }, + ) + async def setup(self, state_id: str, openenv: vf.JsonData) -> vf.JsonData: + config = OpenEnvRuntimeConfig.model_validate(openenv) + existing = self.sessions.get(state_id) + if existing is not None: + await existing.close() + session = OpenEnvSession(config, self.server_for(config)) + self.sessions[state_id] = session + if config.contract == "mcp": + tools = list(config.tools) or await session.tool_defs() + return json_data( + { + "observation": {}, + "openenv_done": False, + "tools": tools, + } + ) + result = await session.reset() + payload: dict[str, object] = { + "observation": json_value(result.observation), + "openenv_done": openenv_result_done(result), + } + if openenv_result_done(result): + payload["finished"] = True + payload["stop_condition"] = "openenv_done" + return json_data(payload) + + @vf.user( + args={ + "state_id": "state.id", + "observation": "state.extras.openenv.observation", + "completion": "state.completion", + }, + sets={ + "observation": "state.extras.openenv.observation", + "openenv_done": "state.extras.openenv.done", + "reward": "state.transcript.last.reward", + "stop_condition": "state.stop_condition", + }, + ) + async def respond( self, - task: vf.Task, - state: vf.State, - messages: Sequence[vf.Message], - session: OpenEnvSession | None = None, - ) -> list[vf.UserMessage]: - assert session is not None + state_id: str, + observation: vf.JsonValue, + completion: list[vf.JsonData], + ) -> vf.JsonData: + session = self.session_for(state_id) config = session.config + if not completion: + return json_data( + { + "messages": await render_messages( + config, session, observation, "reset" + ) + } + ) if config.contract == "mcp": - return [] - assistant_messages = get_messages(messages, role="assistant") - last_message = assistant_messages[-1] if assistant_messages else None - text = str(last_message.content or "").strip() if last_message else "" - action = json.loads(text) - assert isinstance(action, dict) - result = await session.step(cast(vf.JsonData, action)) - state.add_step_reward(result.reward) - state["openenv_done"] = bool(result.done) - if result.done: - state.stop("openenv_done") - renderer = import_config_ref(config.prompt_renderer) - assert callable(renderer) - rendered = await maybe_call_with_named_args( - cast(OpenEnvPromptRenderer, renderer), - observation=result.observation, - context="step", - action_schema=dict(session.action_schema), - contract=config.contract, - seed=config.seed, + return json_data( + {"messages": [], "stop_condition": "openenv_no_tool_calls"} + ) + action = latest_assistant_json(completion) + result = await session.step(action) + return result_payload( + messages=await render_messages(config, session, result.observation, "step"), + result=result, + include_reward=True, ) - response: list[vf.UserMessage] = [] - for message in normalize_messages( - cast(vf.PromptInput, rendered), field_name="openenv" - ): - assert isinstance(message, vf.UserMessage) - response.append(message) - return response + + @vf.tool( + hidden=True, + args={"state_id": "state.id"}, + sets={ + "observation": "state.extras.openenv.observation", + "openenv_done": "state.extras.openenv.done", + "reward": "state.transcript.last.reward", + "finished": "state.is_completed", + "stop_condition": "state.stop_condition", + }, + ) + async def call_tool( + self, state_id: str, name: str, input: vf.JsonData + ) -> vf.JsonData: + """Call an OpenEnv MCP tool by name with JSON input.""" + session = self.session_for(state_id) + result = await session.call_tool(name, input) + content = mcp_tool_content(result.observation) + payload: dict[str, object] = { + "content": content + if isinstance(content, str) + else json.dumps(json_value(content)), + "observation": json_value(result.observation), + "openenv_done": openenv_result_done(result), + } + payload["reward"] = openenv_result_reward(result) + if openenv_result_done(result): + payload["finished"] = True + payload["stop_condition"] = "openenv_done" + return json_data(payload) def load_taskset(config: OpenEnvTasksetConfig) -> OpenEnvTaskset: diff --git a/packages/tasksets/tasksets/openreward.py b/packages/tasksets/tasksets/openreward.py index b71c12b30e..55aad2b142 100644 --- a/packages/tasksets/tasksets/openreward.py +++ b/packages/tasksets/tasksets/openreward.py @@ -11,8 +11,9 @@ TextBlock as OpenRewardTextBlock, ToolOutput as OpenRewardToolOutput, ) -import verifiers as vf -from verifiers.v1.utils.serialization_utils import serializable + +import verifiers.v1 as vf +from verifiers.v1.utils.json_utils import json_data, json_value class OpenRewardSplit(Protocol): @@ -33,14 +34,24 @@ def get_task_range( ) -> Iterable[OpenRewardTask]: ... +class OpenRewardTool(Protocol): + @property + def name(self) -> str: ... + + @property + def description(self) -> str | None: ... + + @property + def input_schema(self) -> vf.JsonValue | None: ... + + +class OpenRewardUserConfig(vf.UserConfig): + pass + + class OpenRewardTasksetConfig(vf.TasksetConfig): - taskset_id: str | None = "openreward" - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - {"session.task": "task"} - ) - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"session": "tasksets.openreward:OpenRewardSession"} - ) + id: str | None = "openreward" + user: vf.UserConfig | None = OpenRewardUserConfig() environment: str variant: str | None = None base_url: str | None = None @@ -49,8 +60,12 @@ class OpenRewardTasksetConfig(vf.TasksetConfig): num_eval_examples: int = 0 +class OpenRewardVFTask(vf.Task, frozen=True): + openreward: vf.JsonData + + class OpenRewardSession: - def __init__(self, task: vf.Task): + def __init__(self, task: vf.JsonData): self.task = task self.client: OpenReward | None = None self.session_context: OpenRewardAPISession | None = None @@ -60,24 +75,37 @@ async def start(self) -> OpenRewardAPISession: if self.session is not None: return self.session spec = self.task["openreward"] - assert isinstance(spec, dict) - task_data = spec["task"] - assert isinstance(task_data, dict) - task_spec = task_data["task_spec"] - assert isinstance(task_spec, dict) + if not isinstance(spec, dict): + raise TypeError("OpenReward task requires openreward config.") + spec_data = spec + task_data = spec_data.get("task") + if not isinstance(task_data, dict): + raise TypeError("OpenReward task requires task data.") + task_spec = task_data.get("task_spec") + if not isinstance(task_spec, dict): + raise TypeError("OpenReward task_spec must be a mapping.") + variant = spec_data.get("variant") + if variant is not None and not isinstance(variant, str): + raise TypeError("OpenReward variant must be a string or null.") + base_url = spec_data.get("base_url") + if base_url is not None and not isinstance(base_url, str): + raise TypeError("OpenReward base_url must be a string or null.") + namespace = task_data.get("namespace") + if namespace is not None and not isinstance(namespace, str): + raise TypeError("OpenReward namespace must be a string or null.") client = OpenReward() self.client = client environment = await asyncio.to_thread( client.environments.get, - name=str(spec["environment"]), - variant=cast(str | None, spec["variant"]), - base_url=cast(str | None, spec["base_url"]), + name=str(spec_data["environment"]), + variant=variant, + base_url=base_url, ) self.session_context = environment.session( task=OpenRewardTask( server_name=str(task_data["server_name"]), environment_name=str(task_data["environment_name"]), - namespace=cast(str | None, task_data["namespace"]), + namespace=namespace, task_spec=cast( OpenRewardJSONObject, {str(key): value for key, value in task_spec.items()}, @@ -91,20 +119,17 @@ async def prompt(self) -> object: session = await self.start() return await asyncio.to_thread(session.get_prompt) - async def tool_specs(self) -> Iterable[object]: + async def call_tool(self, name: str, input: vf.JsonData) -> OpenRewardToolOutput: session = await self.start() return cast( - Iterable[object], await asyncio.to_thread(session.list_tools, "openai") + OpenRewardToolOutput, + await asyncio.to_thread(session.call_tool, name, input), ) - async def call_tool( - self, name: str, arguments: vf.JsonData - ) -> OpenRewardToolOutput: + async def tool_defs(self) -> list[vf.JsonData]: session = await self.start() - result = await asyncio.to_thread( - session.call_tool, name, cast(OpenRewardJSONObject, dict(arguments)) - ) - return cast(OpenRewardToolOutput, result) + tools = await asyncio.to_thread(session.list_tools) + return openreward_tool_defs(tools) def content(self, blocks: object) -> vf.MessageContent: block_list = ( @@ -113,8 +138,9 @@ def content(self, blocks: object) -> vf.MessageContent: else [blocks] ) if all(isinstance(block, OpenRewardTextBlock) for block in block_list): - text_blocks = cast(list[OpenRewardTextBlock], block_list) - return "\n".join(block.text for block in text_blocks) + return "\n".join( + block.text for block in cast(list[OpenRewardTextBlock], block_list) + ) content: list[vf.JsonData] = [] for block in block_list: if isinstance(block, OpenRewardTextBlock): @@ -129,7 +155,7 @@ def content(self, blocks: object) -> vf.MessageContent: } ) else: - assert False, f"Unexpected OpenReward block: {block!r}" + raise TypeError(f"Unexpected OpenReward block: {block!r}") return cast(vf.MessageContent, content) async def close(self) -> None: @@ -142,16 +168,51 @@ async def close(self) -> None: self.client = None +def openreward_tool_defs( + tools: Iterable[OpenRewardTool | vf.JsonData], +) -> list[vf.JsonData]: + tool_defs: list[vf.JsonData] = [] + for tool in tools: + if isinstance(tool, dict): + tool_data = cast(vf.JsonData, tool) + name = tool_data.get("name") + description = tool_data.get("description", "") + schema = tool_data.get("input_schema") + else: + tool_spec = cast(OpenRewardTool, tool) + name = tool_spec.name + description = tool_spec.description + schema = tool_spec.input_schema + if not isinstance(name, str) or not name: + raise TypeError("OpenReward tools must have non-empty names.") + description = description or "" + parameters = ( + json_data(schema, context=f"OpenReward tool {name} schema") + if isinstance(schema, dict) + else {"type": "object", "properties": {}} + ) + tool_defs.append( + json_data( + { + "name": name, + "description": str(description), + "parameters": parameters, + }, + context=f"OpenReward tool {name}", + ) + ) + return tool_defs + + class OpenRewardTaskset(vf.Taskset[OpenRewardTasksetConfig]): - def load_toolsets(self, config: OpenRewardTasksetConfig) -> vf.Toolsets: - return {"openreward": vf.Toolset(scope="rollout", handler=self.call_tool)} + task_type = OpenRewardVFTask def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: if split == "eval": if self.config.num_eval_examples <= 0: return [] return self.openreward_test_tasks( - num_examples=self.config.num_eval_examples, + num_examples=self.config.num_eval_examples ) return self.openreward_tasks( task_split=self.config.split, @@ -208,11 +269,7 @@ def openreward_source_tasks( ) -> Iterable[OpenRewardTask]: if num_examples is None: return environment.list_tasks(split=task_split) - return environment.get_task_range( - split=task_split, - start=0, - stop=num_examples, - ) + return environment.get_task_range(split=task_split, start=0, stop=num_examples) def openreward_task_records( self, @@ -222,96 +279,121 @@ def openreward_task_records( config = self.config data: list[vf.JsonData] = [] for task in tasks: - task_spec = serializable(task.task_spec) - assert isinstance(task_spec, dict) + task_spec = json_value(task.task_spec) + if not isinstance(task_spec, dict): + raise TypeError("OpenReward task_spec must serialize to a mapping.") data.append( - { - "prompt": [ - { - "role": "user", - "content": "OpenReward rollout is initializing.", - } - ], - "openreward": { - "environment": config.environment, - "variant": config.variant, - "base_url": config.base_url, - "split": task_split, - "task": { - "server_name": task.server_name, - "environment_name": task.environment_name, - "namespace": task.namespace, - "task_spec": { - str(key): value for key, value in task_spec.items() + json_data( + { + "prompt": [], + "openreward": { + "environment": config.environment, + "variant": config.variant, + "base_url": config.base_url, + "split": task_split, + "task": { + "server_name": task.server_name, + "environment_name": task.environment_name, + "namespace": task.namespace, + "task_spec": { + str(key): value for key, value in task_spec.items() + }, }, }, }, - } + context="OpenReward task record", + ) ) return data - @vf.setup - async def setup_openreward(self, task: vf.Task, state: vf.State) -> None: - session = await self.get_object("session", task, state) - assert isinstance(session, OpenRewardSession) - prompt = await session.prompt() - state["prompt"] = [vf.UserMessage(content=session.content(prompt))] - for tool_spec in await session.tool_specs(): - tool_value = serializable(tool_spec) - assert isinstance(tool_value, dict) - tool_data = cast(vf.ConfigData, tool_value) - function_data = tool_data.get("function") - data = ( - cast(vf.ConfigData, function_data) - if isinstance(function_data, dict) - else tool_data - ) - name = data["name"] - assert isinstance(name, str) - parameters = ( - data.get("parameters") - or data.get("input_schema") - or data.get("inputSchema") - or {"type": "object", "properties": {}} - ) - assert isinstance(parameters, dict) - state.add_tool( - "openreward", - vf.Tool( - name=name, - description=str(data.get("description") or ""), - parameters={str(key): value for key, value in parameters.items()}, - ), - ) - - @vf.stop - async def openreward_done(self, state: vf.State) -> bool: - return bool(state.get("openreward_finished")) - @vf.reward(weight=1.0) async def openreward_reward(self, state: vf.State) -> float: - return state.total_step_reward() - - async def call_tool( - self, task: vf.Task, state: vf.State, tool: vf.Tool, arguments: vf.JsonData - ) -> vf.MessageContent: - session = await self.get_object("session", task, state) - assert isinstance(session, OpenRewardSession) - tool_arguments = serializable(arguments) - assert isinstance(tool_arguments, dict) - result = await session.call_tool( - tool.name, - cast( - vf.JsonData, {str(key): value for key, value in tool_arguments.items()} - ), + return sum(float(turn.reward or 0.0) for turn in state.transcript) + + +class OpenRewardUser(vf.User[OpenRewardUserConfig]): + session: OpenRewardSession | None + + def start(self) -> None: + self.session = None + + def stop(self) -> None: + if self.session is not None: + asyncio.run(self.session.close()) + self.session = None + + @vf.tool( + hidden=True, + args={"task": "task"}, + sets={ + "prompt": "state.extras.openreward.prompt", + }, + ) + async def setup(self, task: vf.JsonData) -> vf.JsonData: + self.session = OpenRewardSession(task) + prompt = await self.session.prompt() + return json_data( + { + "prompt": json_value(self.session.content(prompt)), + "tools": await self.session.tool_defs(), + } ) - state.add_step_reward(result.reward) - state["openreward_finished"] = result.finished - if result.finished: - state.stop("openreward_done") - if result.metadata is not None: - state["openreward_metadata"] = serializable(result.metadata) - return session.content(result.blocks) + + @vf.user( + args={ + "prompt": "state.extras.openreward.prompt", + "completion": "state.completion", + }, + sets={ + "stop_condition": "state.stop_condition", + }, + ) + async def respond( + self, + prompt: vf.JsonValue, + completion: list[vf.JsonData], + ) -> vf.JsonData: + if self.session is None: + raise RuntimeError("OpenReward setup has not started.") + if completion: + return json_data( + {"messages": [], "stop_condition": "openreward_waiting_for_tools"} + ) + return json_data( + { + "messages": [ + json_data( + vf.UserMessage( + content=cast(vf.MessageContent, prompt) + ).model_dump(mode="json") + ) + ], + } + ) + + @vf.tool( + hidden=True, + sets={ + "reward": "state.transcript.last.reward", + "finished": "state.is_completed", + "stop_condition": "state.stop_condition", + }, + ) + async def call_tool(self, name: str, input: vf.JsonData) -> vf.JsonData: + """Call an OpenReward task tool by name with JSON input.""" + if self.session is None: + raise RuntimeError("OpenReward session has not started.") + output = await self.session.call_tool(name, input) + content = self.session.content(output.blocks) + payload: dict[str, object] = { + "content": content, + } + if output.reward is not None: + payload["reward"] = float(output.reward) + if output.finished: + payload["finished"] = True + payload["stop_condition"] = "openreward_finished" + return json_data(payload) def load_taskset(config: OpenRewardTasksetConfig) -> OpenRewardTaskset: diff --git a/packages/tasksets/tasksets/replay.py b/packages/tasksets/tasksets/replay.py index 22fad90f1c..fb0f395d33 100644 --- a/packages/tasksets/tasksets/replay.py +++ b/packages/tasksets/tasksets/replay.py @@ -1,13 +1,16 @@ import json from pathlib import Path -from typing import ClassVar, cast +from typing import ClassVar from datasets import load_dataset +from pydantic import TypeAdapter -import verifiers as vf +import verifiers.v1 as vf +from verifiers.v1.utils.json_utils import json_data DATA_DIR_FIELD = "data_dir" DATA_FILE_SUFFIX = ".jsonl" +_MESSAGES_ADAPTER = TypeAdapter(vf.Messages) class ReplayTasksetConfig(vf.TasksetConfig): @@ -15,7 +18,12 @@ class ReplayTasksetConfig(vf.TasksetConfig): data_dir: str | None = None +class ReplayTask(vf.Task, frozen=True): + messages: vf.Messages + + class ReplayTaskset(vf.Taskset[ReplayTasksetConfig]): + task_type = ReplayTask data_dir: ClassVar[str | None] = None def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: @@ -31,7 +39,7 @@ def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: def hf_tasks(self, dataset: str) -> list[vf.JsonData]: rows = load_dataset(dataset, split="train") - return [replay_task_record(dict(row)) for row in rows] + return [self.task_record(json_data(dict(row))) for row in rows] def local_tasks(self) -> list[vf.JsonData]: data_dir = self.load_data_dir() @@ -55,7 +63,7 @@ def local_tasks(self) -> list[vf.JsonData]: raise TypeError( f"{item.name}:{line_number} must contain one JSON object." ) - tasks.append(replay_task_record(record)) + tasks.append(self.task_record(json_data(record))) if not tasks: raise FileNotFoundError( f"{DATA_DIR_FIELD} must contain at least one JSONL record." @@ -66,29 +74,27 @@ def load_data_dir(self) -> Path | None: data_dir = self.config.data_dir or self.data_dir return Path(data_dir).expanduser() if data_dir is not None else None - -def replay_task_record(record: dict[str, object]) -> vf.JsonData: - messages = replay_messages(record) - if not any(message.role == "assistant" for message in messages): - raise ValueError("Replay task messages must contain an assistant message.") - data = dict(record) - data["messages"] = [ - cast(vf.JsonData, message.model_dump(mode="json", exclude_none=True)) - for message in messages - ] - return cast(vf.JsonData, data) - - -def replay_messages(record: dict[str, object]) -> vf.Messages: - messages = record.get("messages") - if not isinstance(messages, list): - raise TypeError("Replay task messages must be a list.") - raw_messages: list[dict[str, object]] = [] - for message in messages: - if not isinstance(message, dict): - raise TypeError("Replay task messages must contain JSON objects.") - raw_messages.append(cast(dict[str, object], message)) - return vf.get_messages(raw_messages) + def task_record(self, record: vf.JsonData) -> vf.JsonData: + messages = self.messages(record) + if not any(message.role == "assistant" for message in messages): + raise ValueError("Replay task messages must contain an assistant message.") + data = dict(record) + data["messages"] = [ + json_data(message.model_dump(mode="json", exclude_none=True)) + for message in messages + ] + return json_data(data) + + def messages(self, record: vf.JsonData) -> vf.Messages: + messages = record.get("messages") + if not isinstance(messages, list): + raise TypeError("Replay task messages must be a list.") + raw_messages: list[vf.JsonData] = [] + for message in messages: + if not isinstance(message, dict): + raise TypeError("Replay task messages must contain JSON objects.") + raw_messages.append(json_data(message)) + return _MESSAGES_ADAPTER.validate_python(raw_messages) def load_taskset(config: ReplayTasksetConfig) -> ReplayTaskset: diff --git a/packages/tasksets/tasksets/textarena.py b/packages/tasksets/tasksets/textarena.py index fdd918f1a5..dfacaf592f 100644 --- a/packages/tasksets/tasksets/textarena.py +++ b/packages/tasksets/tasksets/textarena.py @@ -1,11 +1,12 @@ -import asyncio -import random import re +import random from collections.abc import Sequence -from typing import Generic, TypeVar -from typing import Protocol, cast +from typing import Generic, Protocol, TypeVar, cast + +from pydantic import BaseModel -import verifiers as vf +import verifiers.v1 as vf +from verifiers.v1.utils.json_utils import json_data try: import nltk @@ -33,43 +34,36 @@ def step(self, action: str) -> object: ... class TextArenaUserConfig(vf.UserConfig): - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"session": "tasksets.textarena:TextArenaSession"} - ) + pass class TextArenaTasksetConfig(vf.TasksetConfig): - taskset_id: str | None = "textarena" + id: str | None = "textarena" game: str - user: TextArenaUserConfig | None = TextArenaUserConfig() + user: vf.UserConfig | None = TextArenaUserConfig() num_train_examples: int = 2000 num_eval_examples: int = 20 seed: int = 0 answer_state_key: str -TextArenaConfigT = TypeVar("TextArenaConfigT", bound=TextArenaTasksetConfig) +class TextArenaSpec(BaseModel, extra="forbid"): + game: str + answer_state_key: str -def _content_text(content: object) -> str: - if isinstance(content, str): - return content - if isinstance(content, Sequence) and not isinstance( - content, (str, bytes, bytearray) - ): - chunks: list[str] = [] - for part in content: - if isinstance(part, vf.TextContentPart): - chunks.append(part.text) - elif isinstance(part, dict): - text = cast(dict[str, object], part).get("text") - if isinstance(text, str): - chunks.append(text) - return "\n".join(chunks) - return "" +class TextArenaTask(vf.Task, frozen=True): + answer: str + textarena: TextArenaSpec + + +ConfigT = TypeVar("ConfigT", bound=TextArenaTasksetConfig) -class TextArenaTaskset(vf.Taskset[TextArenaConfigT], Generic[TextArenaConfigT]): +class TextArenaTaskset(vf.Taskset[ConfigT], Generic[ConfigT]): + config: ConfigT + task_type = TextArenaTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: if split == "eval": return self.textarena_tasks( @@ -96,29 +90,21 @@ def textarena_tasks( assert isinstance(template, ta.Env) template.reset(num_players=1) _, initial_prompt = template.get_observation() - assert isinstance(initial_prompt, str) - assert initial_prompt - words = template.word_list - if isinstance(words, dict): - words = [ - word - for values in words.values() - for word in (values if isinstance(values, (list, tuple)) else [values]) - ] - word_list = [str(word) for word in words] - assert word_list + if not isinstance(initial_prompt, str) or not initial_prompt: + raise ValueError("TextArena initial prompt must be a non-empty string.") + word_list = textarena_word_list(template) rng = random.Random(config.seed) for _ in range(first_seed_offset): rng.choice(word_list) return [ - { - "prompt": [vf.UserMessage(content=initial_prompt)], - "answer": rng.choice(word_list), - "textarena": { - "game": config.game, - "answer_state_key": config.answer_state_key, - }, - } + TextArenaTask( + prompt=[vf.UserMessage(content=initial_prompt)], + answer=rng.choice(word_list), + textarena=TextArenaSpec( + game=config.game, + answer_state_key=config.answer_state_key, + ), + ) for _ in range(num_examples) ] @@ -137,46 +123,96 @@ def reset(self, game: str) -> TextArenaRuntimeEnv: return self.env +def textarena_word_list(env: object) -> list[str]: + raw_words = getattr(env, "word_list", None) + if isinstance(raw_words, dict): + raw_words = [ + word + for values in raw_words.values() + for word in (values if isinstance(values, list | tuple) else [values]) + ] + if not isinstance(raw_words, Sequence) or isinstance(raw_words, str | bytes): + raise ValueError("TextArena environment must expose a word_list sequence.") + words = [str(word) for word in raw_words] + if not words: + raise ValueError("TextArena word_list must not be empty.") + return words + + +def content_text(content: object) -> str: + if isinstance(content, str): + return content + if isinstance(content, Sequence) and not isinstance( + content, str | bytes | bytearray + ): + chunks: list[str] = [] + for part in content: + if isinstance(part, dict): + text = json_data(part).get("text") + if isinstance(text, str): + chunks.append(text) + return "\n".join(chunks) + return "" + + class TextArenaUser(vf.User[TextArenaUserConfig]): - async def get_response( - self, - task: vf.Task, - state: vf.State, - messages: list[vf.Message], - ) -> list[vf.UserMessage]: - session = await self.get_object("session", task, state) - assert isinstance(session, TextArenaSession) - textarena_config = task["textarena"] - assert isinstance(textarena_config, dict) - game = textarena_config["game"] - assert isinstance(game, str) - answer_state_key = textarena_config["answer_state_key"] - assert isinstance(answer_state_key, str) - ta_env = session.env or session.reset(game) - answer = task["answer"] - assert isinstance(answer, str) - assert answer - ta_env.state.game_state[answer_state_key] = answer - - assistant_messages = vf.get_messages(messages, role="assistant") - last_text = ( - _content_text(assistant_messages[-1].content) if assistant_messages else "" - ) - matches = re.findall(r"(.*?)", last_text, re.DOTALL) - guess = matches[-1].strip() if matches else "" - await asyncio.to_thread(ta_env.step, guess) - if ta_env.state.done: - reason = str(ta_env.state.game_info[0]["reason"]) - state["final_env_response"] = reason - state.stop("textarena_done") - return [vf.UserMessage(content=reason)] - - _, observation = await asyncio.to_thread(ta_env.get_observation) - assert isinstance(observation, str) - return [vf.UserMessage(content=observation)] - - -def load_taskset( - config: TextArenaTasksetConfig, -) -> TextArenaTaskset: + session: TextArenaSession + + def start(self) -> None: + self.session = TextArenaSession() + + @vf.user( + args={ + "textarena": "task.textarena", + "answer": "task.answer", + "completion": "state.completion", + }, + sets={ + "final_env_response": "state.extras.final_env_response", + "stop_condition": "state.stop_condition", + }, + ) + def respond(self, textarena: dict, answer: str, completion: list[dict]) -> dict: + return textarena_respond(self.session, textarena, answer, completion) + + +def textarena_respond( + session: TextArenaSession, textarena: dict, answer: str, completion: list[dict] +) -> dict: + game = textarena.get("game") + answer_state_key = textarena.get("answer_state_key") + if not isinstance(game, str) or not isinstance(answer_state_key, str): + raise TypeError("TextArena task config must contain string fields.") + if not isinstance(answer, str) or not answer: + raise TypeError("TextArena task requires a non-empty answer.") + env = session.env or session.reset(game) + env.state.game_state[answer_state_key] = answer + + assistant_messages = [ + message + for message in completion + if isinstance(message, dict) and message.get("role") == "assistant" + ] + last_text = ( + content_text(assistant_messages[-1].get("content")) + if assistant_messages + else "" + ) + matches = re.findall(r"(.*?)", last_text, re.DOTALL) + guess = matches[-1].strip() if matches else "" + env.step(guess) + if env.state.done: + reason = str(env.state.game_info[0]["reason"]) + return { + "messages": [vf.UserMessage(content=reason).model_dump(mode="json")], + "final_env_response": reason, + "stop_condition": "textarena_done", + } + _, observation = env.get_observation() + if not isinstance(observation, str): + raise TypeError("TextArena observation must be a string.") + return {"messages": [vf.UserMessage(content=observation).model_dump(mode="json")]} + + +def load_taskset(config: TextArenaTasksetConfig) -> TextArenaTaskset: return TextArenaTaskset(config=config) diff --git a/packages/tasksets/tasksets/utils/harbor_utils.py b/packages/tasksets/tasksets/utils/harbor_utils.py index 6525dbae25..bbfcc52765 100644 --- a/packages/tasksets/tasksets/utils/harbor_utils.py +++ b/packages/tasksets/tasksets/utils/harbor_utils.py @@ -4,16 +4,11 @@ import shutil import subprocess import sys -import tarfile -import tempfile from collections.abc import Iterable from importlib.resources import files from pathlib import Path from typing import cast -from verifiers.v1.sandbox import SandboxConfig -from verifiers.v1.utils.sandbox_utils import SandboxClient - TASKS_SUBDIR = "tasks" @@ -29,14 +24,8 @@ def bundle_tasks_root(module_name: str) -> Path: return Path(module_file).resolve().parent / TASKS_SUBDIR -def harbor_sandbox(default: SandboxConfig, configured: SandboxConfig) -> SandboxConfig: - return SandboxConfig.model_validate( - {**default.data(fill_defaults=False), **configured.data(fill_defaults=False)} - ) - - -def harbor_task_dirs(root: Path, task_names: Iterable[str] | None = None) -> list[Path]: - selected = set(task_names or []) +def harbor_task_dirs(root: Path, tasks: Iterable[str] | None = None) -> list[Path]: + selected = set(tasks or []) if not root.exists(): raise FileNotFoundError(f"Harbor tasks path not found: {root}") tasks: list[Path] = [] @@ -95,7 +84,7 @@ def download_harbor_dataset( if harbor_bin is None and uvx_bin is None: raise FileNotFoundError( f"Harbor dataset {dataset_id!r} requires the Harbor CLI or uvx. " - "Install Harbor or uvx before using Harbor Hub datasets." + "Install Harbor or uvx before using Harbor datasets." ) root = cache_dir or Path.home() / ".cache" / "verifiers" / "harbor" dataset_dir = root / ( @@ -131,33 +120,6 @@ def download_harbor_dataset( return task_root -async def upload_harbor_tests( - client: SandboxClient, sandbox_id: str, task_dir: Path -) -> None: - with tempfile.NamedTemporaryFile(suffix=".tar.gz", delete=False) as tmp_file: - tar_path = Path(tmp_file.name) - try: - with tarfile.open(tar_path, "w:gz") as tar: - for dirname, arc_root in (("solution", "oracle"), ("tests", "tests")): - root = task_dir / dirname - if not root.exists(): - continue - for item in root.iterdir(): - tar.add(item, arcname=f"{arc_root}/{item.name}") - remote_tar = "/tmp/harbor_tests.tar.gz" - await client.upload_file(sandbox_id, remote_tar, str(tar_path)) - await client.execute_command( - sandbox_id=sandbox_id, - command=( - f"mkdir -p /oracle /tests /logs/verifier && " - f"tar -xzf {remote_tar} -C / && rm {remote_tar}" - ), - timeout=900, - ) - finally: - tar_path.unlink(missing_ok=True) - - def parse_reward_text(reward_text: str) -> float: if not reward_text: return 0.0 diff --git a/packages/tasksets/tasksets/utils/openenv_utils.py b/packages/tasksets/tasksets/utils/openenv_utils.py index 2ac576d6ae..0740a34568 100644 --- a/packages/tasksets/tasksets/utils/openenv_utils.py +++ b/packages/tasksets/tasksets/utils/openenv_utils.py @@ -68,7 +68,8 @@ def start_container( sandbox_id, max_attempts=self.spec.wait_for_creation_max_attempts, ) - exposure = client.expose( + exposure = self._retry( + client.expose, sandbox_id, port=container_port, name="openenv-env", @@ -161,7 +162,7 @@ def request_schema() -> JsonData: return self._retry(request_schema) - def _retry(self, fn: Callable[..., T], *args: object) -> T: + def _retry(self, fn: Callable[..., T], *args: object, **kwargs: object) -> T: retrying = tc.Retrying( stop=tc.stop_after_attempt(self.spec.max_retries), wait=tc.wait_exponential_jitter( @@ -172,7 +173,7 @@ def _retry(self, fn: Callable[..., T], *args: object) -> T: ), reraise=True, ) - return retrying(fn, *args) + return retrying(fn, *args, **kwargs) def _failure_details( self, client: "PrimeSandboxClient", sandbox_id: str, port: int | None diff --git a/packages/tasksets/tests/test_textarena.py b/packages/tasksets/tests/test_textarena.py index 13fecf0ff3..821d208353 100644 --- a/packages/tasksets/tests/test_textarena.py +++ b/packages/tasksets/tests/test_textarena.py @@ -1,7 +1,7 @@ import sys import pytest -import verifiers as vf +import verifiers.v1 as vf from tasksets import textarena @@ -62,8 +62,7 @@ def fake_textarena(monkeypatch): return fake_ta -@pytest.mark.asyncio -async def test_textarena_user_steps_empty_guess_when_guess_tag_missing(fake_textarena): +def test_textarena_user_steps_empty_guess_when_guess_tag_missing(fake_textarena): taskset = textarena.TextArenaTaskset( config=textarena.TextArenaTasksetConfig( game="FakeWordle-v0", @@ -73,31 +72,35 @@ async def test_textarena_user_steps_empty_guess_when_guess_tag_missing(fake_text ) ) task = taskset.to_task( - vf.Task( - { - "example_id": 0, - "prompt": [], - "answer": "apple", - "textarena": { - "game": "FakeWordle-v0", - "answer_state_key": "secret_word", - }, - } + textarena.TextArenaTask( + row_id=0, + prompt=[], + answer="apple", + textarena={ + "game": "FakeWordle-v0", + "answer_state_key": "secret_word", + }, ) ) - state = vf.State.for_task(task) - state["completion"] = [vf.AssistantMessage(content=None, reasoning_content="think")] - - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - state = await env.harness.setup_state(task, state) - messages = await env.harness.runtime.user_messages(task, state) + session = textarena.TextArenaSession() + completion = [ + vf.AssistantMessage(content=None, reasoning_content="think").model_dump( + mode="json", exclude_none=True + ) + ] + payload = textarena.textarena_respond( + session, + task.textarena.model_dump(mode="json"), + task.answer, + completion, + ) ta_env = fake_textarena.envs[-1] assert ta_env.guesses == [""] - assert messages == [ + assert payload["messages"] == [ { "role": "user", "content": "Board [GAME] Feedback:\nmiss\nY----\ntry again", } ] - assert state.get("done") is None + assert "stop_condition" not in payload diff --git a/pyproject.toml b/pyproject.toml index d7140c136c..24f565b386 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -61,7 +61,7 @@ dependencies = [ dev = [ "ruff", "pre-commit", - "ty>=0.0.1a29,<0.0.22", + "ty>=0.0.40,<0.0.41", "pytest>=7.0.0", "pytest-asyncio>=0.21.0", "pytest-cov>=4.0.0", diff --git a/skills/browse-environments/SKILL.md b/skills/browse-environments/SKILL.md index dab90452d7..67eedd0f20 100644 --- a/skills/browse-environments/SKILL.md +++ b/skills/browse-environments/SKILL.md @@ -23,7 +23,7 @@ prime env list --starred - Prefer environments published by `primeintellect` first. - Keep only candidates with passing latest action/CI status from `--show-actions` or `prime env status`. - Prefer candidates updated in roughly the last 2 months. - - Prefer candidates on version `v0.1.8` or newer. + - Prefer candidates whose latest published version matches current Verifiers docs and package metadata. 4. Inspect details for shortlisted candidates: ```bash prime env info owner/name @@ -61,7 +61,7 @@ prime eval run name -m openai/gpt-4.1-mini -n 5 ```bash prime env install reverse-text --from-repo ``` -4. For v1 Taskset + Harness examples, inspect the environment package for `Taskset` / optional `Harness` classes plus `load_taskset(config: MyTasksetConfig)`, optional `load_harness(config: MyHarnessConfig)`, and the canonical `load_environment(config: vf.EnvConfig) -> vf.Env` shim delegating through `vf.load_taskset(config=config.taskset)` and `vf.load_harness(config=config.harness)`. +4. For v1 Taskset + Harness examples, inspect the environment package for `Taskset` / optional `Harness` classes plus `taskset.py` `load_taskset(config: MyTasksetConfig)` and optional `harness.py` `load_harness(config: MyHarnessConfig)`. The package loader assembles `vf.Env`. ## Anti-Patterns 1. Do not recommend building from scratch if a strong ecosystem option exists. diff --git a/skills/create-environments/SKILL.md b/skills/create-environments/SKILL.md index 6476854214..93624a06bd 100644 --- a/skills/create-environments/SKILL.md +++ b/skills/create-environments/SKILL.md @@ -1,6 +1,6 @@ --- name: create-environments -description: Create or migrate verifiers environments for the Prime Lab ecosystem. Use when asked to build a new environment from scratch, port an eval or benchmark from papers or other libraries, start from an environment on the Hub, or convert existing tasks into a package that exposes load_environment and installs cleanly with prime env install. +description: Create or migrate verifiers environments for the Prime Lab ecosystem. Use when asked to build a new environment from scratch, port an eval or benchmark from papers or other libraries, start from an environment on the Hub, or convert existing tasks into an installable environment package. --- # Create Environments @@ -44,24 +44,16 @@ prime env install math-python --from-repo - `ToolEnv` or `MCPEnv` for stateless tools. - `StatefulToolEnv` for per-rollout resources. - `CliAgentEnv` for running agent binaries in sandboxes with API interception. Override `get_sandbox_resources(state)` for per-instance resources, `build_env_vars(state)` for custom env vars. -- V1 `vf.Env` with explicit `vf.Taskset`/`vf.Harness` objects for the current taskset/harness environment pattern that separates the task collection from the rollout runner. Use this for new taskset/harness work that needs config-driven metrics, rewards, toolsets, user functions, endpoint interception, or sandboxed Python/command programs. Framework programs should build clients from `state.get_endpoint_config(api="chat")`. -3. For v1, start from the generated template. Edit `TasksetConfig` for task settings, `Taskset.load_tasks()` for task records, `Taskset.load_toolsets()` for task-owned tools, `User` subclasses for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Add a harness class only for reusable execution behavior. -4. Keep `load_environment(config: vf.EnvConfig)` as the canonical Taskset/Harness shim: -```python -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` +- V1 `vf.Env` with explicit `vf.Taskset`/`vf.Harness` objects for the current taskset/harness environment pattern that separates the task collection from the rollout runner. Use this for new taskset/harness work that needs config-driven metrics, rewards, toolsets, user servers, endpoint interception, or sandboxed Python/command programs. +3. For v1, start from the generated template. Edit `TasksetConfig` for task settings and `ToolsetConfig` / `UserConfig` entries, `Taskset.load_tasks()` for task records, `servers/toolset.py` for task-owned tools, `servers/user.py` for user behavior, and `@vf.*` methods for lifecycle, metrics, rewards, and advantages. Add a harness class only for reusable execution behavior. +4. Keep `taskset.py` with `load_taskset(config: MyTasksetConfig)` and optional `harness.py` with `load_harness(config: MyHarnessConfig)` as the v1 component entrypoints. The package loader assembles `vf.Env`. 5. For v0 environments, keep the existing `vf.Environment` patterns and preserve v0 compatibility. 6. Add `pyproject.toml` defaults in `[tool.verifiers.eval]` only when stable. ### V1 Authoring Rules -1. Keep v1 environment entrypoints tiny: `import verifiers as vf`, define `TasksetConfig` / optional `HarnessConfig` subclasses for user-facing knobs, define `Taskset` / optional `Harness` classes, then expose typed child loaders and the canonical `load_environment(config: vf.EnvConfig)` shim that delegates through `vf.load_taskset` and `vf.load_harness`. -2. Keep shared dependencies behind the taskset or harness that owns them. Use bindings as the canonical injection path; prefer serializable loader paths for bound objects in config, and use no-arg loader callables only for Python-only construction. Do not pass already-instantiated resource objects through environment loaders. Do not introduce v1 Parser/Rubric wrappers; parsing is ordinary Python. -3. Use `vf.get_messages(state.get("completion") or [], role="assistant")` when reading state completions. The helper returns typed message objects and should not receive `None`. +1. Keep v1 environment entrypoints tiny: `import verifiers.v1 as vf`, define `TasksetConfig` / optional `HarnessConfig` subclasses for user-facing knobs, define `Taskset` / optional `Harness` classes, then expose typed child loaders in `taskset.py` and optional `harness.py`. Do not define package-level `load_environment` for v1 packages. +2. Keep shared dependencies behind the taskset, harness, toolset, user, or runtime that owns them. Configure tools and users with serializable `ToolsetConfig` / `UserConfig` values; implement tool and user behavior on `Toolset` / `User` classes under `servers/`. Do not pass already-instantiated resource objects through environment loaders. Do not introduce v1 Parser/Rubric wrappers; parsing is ordinary Python. +3. Iterate over `state.completion` directly when reading rollout messages. Filter by `message.role` with ordinary Python. 4. Use `program.channels` for v1 program protocol/channel selection. Do not use stale `program.tools` terminology. 5. Use generated child loaders as typed component entrypoints. Add implementation behavior to the taskset or harness class through config fields, `load_*` methods, `User` subclasses, `Toolset`, and `@vf.*` lifecycle methods. 6. Put settings as leaf fields on the taskset or harness config that owns them. @@ -97,14 +89,20 @@ max_turns = 8 ``` 6. In code, use the current class-based config shape: ```python -import verifiers as vf +import verifiers.v1 as vf class MyTasksetConfig(vf.TasksetConfig): system_prompt: vf.SystemPrompt = "Answer exactly." +class MyTask(vf.Task): + answer: str + + class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: """Return serializable task records as a list, generator, or Dataset.""" if split == "eval": @@ -118,24 +116,18 @@ class MyTaskset(vf.Taskset[MyTasksetConfig]): ] @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") + async def correct_answer(self, task: MyTask, state: vf.State) -> float: + messages = [ + message for message in state.completion if message.role == "assistant" + ] if not messages: return 0.0 response = str(messages[-1].content or "").strip() - return float(response == task["answer"]) + return float(response == task.answer) def load_taskset(config: MyTasksetConfig) -> MyTaskset: return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) ``` 7. Use `prime env init my-env --v1` as the reference shape when an implementation starts to drift. @@ -162,9 +154,9 @@ prime env pull owner/name -t ./tmp-env ## Non-Negotiable Quality Rules 1. Use deterministic, well-defined reward checks or LLM judges. 2. Avoid best-effort deterministic heuristics such as keyword style checks except as an explicit last resort with user sign-off. -3. Make environments self-contained after install. Do not require users to run background servers before `load_environment()`. +3. Make environments self-contained after install. Do not require users to run background servers before environment loading. 4. Manage external resources inside the environment lifecycle. -5. Validate required secrets in `load_environment()` via `vf.ensure_keys(...)`. +5. Validate required secrets at the owning loader or config boundary so failures are explicit and early. 6. Surface feature limits directly. Do not ship hacky workarounds without explicit user approval. ## Verification Gate diff --git a/skills/review-environments/SKILL.md b/skills/review-environments/SKILL.md index 5ce7fe2789..e31bcf2d92 100644 --- a/skills/review-environments/SKILL.md +++ b/skills/review-environments/SKILL.md @@ -15,8 +15,8 @@ Find correctness risks and regressions first, then assess maintainability and ec ## Review Workflow 1. Identify environment contract: -- `load_environment(...)` -- base class and rollout behavior (`SingleTurnEnv`, `MultiTurnEnv`, `ToolEnv`/`MCPEnv`/`StatefulToolEnv`, `SandboxEnv`/`PythonEnv`, V1 `vf.Env` with explicit `vf.Taskset`/`vf.Harness` objects for framework programs, `CliAgentEnv` for sandboxed agents) +- v0 `load_environment(...)` or v1 discovered `taskset.py` / optional `harness.py` components +- base class and rollout behavior (`SingleTurnEnv`, `MultiTurnEnv`, `ToolEnv`/`MCPEnv`/`StatefulToolEnv`, `SandboxEnv`/`PythonEnv`, v1 `vf.Env` with explicit `vf.Taskset`/`vf.Harness` objects for framework programs, `CliAgentEnv` for sandboxed agents) - rubric and metrics 2. Verify installability and runtime entrypoint with the canonical eval path. Do not add `--skip-upload` unless the user explicitly requests that deviation; standard runs save automatically for the private Evaluations tab and `prime eval view`: ```bash @@ -39,21 +39,21 @@ prime eval run -m openai/gpt-4.1-mini -n 5 - Flag best-effort keyword or style heuristics unless explicitly approved. - Verify the scoring semantics from code before treating a low reward as an implementation failure. Some environments intentionally complete with `0.0` reward when the model fails the task. 2. Environment self-containment: -- Flag any requirement for user-managed background services before `load_environment()`. +- Flag any requirement for user-managed background services before environment loading. - Require environment-managed lifecycle for sandboxes/sessions. 3. v1 taskset/harness contracts: -- Expect new taskset/harness environments to use the v1 `vf.Env` / `vf.Taskset` / `vf.Harness` boundary, with `load_taskset(config: MyTasksetConfig)` and optional `load_harness(config: MyHarnessConfig)` defining child config types, plus the canonical `load_environment(config: vf.EnvConfig)` shim delegating through `vf.load_taskset(config=config.taskset)` and `vf.load_harness(config=config.harness)`. +- Expect new taskset/harness environments to use the v1 `vf.Env` / `vf.Taskset` / `vf.Harness` boundary, with `taskset.py` exposing `load_taskset(config: MyTasksetConfig)` and optional `harness.py` exposing `load_harness(config: MyHarnessConfig)` to define child config types. The package loader assembles `vf.Env`. - Expect tasksets to own task data, task-owned tools, user behavior, metrics, rewards, and task-specific config. Flag one-off harness classes that only wrap task behavior. -- Review v1 implementations against the generated `prime env init my-env --v1` shape: task settings in `TasksetConfig`, tasks in `load_tasks`, task-owned tools in `load_toolsets`, user behavior in `User` subclasses, lifecycle/metrics/rewards as `@vf.*` methods, and typed component entrypoints through `load_taskset`, optional `load_harness`, and `load_environment`. +- Review v1 implementations against the generated `prime env init my-env --v1` shape: task settings and `ToolsetConfig` / `UserConfig` entries in `TasksetConfig`, tasks in `load_tasks`, task-owned tools in `servers/toolset.py`, user behavior in `servers/user.py`, lifecycle/metrics/rewards/advantages as `@vf.*` methods, and typed component entrypoints through `load_taskset` and optional `load_harness`. - Require environment packages and READMEs to preserve the generated `prime env init` structure. Flag hand-scaffolded environments and freeform environment READMEs; authors should fill in the CLI-generated template sections instead of inventing a new shape. -- Expect shared dependencies to use bindings owned by the taskset, toolset, user, program, or harness that needs them. Flag pre-initialized resource objects passed through environment loaders; object entries should be serializable loader paths or no-arg loader callables. +- Expect shared dependencies to live behind the taskset, harness, toolset, user, or runtime that owns them. Flag pre-initialized resource objects passed through environment loaders; config entries should be serializable. - Verify `Task` data is serializable, `state` remains serializable at rollout boundaries, and model/client controls flow through runtime state rather than top-level dataset columns. -- For V1 harness programs, verify framework clients consume `state.get_endpoint_config(api="chat")` rather than hardcoding an upstream LLM endpoint. For `CliAgentEnv` agents, verify sandboxed agent code consumes the injected interception endpoint; the proxy is what makes rollouts visible to the rubric. +- For v1 harness programs, verify framework clients consume the injected interception endpoint rather than hardcoding an upstream LLM endpoint. For `CliAgentEnv` agents, verify sandboxed agent code consumes the injected interception endpoint; the proxy is what makes rollouts visible to the rubric. 4. Migration fidelity: - For ports, verify one-to-one equivalence of prompts, tool traces, and scoring logic. - Flag any assumptions made without user decision. 5. Secrets handling: -- Ensure required keys are validated in `load_environment()` with `vf.ensure_keys(...)`. +- Ensure required keys are validated at the owning loader or config boundary so failures are explicit and early. 6. Performance and scaling: - Identify obvious bottlenecks in dataset loading, rubric calls, or tool execution. 7. Packaging and repo hygiene: diff --git a/skills/train-with-environments/SKILL.md b/skills/train-with-environments/SKILL.md index 00f23f4dc3..2cab199654 100644 --- a/skills/train-with-environments/SKILL.md +++ b/skills/train-with-environments/SKILL.md @@ -36,7 +36,7 @@ prime lab setup --prime-rl prime env install my-env prime eval run my-env -m openai/gpt-4.1-mini -n 20 -r 3 -s ``` -2. For v1 Taskset + Harness environments, verify the package exposes `load_taskset(config: MyTasksetConfig)`, optional `load_harness(config: MyHarnessConfig)`, and the canonical `load_environment(config: vf.EnvConfig) -> vf.Env` shim delegating through `vf.load_taskset(config=config.taskset)` and `vf.load_harness(config=config.harness)`; trainers interact with the same environment boundary even when the implementation is BYO Harness internally. +2. For v1 Taskset + Harness environments, verify the package exposes `taskset.py` `load_taskset(config: MyTasksetConfig)` and optional `harness.py` `load_harness(config: MyHarnessConfig)`; the package loader assembles `vf.Env`, so trainers interact with the same environment boundary even when the implementation is BYO Harness internally. 3. Confirm reward diversity exists at baseline. 4. Start with conservative run length and inspect samples early. diff --git a/tests/test_environment.py b/tests/test_environment.py index cc30b19bed..a2b70901ac 100644 --- a/tests/test_environment.py +++ b/tests/test_environment.py @@ -65,11 +65,9 @@ async def rollout( state["trajectory"].append(trajectory_step) state["is_completed"] = True - from verifiers.utils.message_utils import concat_messages - last_prompt = state["trajectory"][-1]["prompt"] last_completion = state["trajectory"][-1]["completion"] - full_conversation = concat_messages([last_prompt, last_completion]) + full_conversation = [*last_prompt, *last_completion] state["completion"] = full_conversation[len(state["prompt"]) :] except vf.Error as e: state["error"] = e diff --git a/tests/test_environment_extra.py b/tests/test_environment_extra.py index ab6f5eb672..9dc5b53afc 100644 --- a/tests/test_environment_extra.py +++ b/tests/test_environment_extra.py @@ -78,11 +78,9 @@ async def rollout( state["trajectory"].append(trajectory_step) state["is_completed"] = True - from verifiers.utils.message_utils import concat_messages - last_prompt = state["trajectory"][-1]["prompt"] last_completion = state["trajectory"][-1]["completion"] - full_conversation = concat_messages([last_prompt, last_completion]) + full_conversation = [*last_prompt, *last_completion] state["completion"] = full_conversation[len(state["prompt"]) :] return state diff --git a/tests/test_envs.py b/tests/test_envs.py index 617e8f6592..ba8097f4e7 100644 --- a/tests/test_envs.py +++ b/tests/test_envs.py @@ -1,5 +1,7 @@ import os import importlib.util +import re +import shlex import subprocess import sys from pathlib import Path @@ -29,14 +31,15 @@ "browser_cua_example", # Uses prime-tunnel which is still experimental and has low usage limits "terminus_harbor", - "opencode_harbor", + "harbor_v1", + "opencode_harbor_v1", ] SKIPPED_ENV_LOADING_ENVS = [ # OpenEnv datasets are built by resetting seeds in sandbox-backed env servers. # Skip generic load checks here and cover via dedicated OpenEnv tests. - "openenv_echo", - "openenv_textarena", + "openenv_echo_v1", + "openenv_textarena_v1", # R2E-Gym pulls a full image-backed SWE taskset; cover it with dedicated v1 tests. "rlm_swe_v1", ] @@ -92,40 +95,70 @@ def test_readme_exists(env_dir: Path): assert (env_dir / "README.md").exists(), "README.md does not exist" +@pytest.mark.parametrize( + "env_dir", + sorted(Path("environments").glob("*_v1")), + ids=lambda x: x.name, +) +def test_v1_readme_uses_project_name(env_dir: Path): + with open(env_dir / "pyproject.toml", "rb") as f: + project_name = tomllib.load(f)["project"]["name"] + readme = (env_dir / "README.md").read_text() + tokens: list[str] = [] + tokens.extend(re.findall(r"^#\s+(.+)$", readme, flags=re.MULTILINE)) + tokens.extend(re.findall(r"\*\*Environment ID\*\*: `([^`]+)`", readme)) + tokens.extend(re.findall(r"\bprime eval run ([^\s\\]+)", readme)) + mismatches = [token for token in tokens if token != project_name] + assert not mismatches, ( + f"{env_dir.name} README should use project name {project_name!r}; " + f"found {mismatches!r}." + ) + + def test_alphabet_sort_v1_validates_parameters(): - module_path = Path("environments/alphabet_sort/alphabet_sort_v1.py").resolve() - spec = importlib.util.spec_from_file_location("alphabet_sort_v1_test", module_path) - assert spec is not None and spec.loader is not None - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - spec.loader.exec_module(module) + env_dir = Path("environments/alphabet_sort_v1").resolve() + sys.path.insert(0, str(env_dir)) + try: + module = importlib.import_module("alphabet_sort_v1.taskset") + finally: + sys.path.remove(str(env_dir)) with pytest.raises(ValueError, match="min_turns must be at least 1"): - module.AlphabetSortTaskset(config=module.AlphabetSortTasksetConfig(min_turns=0)) + list( + module.AlphabetSortTaskset( + config=module.AlphabetSortTasksetConfig(min_turns=0) + ).load_tasks() + ) with pytest.raises( ValueError, match="min_turns must be less than or equal to max_turns" ): - module.AlphabetSortTaskset( - config=module.AlphabetSortTasksetConfig(min_turns=3, max_turns=2) + list( + module.AlphabetSortTaskset( + config=module.AlphabetSortTasksetConfig(min_turns=3, max_turns=2) + ).load_tasks() ) with pytest.raises(ValueError, match="min_names_per_turn must be at least 1"): - module.AlphabetSortTaskset( - config=module.AlphabetSortTasksetConfig(min_names_per_turn=0) + list( + module.AlphabetSortTaskset( + config=module.AlphabetSortTasksetConfig(min_names_per_turn=0) + ).load_tasks() ) with pytest.raises( ValueError, match="min_names_per_turn must be less than or equal to max_names_per_turn", ): - module.AlphabetSortTaskset( - config=module.AlphabetSortTasksetConfig( - min_names_per_turn=3, - max_names_per_turn=2, - ) + list( + module.AlphabetSortTaskset( + config=module.AlphabetSortTasksetConfig( + min_names_per_turn=3, + max_names_per_turn=2, + ) + ).load_tasks() ) @pytest.mark.parametrize("env_name", ["alphabet_sort", "math_python"]) -def test_v1_wrapper_rejects_unknown_kwargs(env_name: str): +def test_v0_wrappers_reject_v1_kwargs(env_name: str): module_path = Path("environments") / env_name / f"{env_name}.py" spec = importlib.util.spec_from_file_location( f"{env_name}_wrapper_test", module_path @@ -135,10 +168,8 @@ def test_v1_wrapper_rejects_unknown_kwargs(env_name: str): sys.modules[spec.name] = module spec.loader.exec_module(module) - with pytest.raises( - TypeError, match="Unsupported v1 load_environment kwargs: extra" - ): - module.load_environment(v1=True, extra=True) + with pytest.raises(TypeError): + module.load_environment(v1=True) @pytest.mark.slow @@ -149,10 +180,18 @@ def test_env(env_dir: Path, tmp_path_factory: pytest.TempPathFactory): pytest.skip(f"Skipping {env_dir.name}") if env_dir.name in SKIPPED_ENV_LOADING_ENVS: pytest.skip(f"Skipping dedicated-runtime smoke test for {env_dir.name}") + if env_dir.name in {"toxicity_explanation", "wiki_search"} and not os.getenv( + "OPENAI_API_KEY" + ): + pytest.skip(f"Skipping {env_dir.name} load test without OPENAI_API_KEY") + if env_dir.name == "nemo_gym_env_v1" and sys.version_info < (3, 12): + pytest.skip("Skipping nemo_gym_env_v1 install test on Python < 3.12") tmp_venv_dir = tmp_path_factory.mktemp(f"venv_{env_dir.name}") repo_root = Path(__file__).parent.parent + python = shlex.quote(sys.executable) cmd = ( - f"cd {tmp_venv_dir} && uv venv --clear && source .venv/bin/activate && " + f"cd {tmp_venv_dir} && uv venv --clear --python {python} && " + "source .venv/bin/activate && " "uv pip install " "--exclude-newer-package prime-pydantic-config=2026-05-20T00:00:00Z " f"{repo_root.as_posix()} && " diff --git a/tests/test_eval_cli.py b/tests/test_eval_cli.py index 1b76cdc9b9..f57d66d0b9 100644 --- a/tests/test_eval_cli.py +++ b/tests/test_eval_cli.py @@ -199,7 +199,7 @@ def test_cli_v1_env_config_overrides_preserve_env_args_config( module_name = f"cli_override_env_{time.time_ns()}" (tmp_path / f"{module_name}.py").write_text( """ -import verifiers as vf +import verifiers.v1 as vf class DemoTasksetConfig(vf.TasksetConfig): @@ -217,10 +217,6 @@ def load_taskset(config: DemoTasksetConfig): def load_harness(config: DemoHarnessConfig): raise RuntimeError("not used") - - -def load_environment(config: vf.EnvConfig): - raise RuntimeError("not used") """, encoding="utf-8", ) @@ -247,11 +243,11 @@ def load_environment(config: vf.EnvConfig): assert captured["configs"][0].env_args == { "config": { "taskset": { - "taskset_id": "override-id", + "id": "override-id", "count": 2, "enabled": False, }, - "harness": {"harness_id": "demo-harness", "max_turns": 4}, + "harness": {"id": "demo-harness", "max_turns": 4}, } } @@ -1010,7 +1006,7 @@ def test_load_toml_config_with_args_taskset_harness(): assert "harness" not in result[0] -def test_load_toml_config_allows_taskset_id_without_env_id(): +def test_load_toml_config_allows_taskset_without_env_id(): with tempfile.NamedTemporaryFile(suffix=".toml", delete=False, mode="w") as f: f.write( "[[eval]]\n" diff --git a/tests/test_harbor_v1.py b/tests/test_harbor_v1.py new file mode 100644 index 0000000000..b65323a693 --- /dev/null +++ b/tests/test_harbor_v1.py @@ -0,0 +1,57 @@ +import importlib +import sys +from pathlib import Path +from typing import Any + +import verifiers.v1 as vf +from harnesses import OpenCode +from tasksets import HarborTaskset +from verifiers.v1.loaders import load_environment_from_components + + +def _load_harbor_modules(monkeypatch: Any) -> tuple[Any, Any]: + env_dir = Path(__file__).resolve().parent.parent / "environments" / "harbor_v1" + monkeypatch.syspath_prepend(str(env_dir)) + for name in ("harbor_v1", "harbor_v1.taskset", "harbor_v1.harness"): + sys.modules.pop(name, None) + return ( + importlib.import_module("harbor_v1"), + importlib.import_module("harbor_v1.taskset"), + ) + + +def test_harbor_v1_loads_thin_taskset_harness_package(monkeypatch: Any) -> None: + package, module = _load_harbor_modules(monkeypatch) + + env = load_environment_from_components(package, {"config": {}}) + + assert isinstance(env, vf.Env) + assert isinstance(env.taskset, HarborTaskset) + assert isinstance(env.harness, OpenCode) + assert env.taskset.id == "harbor" + assert env.taskset.config.source == "package" + assert env.taskset.config.dataset == "harbor_v1" + task = next(iter(env.taskset)) + assert task.task_name == "hello-world" + assert task.image == "python:3.11-slim" + assert Path(task.task_dir).parent == Path(module.__file__).parent / "tasks" + + +def test_harbor_v1_allows_package_dataset_override(monkeypatch: Any) -> None: + package, _ = _load_harbor_modules(monkeypatch) + + env = load_environment_from_components( + package, + { + "config": { + "taskset": { + "source": "package", + "dataset": "harbor_v1", + } + } + }, + ) + + task = next(iter(env.taskset)) + assert task.task_name == "hello-world" + assert Path(task.task_dir).name == "hello-world" diff --git a/tests/test_init_script.py b/tests/test_init_script.py index 26e11c751f..747b16a613 100644 --- a/tests/test_init_script.py +++ b/tests/test_init_script.py @@ -1,13 +1,26 @@ from pathlib import Path import verifiers as vf +import verifiers.v1 as vf1 +from verifiers.scripts.build import _resolve_project_dir from verifiers.scripts.init import init_environment def read_env_file(root: Path, env_id: str) -> str: module_name = env_id.replace("-", "_") + taskset_file = root / module_name / module_name / "taskset.py" + if taskset_file.exists(): + return taskset_file.read_text() + package_file = root / module_name / module_name / f"{module_name}.py" + if package_file.exists(): + return package_file.read_text() return (root / module_name / f"{module_name}.py").read_text() +def package_dir(root: Path, env_id: str) -> Path: + module_name = env_id.replace("-", "_") + return root / module_name / module_name + + def test_init_default_writes_v0_stub(tmp_path: Path) -> None: root = init_environment("foo", path=str(tmp_path)) content = read_env_file(tmp_path, "foo") @@ -22,9 +35,19 @@ def test_init_default_writes_v0_stub(tmp_path: Path) -> None: def test_init_v1_writes_taskset_template(tmp_path: Path) -> None: init_environment("bar", path=str(tmp_path), v1=True) content = read_env_file(tmp_path, "bar") + package = package_dir(tmp_path, "bar") + init_content = (package / "__init__.py").read_text() + pyproject = (tmp_path / "bar" / "pyproject.toml").read_text() + user_config = (package / "servers" / "user" / "config.py").read_text() + user_server = (package / "servers" / "user" / "user.py").read_text() + toolset_config = (package / "servers" / "example" / "config.py").read_text() + toolset_server = (package / "servers" / "example" / "toolset.py").read_text() assert "class BarTasksetConfig(vf.TasksetConfig):" in content + assert "class BarTask(vf.Task):" in content + assert "answer: str" in content assert "class BarTaskset(vf.Taskset[BarTasksetConfig]):" in content + assert "task_type = BarTask" in content assert 'system_prompt: vf.SystemPrompt = "Answer exactly."' in content assert '"""Taskset implementation for bar.' in content assert 'def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks:' in content @@ -33,13 +56,36 @@ def test_init_v1_writes_taskset_template(tmp_path: Path) -> None: in content ) assert "def load_system_prompt" not in content - assert "async def correct_answer(self, task: vf.Task, state: vf.State)" in content + assert "async def correct_answer(self, task: BarTask, state: vf.State)" in content + assert "task.answer" in content + assert "task.data" not in content assert "def load_taskset(config: BarTasksetConfig) -> BarTaskset:" in content assert '"""Typed taskset loader used by vf.load_taskset."""' in content assert "return BarTaskset(config=config)" in content - assert "taskset=vf.load_taskset(config=config.taskset)" in content - assert '"""Loader pattern for all Taskset/Harness environments."""' in content - assert "harness=vf.load_harness(config=config.harness)" in content + assert "from .servers.user import UserConfig" in content + assert "from .servers.example import ExampleToolsetConfig" in content + assert "UserConfig()" in content + assert ( + 'toolsets: vf.ToolsetConfigs = {"example": ExampleToolsetConfig()}' in content + ) + assert "def load_user(" not in content + assert "def load_toolsets(" not in content + assert "def load_environment" not in content + assert '"""bar environment package."""' in init_content + assert "load_environment" not in init_content + assert 'include = ["bar/**/*", "pyproject.toml", "README.md"]' in pyproject + assert "class UserConfig(vf.UserConfig):" in user_config + assert "class User(vf.User[UserConfig]):" in user_server + assert "@vf.user(" in user_server + assert ( + "def respond(self, task: dict, state: dict, transcript: list[dict]) -> dict:" + in (user_server) + ) + assert "class ExampleToolsetConfig(vf.ToolsetConfig):" in toolset_config + assert "name:" not in toolset_config + assert "class ExampleToolset(vf.Toolset[ExampleToolsetConfig]):" in toolset_server + assert "@vf.tool" in toolset_server + assert "def reverse_text(self, text: str) -> str:" in toolset_server assert "class EnvTaskset(" not in content assert "_default_" not in content assert 'tasks: str = "load_tasks"' not in content @@ -53,23 +99,28 @@ def test_init_v1_template_loads_with_vf_load_environment( monkeypatch.syspath_prepend(str(tmp_path / "loadable_v1")) env = vf.load_environment("loadable-v1") + taskset = vf1.load_taskset("loadable-v1") dataset = env.get_dataset() assert len(dataset) == 1 assert dataset[0]["answer"] == "cba" + assert taskset.config.system_prompt == "Answer exactly." def test_init_v1_with_harness_writes_harness_stub(tmp_path: Path) -> None: init_environment("baz", path=str(tmp_path), v1=True, with_harness=True) - content = read_env_file(tmp_path, "baz") - - assert "class BazTaskset(vf.Taskset[BazTasksetConfig]):" in content - assert "class BazHarnessConfig(vf.HarnessConfig):" in content - assert "class BazHarness(vf.Harness[BazHarnessConfig]):" in content - assert "def load_harness(config: BazHarnessConfig) -> BazHarness:" in content - assert "taskset=vf.load_taskset(config=config.taskset)" in content - assert "harness=vf.load_harness(config=config.harness)" in content + taskset_content = read_env_file(tmp_path, "baz") + harness_content = (package_dir(tmp_path, "baz") / "harness.py").read_text() + + assert "class BazTaskset(vf.Taskset[BazTasksetConfig]):" in taskset_content + assert "class BazHarnessConfig(vf.HarnessConfig):" in harness_content + assert "class BazHarness(vf.Harness[BazHarnessConfig]):" in harness_content + assert "def load_harness(config: BazHarnessConfig) -> BazHarness:" in ( + harness_content + ) + assert "def load_environment" not in taskset_content + assert "def load_environment" not in harness_content def test_init_with_harness_without_v1_warns_and_uses_v0(tmp_path: Path, capsys) -> None: @@ -84,32 +135,40 @@ def test_init_with_harness_without_v1_warns_and_uses_v0(tmp_path: Path, capsys) def test_init_v1_multifile_exports_component_loaders(tmp_path: Path) -> None: init_environment("pkg-env", path=str(tmp_path), v1=True, multi_file=True) - package_dir = tmp_path / "pkg_env" / "pkg_env" - init_content = (package_dir / "__init__.py").read_text() - env_content = (package_dir / "pkg_env.py").read_text() + package = package_dir(tmp_path, "pkg-env") + init_content = (package / "__init__.py").read_text() + taskset_content = (package / "taskset.py").read_text() - assert "from .pkg_env import load_environment, load_taskset" in init_content - assert "__all__ = ['load_environment', 'load_taskset']" in init_content - assert "class PkgEnvTaskset(vf.Taskset[PkgEnvTasksetConfig]):" in env_content - assert "return PkgEnvTaskset(config=config)" in env_content + assert '"""pkg-env environment package."""' in init_content + assert "load_environment" not in init_content + assert "class PkgEnvTaskset(vf.Taskset[PkgEnvTasksetConfig]):" in taskset_content + assert "return PkgEnvTaskset(config=config)" in taskset_content + assert (package / "servers" / "user" / "config.py").exists() + assert (package / "servers" / "user" / "user.py").exists() + assert (package / "servers" / "example" / "config.py").exists() + assert (package / "servers" / "example" / "toolset.py").exists() def test_init_openenv_writes_v1_taskset_template(tmp_path: Path) -> None: init_environment("openenv-sample", path=str(tmp_path), openenv=True) content = read_env_file(tmp_path, "openenv-sample") + package = package_dir(tmp_path, "openenv-sample") pyproject = (tmp_path / "openenv_sample" / "pyproject.toml").read_text() assert "from tasksets import OpenEnvTaskset, OpenEnvTasksetConfig" in content assert ( "def load_taskset(config: OpenEnvTasksetConfig) -> OpenEnvTaskset:" in content ) - assert "taskset=vf.load_taskset(config=config.taskset)" in content - assert "harness=vf.load_harness(config=config.harness)" in content + assert "def load_environment" not in content assert "vf.OpenEnvEnv" not in content assert '"tasksets[openenv]>=0.1.5"' in pyproject + assert 'include = ["openenv_sample/**/*", "pyproject.toml", "README.md"]' in ( + pyproject + ) + assert (package / "proj" / "openenv.yaml").exists() -def test_init_openenv_multifile_exports_taskset_loader(tmp_path: Path) -> None: +def test_init_openenv_multifile_uses_component_package(tmp_path: Path) -> None: init_environment( "openenv-pkg", path=str(tmp_path), @@ -120,5 +179,20 @@ def test_init_openenv_multifile_exports_taskset_loader(tmp_path: Path) -> None: tmp_path / "openenv_pkg" / "openenv_pkg" / "__init__.py" ).read_text() - assert "from .openenv_pkg import load_environment, load_taskset" in init_content - assert "__all__ = ['load_environment', 'load_taskset']" in init_content + assert '"""openenv-pkg environment package."""' in init_content + assert "load_environment" not in init_content + assert (tmp_path / "openenv_pkg" / "openenv_pkg" / "taskset.py").exists() + assert (tmp_path / "openenv_pkg" / "openenv_pkg" / "proj").is_dir() + + +def test_vf_build_resolves_openenv_component_project(tmp_path: Path) -> None: + init_environment( + "openenv-pkg", + path=str(tmp_path), + openenv=True, + multi_file=True, + ) + + project_dir = _resolve_project_dir(tmp_path, "openenv_pkg") + + assert project_dir == tmp_path / "openenv_pkg" / "openenv_pkg" / "proj" diff --git a/tests/test_langchain_deep_agents_wikispeedia.py b/tests/test_langchain_deep_agents_wikispeedia.py index 4632a8423e..7fef8d1eef 100644 --- a/tests/test_langchain_deep_agents_wikispeedia.py +++ b/tests/test_langchain_deep_agents_wikispeedia.py @@ -1,37 +1,31 @@ import importlib -import inspect import sys -import types from pathlib import Path import pytest -import verifiers as vf +import verifiers.v1 as vf +from verifiers.v1.loaders import load_environment_from_components -def load_module(monkeypatch: pytest.MonkeyPatch): +def load_modules(monkeypatch: pytest.MonkeyPatch): env_dir = ( - Path(__file__).parents[1] / "environments" / "langchain_deep_agents_wikispeedia" + Path(__file__).parents[1] + / "environments" + / "langchain_deep_agents_wikispeedia_v1" ) monkeypatch.syspath_prepend(str(env_dir)) - sys.modules.pop("langchain_deep_agents_wikispeedia", None) - sys.modules.pop("wiki_graph", None) - return importlib.import_module("langchain_deep_agents_wikispeedia") - - -class FakeWiki: - articles = {"A": "Article A", "B": "Article B"} - links = {"A": ["B"], "B": []} - distances = {"A": {"B": 1}} - - def get_text(self, article: str) -> str: - return self.articles[article] - - def get_links(self, article: str) -> list[str]: - return self.links[article] - - def get_human_stats(self, source: str, target: str): - return None + for name in ( + "langchain_deep_agents_wikispeedia_v1", + "langchain_deep_agents_wikispeedia_v1.taskset", + "langchain_deep_agents_wikispeedia_v1.harness", + "langchain_deep_agents_wikispeedia_v1.wiki_graph", + ): + sys.modules.pop(name, None) + return ( + importlib.import_module("langchain_deep_agents_wikispeedia_v1"), + importlib.import_module("langchain_deep_agents_wikispeedia_v1.taskset"), + ) def make_small_wiki(module): @@ -55,64 +49,67 @@ def make_small_wiki(module): def test_wikispeedia_loads_as_v1_taskset_harness( monkeypatch: pytest.MonkeyPatch, ) -> None: - module = load_module(monkeypatch) + package, _ = load_modules(monkeypatch) - env = module.load_environment(config=module.WikispeediaEnvConfig()) + env = load_environment_from_components(package, {}) assert isinstance(env, vf.Env) assert isinstance(env.taskset, vf.Taskset) assert isinstance(env.harness, vf.Harness) - assert env.taskset.taskset_id == "langchain-deep-agents-wikispeedia" + assert env.taskset.id == "langchain-deep-agents-wikispeedia" def test_wikispeedia_env_config_reaches_taskset_and_harness( monkeypatch: pytest.MonkeyPatch, ) -> None: - module = load_module(monkeypatch) + package, module = load_modules(monkeypatch) wiki = make_small_wiki(module) monkeypatch.setattr(module, "load_wiki_graph", lambda cache_dir=None: wiki) - env = module.load_environment( - config=module.WikispeediaEnvConfig( - taskset={ - "train_size": 2, - "eval_size": 1, - "min_path_length": 1, - "max_path_length": 1, - "eval_target_fraction": 0.5, - "allow_go_back": False, - "links_only": True, - "max_turns": 7, - }, - harness={ - "max_turns": 8, - "timeout_seconds": 9.0, - }, - ) + env = load_environment_from_components( + package, + { + "config": { + "taskset": { + "train_size": 2, + "eval_size": 1, + "min_path_length": 1, + "max_path_length": 1, + "eval_target_fraction": 0.5, + "allow_go_back": False, + "links_only": True, + "max_turns": 7, + }, + "harness": {"max_turns": 8, "timeout_seconds": 9.0}, + } + }, ) train_rows = list(env.taskset) eval_rows = [ - env.taskset.to_task(vf.Task(dict(row))) - for row in env.taskset.get_eval_dataset() + env.taskset.to_task(dict(row)) for row in env.taskset.get_eval_dataset() ] assert len(train_rows) == 2 assert len(eval_rows) == 1 - assert train_rows[0]["max_turns"] == 7 + assert train_rows[0].max_turns == 7 + assert train_rows[0].links_only is True + assert train_rows[0].allow_go_back is False assert env.harness.config.max_turns == 8 assert env.harness.config.timeout_seconds == 9.0 - assert [tool.__name__ for tool in env.taskset.toolsets[0].tools] == ["click_link"] def test_wikispeedia_rows_use_v1_task_shape( monkeypatch: pytest.MonkeyPatch, ) -> None: - module = load_module(monkeypatch) + _, module = load_modules(monkeypatch) + wiki = make_small_wiki(module) dataset = module.build_dataset( - FakeWiki(), + wiki, [("A", "B", 1)], + cache_dir="/tmp/wiki", links_only=False, + allow_go_back=True, max_turns=7, ) row = dataset[0] @@ -120,276 +117,71 @@ def test_wikispeedia_rows_use_v1_task_shape( assert "task" not in row assert row["task_id"] == "A->B" assert row["max_turns"] == 7 - assert row["info"] == {"source": "A", "target": "B", "shortest_path": 1} - - -def test_wikispeedia_taskset_sources_use_disjoint_target_split( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = load_module(monkeypatch) - wiki = make_small_wiki(module) - monkeypatch.setattr(module, "load_wiki_graph", lambda cache_dir=None: wiki) - taskset = module.load_taskset( - config=module.WikispeediaTasksetConfig( - train_size=2, - eval_size=1, - min_path_length=1, - max_path_length=1, - eval_target_fraction=0.5, - ) - ) - - train_rows = list(taskset) - eval_rows = [ - taskset.to_task(vf.Task(dict(row))) for row in taskset.get_eval_dataset() - ] - - assert len(train_rows) == 2 - assert len(eval_rows) == 1 - assert {row["answer"] for row in train_rows}.isdisjoint( - {row["answer"] for row in eval_rows} - ) + assert row["cache_dir"] == "/tmp/wiki" + assert row["source"] == "A" + assert row["target"] == "B" + assert row["shortest_path"] == 1 -def test_wikispeedia_efficiency_weight_uses_fresh_reward_wrapper( +def test_wikispeedia_navigation_uses_state_extras( monkeypatch: pytest.MonkeyPatch, ) -> None: - module = load_module(monkeypatch) + _, module = load_modules(monkeypatch) wiki = make_small_wiki(module) - monkeypatch.setattr(module, "load_wiki_graph", lambda cache_dir=None: wiki) - - weighted = module.load_taskset( - config=module.WikispeediaTasksetConfig(efficiency_weight=0.5) + task = module.WikispeediaTask( + prompt=[{"role": "user", "content": "start"}], + answer="B", + source="A", + target="B", + shortest_path=1, ) - plain = module.load_taskset( - config=module.WikispeediaTasksetConfig(efficiency_weight=0.0) - ) - - assert any(fn.__name__ == "path_efficiency" for fn in weighted.rewards) - assert any(fn is module.path_efficiency for fn in plain.metrics) - assert not getattr(module.path_efficiency, "reward", False) + state = vf.State(task_id=task.task_id) + module.init_navigation_state(task, state) + result = module.click_link_result("B", wiki, state) -def test_wikispeedia_taskset_owns_navigation_tools( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = load_module(monkeypatch) - - taskset = module.load_taskset( - config=module.WikispeediaTasksetConfig(allow_go_back=True) - ) - names = [tool.__name__ for tool in taskset.toolsets[0].tools] - no_back = module.load_taskset( - config=module.WikispeediaTasksetConfig(allow_go_back=False) - ) - - assert names == ["click_link", "go_back"] - assert [tool.__name__ for tool in no_back.toolsets[0].tools] == ["click_link"] - assert module.load_harness(config=module.WikispeediaHarnessConfig()).toolsets == [] - - -def test_wikispeedia_system_prompt_matches_available_tools( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = load_module(monkeypatch) - - with_back = module.load_taskset( - config=module.WikispeediaTasksetConfig(allow_go_back=True) - ) - without_back = module.load_taskset( - config=module.WikispeediaTasksetConfig(allow_go_back=False) - ) - - assert "go_back" in with_back.system_prompt[0]["content"] - assert "go_back" not in without_back.system_prompt[0]["content"] - assert "Backtracking is disabled" in without_back.system_prompt[0]["content"] - - -@pytest.mark.asyncio -async def test_wikispeedia_tools_resolve_through_v1_runtime( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = load_module(monkeypatch) - wiki = make_small_wiki(module) - monkeypatch.setattr(module, "load_wiki_graph", lambda cache_dir=None: wiki) - env = vf.Env( - taskset=module.load_taskset( - config=module.WikispeediaTasksetConfig( - train_size=2, - eval_size=1, - min_path_length=1, - max_path_length=1, - ) - ), - harness=module.load_harness(config=module.WikispeediaHarnessConfig()), - ) - task = env.taskset.to_task(env.taskset.get_dataset()[0]) - state = module.vf.State.for_task(task) - state = await env.harness.setup_state(task, state) - - tools = state.get_tools() - state["current_article"] = state["info"]["source"] - state["path"] = [state["info"]["source"]] - state["reached_target"] = False - state["links_only"] = False - - result = await tools["click_link"](article=state["info"]["target"]) - - assert sorted(tools) == ["click_link", "go_back"] assert result.startswith("TARGET REACHED") - assert state["reached_target"] is True + assert state.extras["current_article"] == "B" + assert state.extras["path"] == ["A", "B"] + assert state.extras["reached_target"] is True + assert state.stop_condition == "target_reached" @pytest.mark.asyncio -async def test_wikispeedia_langchain_tools_keep_explicit_schema( +async def test_wikispeedia_scores_from_extras_and_transcript( monkeypatch: pytest.MonkeyPatch, ) -> None: - module = load_module(monkeypatch) - fake_langchain_core = types.ModuleType("langchain_core") - fake_tools_module = types.ModuleType("langchain_core.tools") - - def tool(func): - return func - - fake_tools_module.tool = tool - fake_langchain_core.tools = fake_tools_module - monkeypatch.setitem(sys.modules, "langchain_core", fake_langchain_core) - monkeypatch.setitem(sys.modules, "langchain_core.tools", fake_tools_module) - calls = [] - - async def runtime_click_link(**kwargs): - calls.append(kwargs) - return "clicked" - - async def runtime_go_back(): - return "back" - - tools = module.langchain_navigation_tools( - {"click_link": runtime_click_link, "go_back": runtime_go_back} + _, module = load_modules(monkeypatch) + taskset = module.WikispeediaTaskset( + module.WikispeediaTasksetConfig(efficiency_weight=0.5) ) - - assert [tool.__name__ for tool in tools] == ["click_link", "go_back"] - assert list(inspect.signature(tools[0]).parameters) == ["article"] - assert list(inspect.signature(tools[1]).parameters) == [] - assert await tools[0](article="B") == "clicked" - assert calls == [{"article": "B"}] - - -@pytest.mark.asyncio -async def test_wikispeedia_graph_recursion_limit_stops_rollout( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = load_module(monkeypatch) - - class GraphRecursionError(Exception): - pass - - class FakeState(dict): - def get_endpoint_config(self, api: str): - _ = api - - class EndpointConfig: - model = "model" - base_url = "https://example.invalid/v1" - - return EndpointConfig() - - def get_client(self, api: str, *, sync: bool = False): - _ = api, sync - - class Client: - api_key = "key" - - def close(self) -> None: - return None - - return Client() - - def get_tools(self): - return {} - - def get_max_turns(self, default: int): - return default - - def stop(self, reason: str): - self["stop_reason"] = reason - - class FakeChatOpenAI: - def __init__(self, **kwargs): - self.kwargs = kwargs - - class FakeAgent: - async def ainvoke(self, payload, config=None): - raise GraphRecursionError("recursion limit") - - created_system_prompts = [] - - def fake_create_deep_agent(**kwargs): - created_system_prompts.append(kwargs["system_prompt"]) - return FakeAgent() - - fake_deepagents = types.ModuleType("deepagents") - fake_langchain_openai = types.ModuleType("langchain_openai") - fake_langgraph = types.ModuleType("langgraph") - fake_langgraph_errors = types.ModuleType("langgraph.errors") - fake_langchain_core = types.ModuleType("langchain_core") - fake_tools_module = types.ModuleType("langchain_core.tools") - - fake_deepagents.create_deep_agent = fake_create_deep_agent - fake_langchain_openai.ChatOpenAI = FakeChatOpenAI - fake_langgraph_errors.GraphRecursionError = GraphRecursionError - fake_langgraph.errors = fake_langgraph_errors - fake_tools_module.tool = lambda func: func - fake_langchain_core.tools = fake_tools_module - monkeypatch.setitem(sys.modules, "deepagents", fake_deepagents) - monkeypatch.setitem(sys.modules, "langchain_openai", fake_langchain_openai) - monkeypatch.setitem(sys.modules, "langgraph", fake_langgraph) - monkeypatch.setitem(sys.modules, "langgraph.errors", fake_langgraph_errors) - monkeypatch.setitem(sys.modules, "langchain_core", fake_langchain_core) - monkeypatch.setitem(sys.modules, "langchain_core.tools", fake_tools_module) - - program = module.make_langchain_deep_agents_program( - max_turns=50, - timeout_seconds=30, + task = module.WikispeediaTask( + prompt=[{"role": "user", "content": "start"}], + answer="B", + source="A", + target="B", + shortest_path=1, ) - state = FakeState( + state = vf.State(task_id=task.task_id) + state.extras.update( { - "info": {"source": "A"}, - "prompt": [{"role": "user", "content": "start"}], - "system_prompt": [ - {"role": "user", "content": "first prompt chunk"}, - {"role": "system", "content": "second prompt chunk"}, - ], + "path": ["A", "B"], + "reached_target": True, + "agent_timeout": False, + "shortest_path": 1, } ) + state.transcript.append( + vf.Turn( + prompt=[], + completion=[vf.AssistantMessage(content="done")], + tool_calls=[ + vf.ToolCall(id="call_1", name="click_link", arguments='{"article":"B"}') + ], + ) + ) - result = await program({}, state) - - assert created_system_prompts == ["first prompt chunk\n\nsecond prompt chunk"] - assert result["agent_timeout"] is True - assert result["stop_reason"] == "agent_recursion_limit" - assert result["agent_completion"] == [] - - -@pytest.mark.asyncio -async def test_wikispeedia_tool_metrics_use_agent_completion( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = load_module(monkeypatch) - task = vf.Task({"prompt": [], "info": {"shortest_path": 1}}).freeze() - state = vf.State.for_task(task) - state["completion"] = [ - { - "role": "assistant", - "content": "", - "tool_calls": [{"id": "call_1", "name": "click_link", "arguments": "{}"}], - }, - { - "role": "tool", - "tool_call_id": "call_1", - "content": "'C' is not a valid link from 'A'.", - }, - ] - - assert await module.total_tool_calls(task, state) == 1.0 - assert await module.invalid_link_rate(task, state) == 1.0 + assert await taskset.reached_target(state) == 1.0 + assert await taskset.path_efficiency(state) == 1.0 + assert await taskset.path_efficiency_reward(state) == 0.5 + assert await taskset.total_tool_calls(state) == 1.0 diff --git a/tests/test_mcp_search_env.py b/tests/test_mcp_search_env.py index 666e7f5fe4..ee83d71d46 100644 --- a/tests/test_mcp_search_env.py +++ b/tests/test_mcp_search_env.py @@ -1,59 +1,32 @@ -import importlib.util -import inspect -import sys -from pathlib import Path -from typing import Any +from environments.mcp_search_env_v1.mcp_search_env_v1 import taskset as module import pytest -import verifiers as vf - - -def _load_mcp_search_module() -> Any: - module_path = ( - Path(__file__).resolve().parent.parent - / "environments" - / "mcp_search_env" - / "mcp_search_env.py" - ) - spec = importlib.util.spec_from_file_location("test_mcp_search_env", module_path) - assert spec is not None - assert spec.loader is not None - - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - spec.loader.exec_module(module) - return module +import verifiers.v1 as vf +from verifiers.v1.loaders import load_environment_from_components def test_mcp_search_env_is_v1_only() -> None: - module = _load_mcp_search_module() - - env = module.load_environment( - config=module.MCPSearchEnvConfig(taskset={"max_turns": 4}) + env = load_environment_from_components( + module, {"config": {"taskset": {"max_turns": 4}}} ) assert isinstance(env, vf.Env) assert isinstance(env.taskset, vf.Taskset) assert isinstance(env.harness, vf.Harness) - assert "v1" not in inspect.signature(module.load_environment).parameters + assert not hasattr(module, "load_environment") assert not hasattr(module, "load_v1_environment") - assert not (Path(module.__file__).parent / "mcp_search_v1.py").exists() assert env.taskset.config.max_turns == 4 def test_mcp_search_env_preserves_harness_config() -> None: - module = _load_mcp_search_module() - - env = module.load_environment( - config=module.MCPSearchEnvConfig(harness={"max_turns": 7}) + env = load_environment_from_components( + module, {"config": {"harness": {"max_turns": 7}}} ) assert env.harness.config.max_turns == 7 def test_mcp_search_default_taskset_has_stable_non_doc_fixture() -> None: - module = _load_mcp_search_module() - rows = list(module.load_tasks()) assert len(rows) >= 10 @@ -63,27 +36,26 @@ def test_mcp_search_default_taskset_has_stable_non_doc_fixture() -> None: def test_mcp_search_taskset_accepts_v1_taskset_config() -> None: - module = _load_mcp_search_module() - - env = module.load_environment( - config=module.MCPSearchEnvConfig(taskset={"max_turns": 3}), + env = load_environment_from_components( + module, {"config": {"taskset": {"max_turns": 3}}} ) tasks = list(env.taskset) assert env.taskset.config.max_turns == 3 - assert all(task["max_turns"] == 3 for task in tasks) + assert all(task.max_turns == 3 for task in tasks) @pytest.mark.asyncio async def test_mcp_search_reward_handles_missing_assistant() -> None: - module = _load_mcp_search_module() - - task = vf.Task({"answer": "expected"}) - assert await module.exact_title_reward(task, vf.State({"completion": []})) == 0.0 - assert ( - await module.exact_title_reward( - task, - vf.State({"completion": [{"role": "user", "content": "expected"}]}), - ) - == 0.0 + task = module.MCPSearchTask( + query="expected", + question="find expected", + answer="expected", + ) + taskset = module.MCPSearchTaskset(module.MCPSearchTasksetConfig()) + state = vf.State(task_id=task.task_id) + assert await taskset.exact_title_reward(task, state) == 0.0 + state.transcript.append( + vf.Turn(prompt=task.prompt, completion=[vf.UserMessage(content="expected")]) ) + assert await taskset.exact_title_reward(task, state) == 0.0 diff --git a/tests/test_message_utils.py b/tests/test_message_utils.py index 18f91909a6..c032d28f54 100644 --- a/tests/test_message_utils.py +++ b/tests/test_message_utils.py @@ -1,39 +1,29 @@ -from verifiers.types import AssistantMessage, UserMessage -from verifiers.utils.message_utils import ( - from_raw_message, - get_messages, - normalize_messages, -) +from pydantic import TypeAdapter +from verifiers.types import AssistantMessage, Messages, ToolCall -def test_from_raw_message_normalizes_oai_tool_calls(): +MESSAGES_ADAPTER = TypeAdapter(Messages) + + +def test_tool_call_accepts_openai_shape(): raw = { - "role": "assistant", - "content": "calling tool", - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": { - "name": "echo", - "arguments": '{"x": 1}', - }, - } - ], + "id": "call_1", + "type": "function", + "function": { + "name": "echo", + "arguments": '{"x": 1}', + }, } - message = from_raw_message(raw) + tool_call = ToolCall.model_validate(raw) - assert isinstance(message, AssistantMessage) - assert message.tool_calls is not None - assert len(message.tool_calls) == 1 - assert message.tool_calls[0].id == "call_1" - assert message.tool_calls[0].name == "echo" - assert message.tool_calls[0].arguments == '{"x": 1}' + assert tool_call.id == "call_1" + assert tool_call.name == "echo" + assert tool_call.arguments == '{"x": 1}' -def test_normalize_messages_accepts_oai_tool_call_dicts(): - messages = normalize_messages( +def test_messages_adapter_accepts_openai_tool_call_dicts(): + messages = MESSAGES_ADAPTER.validate_python( [ { "role": "assistant", @@ -59,30 +49,3 @@ def test_normalize_messages_accepts_oai_tool_call_dicts(): assert assistant.tool_calls[0].id == "call_2" assert assistant.tool_calls[0].name == "lookup" assert assistant.tool_calls[0].arguments == '{"q": "hello"}' - - -def test_get_messages_returns_typed_messages(): - messages = get_messages( - [ - {"role": "user", "content": "question"}, - {"role": "assistant", "content": "answer"}, - ] - ) - - assert isinstance(messages[0], UserMessage) - assert isinstance(messages[1], AssistantMessage) - assert messages[-1].content == "answer" - - -def test_get_messages_filters_by_role_with_typed_return(): - messages = get_messages( - [ - {"role": "user", "content": "question"}, - {"role": "assistant", "content": "answer"}, - ], - role="assistant", - ) - - assert len(messages) == 1 - assert isinstance(messages[0], AssistantMessage) - assert messages[0].content == "answer" diff --git a/tests/test_opencode_harbor.py b/tests/test_opencode_harbor.py index df6c92d6e4..2115f62832 100644 --- a/tests/test_opencode_harbor.py +++ b/tests/test_opencode_harbor.py @@ -1,40 +1,45 @@ -import importlib.util +import importlib import sys from pathlib import Path -from typing import Any, cast +from typing import Any -import verifiers as vf -from harnesses import OpenCode, OpenCodeConfig, OpenCodeProgramConfig +import pytest +import verifiers.v1 as vf +from harnesses import OpenCode, OpenCodeConfig from tasksets import HarborTaskset +from verifiers.v1.loaders import load_environment_from_components -def _load_opencode_module() -> Any: - module_path = ( - Path(__file__).resolve().parent.parent - / "environments" - / "opencode_harbor" - / "opencode_harbor.py" +def _load_opencode_modules(monkeypatch: pytest.MonkeyPatch) -> tuple[Any, Any]: + env_dir = ( + Path(__file__).resolve().parent.parent / "environments" / "opencode_harbor_v1" ) - spec = importlib.util.spec_from_file_location( - "test_opencode_harbor_module", module_path + monkeypatch.syspath_prepend(str(env_dir)) + for name in ( + "opencode_harbor_v1", + "opencode_harbor_v1.taskset", + "opencode_harbor_v1.harness", + ): + sys.modules.pop(name, None) + return ( + importlib.import_module("opencode_harbor_v1"), + importlib.import_module("opencode_harbor_v1.taskset"), ) - assert spec is not None - assert spec.loader is not None - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - spec.loader.exec_module(module) - return module +def test_load_environment_uses_v1_taskset_and_harness( + monkeypatch: pytest.MonkeyPatch, +) -> None: + package, module = _load_opencode_modules(monkeypatch) -def test_load_environment_uses_v1_taskset_and_harness() -> None: - module = _load_opencode_module() - - env = module.load_environment( - config=vf.EnvConfig( - taskset=module.HarborTasksetConfig(), - harness=module.OpenCodeConfig(), - ) + env = load_environment_from_components( + package, + { + "config": { + "taskset": {}, + "harness": {}, + } + }, ) assert isinstance(env, vf.Env) @@ -43,64 +48,59 @@ def test_load_environment_uses_v1_taskset_and_harness() -> None: assert isinstance(env.harness.config, OpenCodeConfig) assert not hasattr(module, "OpenCodeHarborHarnessConfig") assert not hasattr(module, "TERMINAL_BENCH_SAMPLE_TASKS") - assert env.taskset.config.bundle_package == module.__name__ + assert env.taskset.config.source == "package" + assert env.taskset.config.dataset == "opencode_harbor_v1" task = next(iter(env.taskset)) - assert ( - Path(cast(str, task["task_dir"])).parent - == Path(module.__file__).parent / "tasks" - ) + assert Path(task.task_dir).parent == Path(module.__file__).parent / "tasks" assert env.harness.config.max_turns == 4 - assert env.harness.config.program.disabled_tools == ( - OpenCodeConfig().program.disabled_tools - ) - assert "webfetch" in env.harness.config.program.disabled_tools - assert "question" in env.harness.config.program.disabled_tools - - program = cast(dict[str, object], env.harness.config.program.data()) - mcp_setup = cast(dict[str, object], program["channels"])["mcp"] - assert '"webfetch": false' in cast(str, mcp_setup) - assert '"question": false' in cast(str, mcp_setup) - - -def test_load_environment_accepts_v1_taskset_and_harness_config() -> None: - module = _load_opencode_module() - - env = module.load_environment( - config=vf.EnvConfig( - taskset=module.HarborTasksetConfig( - task_names=["hello-world"], - sandbox=vf.SandboxConfig(cpu_cores=1.5), - ), - harness=module.OpenCodeConfig( - program=OpenCodeProgramConfig( - agent_workdir="/workspace", - disabled_tools=["webfetch"], - ), - max_turns=2, - ), - ) + assert env.harness.config.disabled_tools == OpenCodeConfig().disabled_tools + assert "webfetch" in env.harness.config.disabled_tools + assert "question" in env.harness.config.disabled_tools + + command = env.harness.command(task, vf.State(task_id=task.task_id)) + assert '"webfetch": false' in command[2] + assert '"question": false' in command[2] + + +def test_load_environment_accepts_v1_taskset_and_harness_config( + monkeypatch: pytest.MonkeyPatch, +) -> None: + package, module = _load_opencode_modules(monkeypatch) + + env = load_environment_from_components( + package, + { + "config": { + "taskset": { + "tasks": ["hello-world"], + }, + "harness": { + "cwd": "/workspace", + "disabled_tools": ["webfetch"], + "max_turns": 2, + }, + } + }, ) - assert env.taskset.config.bundle_package == module.__name__ + assert env.taskset.config.source == "package" + assert env.taskset.config.dataset == "opencode_harbor_v1" + assert isinstance(env.harness, OpenCode) task = next(iter(env.taskset)) - assert task["task_dir"] == str( - Path(module.__file__).parent / "tasks" / "hello-world" - ) - assert env.taskset.config.task_names == ["hello-world"] - assert env.taskset.config.sandbox.cpu_cores == 1.5 - assert env.harness.config.program.agent_workdir == "/workspace" + assert task.task_dir == str(Path(module.__file__).parent / "tasks" / "hello-world") + assert env.taskset.config.tasks == ["hello-world"] + assert env.harness.config.cwd == "/workspace" assert env.harness.config.max_turns == 2 - program = cast(dict[str, object], env.harness.config.program.data()) - command = cast(list[object], program["command"]) - mcp_setup = cast(dict[str, object], program["channels"])["mcp"] - assert "/workspace" in cast(str, command[2]) - assert '"webfetch": false' in cast(str, mcp_setup) - assert '"question": false' not in cast(str, mcp_setup) + command = env.harness.command(task, vf.State(task_id=task.task_id)) + assert '"webfetch": false' in command[2] + assert '"question": false' not in command[2] -def test_pyproject_does_not_define_unsupported_harness_defaults() -> None: - module = _load_opencode_module() - pyproject = Path(module.__file__).parent / "pyproject.toml" +def test_pyproject_does_not_define_unsupported_harness_defaults( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, module = _load_opencode_modules(monkeypatch) + pyproject = Path(module.__file__).parents[1] / "pyproject.toml" assert "[tool.verifiers.harness]" not in pyproject.read_text() diff --git a/tests/test_openenv_echo_v1.py b/tests/test_openenv_echo_v1.py new file mode 100644 index 0000000000..a43a06ee71 --- /dev/null +++ b/tests/test_openenv_echo_v1.py @@ -0,0 +1,29 @@ +from pathlib import Path +import importlib.util + +import pytest + + +@pytest.mark.asyncio +async def test_openenv_echo_async_step_sets_tool_reward() -> None: + pytest.importorskip("openenv") + from openenv.core.env_server.mcp_types import CallToolAction + + path = ( + Path(__file__).parents[1] + / "environments/openenv_echo_v1/openenv_echo_v1/proj/server/echo_environment.py" + ) + spec = importlib.util.spec_from_file_location("openenv_echo_environment", path) + assert spec is not None + assert spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + observation = await module.EchoEnvironment().step_async( + CallToolAction( + tool_name="echo_message", + arguments={"message": "hello from openenv"}, + ) + ) + + assert observation.reward == pytest.approx(1.8) diff --git a/tests/test_trajectory_processing.py b/tests/test_trajectory_processing.py index 386e4fd947..d0080fb5cb 100644 --- a/tests/test_trajectory_processing.py +++ b/tests/test_trajectory_processing.py @@ -248,7 +248,7 @@ async def test_parsed_prompt_attribution_survives_v1_assert_serializable(): """ from renderers.base import RenderedTokens - from verifiers.v1.utils.serialization_utils import serializable + from verifiers.v1.utils.json_utils import jsonable response = Response( id="t", @@ -278,7 +278,7 @@ async def test_parsed_prompt_attribution_survives_v1_assert_serializable(): ), ) parsed = await parse_response_tokens(response) - step = {"tokens": serializable(parsed), "response": serializable(response)} + step = {"tokens": jsonable(parsed), "response": jsonable(response)} State({"trajectory": [step]}).assert_serializable() diff --git a/tests/test_v1_bfcl.py b/tests/test_v1_bfcl.py deleted file mode 100644 index cd8ad56f74..0000000000 --- a/tests/test_v1_bfcl.py +++ /dev/null @@ -1,178 +0,0 @@ -import importlib.util -import sys -from pathlib import Path -from types import ModuleType - -import pytest - -import verifiers as vf - - -def load_bfcl_module() -> ModuleType: - path = Path(__file__).parents[1] / "environments" / "bfcl_v3" / "bfcl_v3.py" - spec = importlib.util.spec_from_file_location("bfcl_v3_test_module", path) - assert spec is not None - assert spec.loader is not None - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - spec.loader.exec_module(module) - return module - - -def test_bfcl_prefers_hinted_function_schemas() -> None: - bfcl = load_bfcl_module() - task = { - "function": [{"name": "plain"}], - "function_with_hints": [{"name": "hinted"}], - } - - assert bfcl.bfcl_functions(task) == [{"name": "hinted"}] - - -def test_bfcl_prefers_hinted_holdout_function_schemas() -> None: - bfcl = load_bfcl_module() - task = { - "missed_function": {"1": [{"name": "plain"}]}, - "missed_function_with_hints": {"1": [{"name": "hinted"}]}, - } - - assert bfcl.bfcl_missed_function(task) == {"1": [{"name": "hinted"}]} - - -def test_bfcl_row_preserves_hinted_holdout_functions() -> None: - bfcl = load_bfcl_module() - entry = { - "id": "case", - "question": ["call tools"], - "function": [{"name": "plain"}], - "missed_function": {"1": [{"name": "plain_holdout"}]}, - } - hinted_entry = { - "function": [{"name": "hinted"}], - "missed_function": {"1": [{"name": "hinted_holdout"}]}, - } - - row = bfcl.bfcl_row("multi_turn", entry, hinted_entry, None) - - assert row["function_with_hints"] == [{"name": "hinted"}] - assert row["missed_function"] == {"1": [{"name": "plain_holdout"}]} - assert row["missed_function_with_hints"] == {"1": [{"name": "hinted_holdout"}]} - assert "toolsets" not in row - - -@pytest.mark.asyncio -async def test_bfcl_single_turn_tools_are_rollout_scoped( - monkeypatch: pytest.MonkeyPatch, -) -> None: - bfcl = load_bfcl_module() - utils_module = ModuleType("bfcl_eval.utils") - utils_module.is_multi_turn = lambda category: False - monkeypatch.setitem(sys.modules, "bfcl_eval", ModuleType("bfcl_eval")) - monkeypatch.setitem(sys.modules, "bfcl_eval.utils", utils_module) - monkeypatch.setattr(bfcl, "patch_bfcl_eval", lambda: None) - monkeypatch.setattr( - bfcl, - "bfcl_tool_defs", - lambda functions: [ - vf.Tool( - name=str(functions[0]["name"]), - description="", - parameters={"type": "object", "properties": {}}, - ) - ], - ) - taskset = bfcl.BFCLTaskset(config=bfcl.BFCLTasksetConfig(examples_per_category=0)) - env = vf.Env(taskset=taskset) - task = vf.Task( - { - "prompt": [{"role": "user", "content": "call tool"}], - "category": "simple_python", - "function": [{"name": "lookup"}], - } - ).freeze() - - state = await env.harness.setup_state(task, vf.State.for_task(task)) - await env.harness.runtime.setup_rollout(task, state) - - assert list(taskset.named_toolsets) == ["bfcl"] - assert state["tools"] == ["lookup"] - - -def test_bfcl_empty_completion_has_no_tool_calls() -> None: - bfcl = load_bfcl_module() - - assert bfcl.assistant_tool_calls({"completion": []}) == [] - assert ( - bfcl.assistant_tool_calls( - {"completion": [{"role": "user", "content": "no assistant"}]} - ) - == [] - ) - - -def test_bfcl_public_loader_is_v1_only(monkeypatch: pytest.MonkeyPatch) -> None: - bfcl = load_bfcl_module() - seen_harness_config: vf.HarnessConfig | None = None - - def fake_harness(config: vf.HarnessConfig) -> vf.Harness: - nonlocal seen_harness_config - seen_harness_config = config - return vf.Harness(config=config) - - monkeypatch.setattr(bfcl, "load_harness", fake_harness) - - env = bfcl.load_environment( - config=bfcl.BFCLEnvConfig( - taskset=bfcl.BFCLTasksetConfig( - test_category="simple_python", - examples_per_category=0, - ), - harness=bfcl.BFCLHarnessConfig(), - ) - ) - - assert isinstance(env, vf.Env) - seen_taskset_config = env.taskset.config - assert isinstance(seen_taskset_config, bfcl.BFCLTasksetConfig) - assert isinstance(seen_harness_config, bfcl.BFCLHarnessConfig) - assert seen_taskset_config.test_category == "simple_python" - assert seen_taskset_config.examples_per_category == 0 - assert "rewards" not in seen_taskset_config.model_fields_set - assert [reward.__name__ for reward in env.taskset.rewards] == ["bfcl_reward"] - assert seen_harness_config.test_category == "simple_python" - assert not hasattr(bfcl, "load_v1_environment") - - -def test_bfcl_loader_supports_category_groups( - monkeypatch: pytest.MonkeyPatch, -) -> None: - bfcl = load_bfcl_module() - seen_harness_categories = [] - - def fake_load_tasks(test_category: str, **kwargs: object): - _ = kwargs - return [{"question": test_category, "answer": "a"}] - - def fake_harness(config: vf.HarnessConfig) -> vf.Harness: - assert isinstance(config, bfcl.BFCLHarnessConfig) - seen_harness_categories.append(config.test_category) - return vf.Harness(config=config) - - monkeypatch.setattr(bfcl, "load_tasks", fake_load_tasks) - monkeypatch.setattr(bfcl, "load_harness", fake_harness) - - env = bfcl.load_environment( - config=bfcl.BFCLEnvConfig( - taskset=bfcl.BFCLTasksetConfig( - test_categories=["simple_python", "simple_java"], - examples_per_category=0, - ), - harness=bfcl.BFCLHarnessConfig(), - ) - ) - - assert isinstance(env, vf.EnvGroup) - assert env.env_names == ["simple_python", "simple_java"] - seen_taskset_categories = [item.taskset.config.test_category for item in env.envs] - assert seen_taskset_categories == ["simple_python", "simple_java"] - assert seen_harness_categories == ["simple_python", "simple_java"] diff --git a/tests/test_v1_config_extension.py b/tests/test_v1_config_extension.py deleted file mode 100644 index 955871d290..0000000000 --- a/tests/test_v1_config_extension.py +++ /dev/null @@ -1,3707 +0,0 @@ -import importlib -import sys -import types -from typing import Any, cast - -import pytest -from datasets import Dataset -from pydantic import ValidationError - -import verifiers as vf -from verifiers import ( - Config, - Env, - EnvConfig, - Harness, - HarnessConfig, - State, - Task, - Taskset, - TasksetConfig, - Toolset, -) -from verifiers.v1.toolset import normalize_toolset -from verifiers.v1.types import ModelClient -from harnesses import OpenCode -from verifiers.utils.import_utils import load_toml -from verifiers.v1.utils.config_utils import coerce_config, explicit_config_data - - -REF_MODULE = "v1_config_extension_refs" -EVAL_FALLBACK_WARNING = "eval_dataset is not set, falling back to train dataset" - - -def load_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - answer = "eval ok" if split == "eval" else "ok" - return [ - { - "prompt": [{"role": "user", "content": f"Say {answer}."}], - "answer": answer, - } - ] - - -def load_other_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - _ = split - return [ - { - "prompt": [{"role": "user", "content": "Say other ok."}], - "answer": "other ok", - } - ] - - -def load_dataset_tasks(split: vf.TaskSplit = "train") -> Dataset: - _ = split - return Dataset.from_list( - [ - { - "prompt": [{"role": "user", "content": "Say dataset ok."}], - "answer": "dataset ok", - } - ] - ) - - -def load_system_prompt() -> vf.SystemPrompt: - return "loaded system prompt" - - -@vf.metric -async def config_metric(task: dict[str, object], state: dict[str, object]) -> float: - return float(task.get("answer") == "ok" and state.get("answer") == "ok") - - -@vf.reward(weight=0.25) -async def config_reward(task: dict[str, object], state: dict[str, object]) -> float: - return float(task.get("answer") == state.get("answer")) - - -@vf.metric(stage="group") -async def group_config_metric( - tasks: list[dict[str, object]], states: list[dict[str, object]] -) -> list[float]: - _ = tasks - return [float(index) for index, _ in enumerate(states)] - - -@vf.reward(stage="group") -async def group_config_reward( - tasks: list[dict[str, object]], states: list[dict[str, object]] -) -> list[float]: - return [ - float(task.get("answer") == state.get("answer")) - for task, state in zip(tasks, states) - ] - - -@vf.advantage -async def config_advantage( - tasks: list[dict[str, object]], states: list[dict[str, object]] -) -> list[float]: - _ = tasks - return [float(index) for index, _ in enumerate(states)] - - -@vf.cleanup(priority=5) -async def config_cleanup(task: dict[str, object], state: dict[str, object]) -> None: - state["cleaned"] = True - - -@vf.cleanup(stage="group", priority=5) -async def config_group_cleanup( - tasks: list[dict[str, object]], states: list[dict[str, object]] -) -> None: - _ = tasks - for state in states: - state["group_cleaned"] = True - - -TEARDOWN_EVENTS: list[str] = [] - - -@vf.teardown -async def config_taskset_teardown() -> None: - TEARDOWN_EVENTS.append("taskset") - - -@vf.teardown -async def config_harness_teardown() -> None: - TEARDOWN_EVENTS.append("harness") - - -@vf.setup(priority=5) -async def config_setup(task: dict[str, object], state: dict[str, object]) -> None: - _ = task - state["configured_setup"] = True - cast(list[str], state.setdefault("setup_order", [])).append("config_setup") - - -@vf.update(priority=5) -async def config_update(task: dict[str, object], state: dict[str, object]) -> None: - _ = task - state["updated"] = True - - -@vf.reward -async def updated_reward(task: dict[str, object], state: dict[str, object]) -> float: - _ = task - return float(state.get("updated") is True) - - -@vf.update(stage="group") -async def config_group_update( - tasks: list[dict[str, object]], states: list[dict[str, object]] -) -> None: - _ = tasks - for state in states: - state["group_updated"] = True - - -@vf.update(stage="group") -async def bad_group_update(tasks, states, extra) -> None: - _ = tasks, states, extra - - -@vf.reward(stage="group") -async def group_updated_reward( - tasks: list[dict[str, object]], states: list[dict[str, object]] -) -> list[float]: - _ = tasks - return [float(state.get("group_updated") is True) for state in states] - - -async def config_tool(query: str, prefix: str) -> str: - return f"{prefix}:{query}" - - -async def direct_tool() -> str: - return "direct" - - -async def hidden_tool() -> str: - return "hidden" - - -async def object_tool(value: str, box: dict[str, object]) -> str: - values = cast(list[str], box.setdefault("values", [])) - values.append(value) - return value - - -async def object_prefix_tool(value: str, box: dict[str, object]) -> str: - prefix = box["prefix"] - assert isinstance(prefix, str) - return f"{prefix}:{value}" - - -def load_object_box() -> dict[str, object]: - return {"values": []} - - -closed_objects: list[str] = [] - - -class ClosableObject: - def __init__(self, name: str): - self.name = name - - async def close(self) -> None: - closed_objects.append(self.name) - - -def load_rollout_closable_object() -> ClosableObject: - return ClosableObject("rollout") - - -def load_prefixed_object_box(prefix: str) -> dict[str, object]: - return {"prefix": prefix} - - -async def update_from_binding( - task: dict[str, object], state: dict[str, object], expected: str -) -> None: - _ = task - state["expected"] = expected - - -@vf.update(stage="group") -async def group_update_from_binding( - tasks: list[dict[str, object]], states: list[dict[str, object]], expected: str -) -> None: - _ = tasks - for state in states: - state["group_expected"] = expected - - -@vf.reward -async def reward_from_binding( - task: dict[str, object], state: dict[str, object], expected: str -) -> float: - _ = state - return float(task.get("answer") == expected) - - -@vf.reward(stage="group") -async def group_reward_from_binding( - tasks: list[dict[str, object]], states: list[dict[str, object]], expected: str -) -> list[float]: - _ = tasks - return [float(state.get("answer") == expected) for state in states] - - -async def colliding_tool(value: str, token: str) -> str: - return f"{token}:{value}" - - -async def colliding_update( - task: dict[str, object], state: dict[str, object], token: str -) -> None: - _ = task - state["colliding_update_token"] = token - - -colliding_update.__name__ = "colliding_tool" - - -def dynamic_tool(task: dict[str, object]) -> vf.Tool: - tool = task["dynamic_tool"] - if not isinstance(tool, dict): - raise TypeError("dynamic_tool must be a mapping.") - tool = cast(dict[str, Any], tool) - return vf.Tool( - name=str(tool["name"]), - description=str(tool["description"]), - parameters=dict(cast(dict[str, Any], tool["parameters"])), - ) - - -async def dynamic_tool_handler( - state: dict[str, object], tool: vf.Tool, arguments: dict[str, object] -) -> str: - calls = cast(list[object], state.setdefault("dynamic_tool_calls", [])) - calls.append({tool.name: arguments}) - return "recorded" - - -async def config_user( - task: dict[str, object], state: dict[str, object] -) -> list[dict[str, str]]: - _ = task - if state.get("user_called"): - return [] - state["user_called"] = True - return [{"role": "user", "content": "continue"}] - - -def token_factory() -> str: - return "secret-token" - - -async def config_user_with_bindings( - task: dict[str, object], - state: dict[str, object], - token: str, - messages: list[object], -) -> list[dict[str, str]]: - _ = task - state["token_seen"] = token - state["messages_len"] = len(messages) - return [{"role": "user", "content": token}] - - -async def direct_user_with_messages( - task: dict[str, object], - state: dict[str, object], - messages: list[object], -) -> list[dict[str, str]]: - _ = task - state["direct_messages_len"] = len(messages) - return [{"role": "user", "content": "continue"}] - - -async def sandbox_user( - task: dict[str, object], state: dict[str, object], sandbox: object -) -> list[dict[str, str]]: - _ = task - state["sandbox_seen"] = sandbox - return [{"role": "user", "content": "sandbox ok"}] - - -class ConfigUserConfig(vf.UserConfig): - pass - - -class ConfigUser(vf.User[ConfigUserConfig]): - async def get_response( - self, task: dict[str, object], state: dict[str, object] - ) -> list[dict[str, str]]: - return await config_user(task, state) - - -class ConfigUserWithBindingsConfig(vf.UserConfig): - pass - - -class ConfigUserWithBindings(vf.User[ConfigUserWithBindingsConfig]): - async def get_response( - self, - task: dict[str, object], - state: dict[str, object], - token: str, - messages: list[object], - ) -> list[dict[str, str]]: - return await config_user_with_bindings(task, state, token, messages) - - -class DirectUserWithMessagesConfig(vf.UserConfig): - pass - - -class DirectUserWithMessages(vf.User[DirectUserWithMessagesConfig]): - async def get_response( - self, - task: dict[str, object], - state: dict[str, object], - messages: list[object], - ) -> list[dict[str, str]]: - return await direct_user_with_messages(task, state, messages) - - -class SandboxUserConfig(vf.UserConfig): - pass - - -class SandboxUser(vf.User[SandboxUserConfig]): - async def get_response( - self, task: dict[str, object], state: dict[str, object], sandbox: object - ) -> list[dict[str, str]]: - return await sandbox_user(task, state, sandbox) - - -async def config_program( - task: dict[str, object], state: dict[str, object] -) -> dict[str, object]: - state["answer"] = task["answer"] - return {"program": "ran"} - - -async def setup_aware_program( - task: dict[str, object], state: dict[str, object] -) -> dict[str, object]: - _ = task - if state.get("configured_setup") is not True: - raise AssertionError("setup did not run before program") - return {"program_saw_setup": True} - - -def config_toolset(prefix: str = "cfg") -> Toolset: - def prefix_value() -> str: - return prefix - - return Toolset( - tools=[config_tool], - bindings={"config_tool.prefix": prefix_value}, - ) - - -def load_another_harness_config() -> HarnessConfig: - return HarnessConfig(max_turns=6, rewards=[ref("config_reward")]) - - -ref_module = types.ModuleType(REF_MODULE) -setattr(ref_module, "load_tasks", load_tasks) -setattr(ref_module, "load_other_tasks", load_other_tasks) -setattr(ref_module, "load_dataset_tasks", load_dataset_tasks) -setattr(ref_module, "load_system_prompt", load_system_prompt) -setattr(ref_module, "config_metric", config_metric) -setattr(ref_module, "group_config_metric", group_config_metric) -setattr(ref_module, "config_reward", config_reward) -setattr(ref_module, "group_config_reward", group_config_reward) -setattr(ref_module, "config_advantage", config_advantage) -setattr(ref_module, "config_cleanup", config_cleanup) -setattr(ref_module, "config_group_cleanup", config_group_cleanup) -setattr(ref_module, "config_taskset_teardown", config_taskset_teardown) -setattr(ref_module, "config_harness_teardown", config_harness_teardown) -setattr(ref_module, "config_setup", config_setup) -setattr(ref_module, "config_update", config_update) -setattr(ref_module, "updated_reward", updated_reward) -setattr(ref_module, "config_group_update", config_group_update) -setattr(ref_module, "bad_group_update", bad_group_update) -setattr(ref_module, "group_updated_reward", group_updated_reward) -setattr(ref_module, "config_tool", config_tool) -setattr(ref_module, "config_toolset", config_toolset) -setattr(ref_module, "dynamic_tool", dynamic_tool) -setattr(ref_module, "direct_tool", direct_tool) -setattr(ref_module, "hidden_tool", hidden_tool) -setattr(ref_module, "object_tool", object_tool) -setattr(ref_module, "load_object_box", load_object_box) -setattr(ref_module, "load_rollout_closable_object", load_rollout_closable_object) -setattr(ref_module, "load_prefixed_object_box", load_prefixed_object_box) -setattr(ref_module, "reward_from_binding", reward_from_binding) -setattr(ref_module, "group_reward_from_binding", group_reward_from_binding) -setattr(ref_module, "update_from_binding", update_from_binding) -setattr(ref_module, "group_update_from_binding", group_update_from_binding) -setattr(ref_module, "colliding_update", colliding_update) -setattr(ref_module, "config_user", config_user) -setattr(ref_module, "direct_user_with_messages", direct_user_with_messages) -setattr(ref_module, "token_factory", token_factory) -setattr(ref_module, "config_user_with_bindings", config_user_with_bindings) -setattr(ref_module, "sandbox_user", sandbox_user) -setattr(ref_module, "config_program", config_program) -setattr(ref_module, "setup_aware_program", setup_aware_program) -setattr(ref_module, "load_another_harness_config", load_another_harness_config) -sys.modules[REF_MODULE] = ref_module - - -def test_explicit_config_data_preserves_explicit_none_values() -> None: - class NestedConfig(Config): - sandbox: vf.SandboxConfig | None = None - - class OuterConfig(Config): - nested: NestedConfig = NestedConfig() - label: str | None = "default" - - config = OuterConfig( - nested=NestedConfig( - sandbox=vf.SandboxConfig(image="python:3.12-slim", workdir=None) - ), - label=None, - ) - - assert explicit_config_data(config) == { - "nested": {"sandbox": {"image": "python:3.12-slim", "workdir": None}}, - "label": None, - } - - -def test_toolset_mapping_uses_explicit_show_hide_lists() -> None: - shown = normalize_toolset( - { - "tools": [ref("direct_tool"), ref("hidden_tool")], - "show": ["direct_tool"], - } - ) - hidden = normalize_toolset( - { - "tools": [ref("direct_tool"), ref("hidden_tool")], - "hide": ["hidden_tool"], - } - ) - - assert shown.show == ("direct_tool",) - assert hidden.hide == ("hidden_tool",) - - -def test_inline_toolset_mapping_rejects_non_boolean_write() -> None: - with pytest.raises(ValidationError): - normalize_toolset({"tools": [ref("direct_tool")], "write": "false"}) - - -def test_inline_toolset_mapping_rejects_unknown_keys() -> None: - with pytest.raises(ValueError, match="Unknown toolset config keys"): - normalize_toolset({"tools": [ref("direct_tool")], "bindngs": {}}) - - -def ref(name: str) -> str: - return f"{REF_MODULE}:{name}" - - -def has_runtime_toolset(value: object) -> bool: - if isinstance(value, Toolset): - return True - if isinstance(value, dict): - return any(has_runtime_toolset(item) for item in value.values()) - if isinstance(value, list | tuple): - return any(has_runtime_toolset(item) for item in value) - return False - - -class ConfigExtensionTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -class ConfigExtensionDatasetTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_dataset_tasks(split) - - -def make_taskset(config: object | None = None, **values: object) -> Taskset: - base_config = coerce_config(TasksetConfig, config) - data = {**base_config.model_dump(exclude_none=True, exclude_unset=True), **values} - runtime_toolsets = data.pop("toolsets", None) - if runtime_toolsets is not None and not has_runtime_toolset(runtime_toolsets): - data["toolsets"] = runtime_toolsets - runtime_toolsets = None - taskset = ConfigExtensionTaskset(config=coerce_config(type(base_config), data)) - if runtime_toolsets is not None: - taskset.add_toolset(runtime_toolsets) - return taskset - - -def make_harness(config: object | None = None, **values: object) -> Harness: - base_config = coerce_config(HarnessConfig, config) - data = {**base_config.model_dump(exclude_none=True, exclude_unset=True), **values} - runtime_client = data.pop("client", None) - model_value = data.pop("model", None) - sampling_args = data.pop("sampling_args", None) - if model_value is not None or sampling_args is not None: - if model_value is None: - model_data: dict[str, object] = {} - elif isinstance(model_value, str): - model_data = {"name": model_value} - elif isinstance(model_value, vf.ModelConfig): - model_data = model_value.model_dump(exclude_none=True, exclude_unset=True) - elif isinstance(model_value, dict): - model_data = dict(model_value) - else: - raise TypeError("test harness model config must be a mapping.") - if sampling_args is not None: - model_data["sampling_args"] = sampling_args - data["model"] = model_data - runtime_toolsets = data.pop("toolsets", None) - if runtime_toolsets is not None and not has_runtime_toolset(runtime_toolsets): - data["toolsets"] = runtime_toolsets - runtime_toolsets = None - harness = Harness(config=coerce_config(type(base_config), data)) - if runtime_client is not None: - harness.model_client = cast(ModelClient, runtime_client) - if runtime_toolsets is not None: - harness.add_toolset(runtime_toolsets) - return harness - - -def test_taskset_config_extends_constructor_surface() -> None: - class ConfiguredTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - taskset = ConfiguredTaskset( - config={ - "taskset_id": "configured", - "metrics": [ref("config_metric")], - "rewards": [ref("config_reward")], - "advantages": [ref("config_advantage")], - "setups": [ref("config_setup")], - "cleanups": [ref("config_cleanup")], - "teardowns": [ref("config_taskset_teardown")], - "toolsets": [ - { - "tools": [ref("config_tool")], - "bindings": {"config_tool.prefix": "task.answer"}, - } - ], - "user": ConfigUserConfig(), - } - ) - - eval_rows = taskset.get_eval_dataset() - task = next(iter(taskset)) - - assert task["taskset_id"] == "configured" - assert task["answer"] == "ok" - assert eval_rows[0]["answer"] == "eval ok" - assert taskset.metrics == [config_metric] - assert taskset.rewards == [config_reward] - assert taskset.advantages == [config_advantage] - assert taskset.setups == [config_setup] - assert taskset.cleanups == [config_cleanup] - assert taskset.teardowns == [config_taskset_teardown] - assert taskset.user is not None - assert len(taskset.toolsets) == 1 - assert taskset.toolsets[0].tools == (config_tool,) - assert taskset.toolsets[0].bindings == {"config_tool.prefix": "task.answer"} - - -def test_taskset_to_task_normalizes_task_input() -> None: - taskset = Taskset(config={"taskset_id": "configured"}) - original = Task({"prompt": [], "taskset_id": "original"}) - - task = taskset.to_task(original) - - assert task is not original - assert task["taskset_id"] == "configured" - assert task.frozen - assert not original.frozen - - -def test_taskset_to_task_copies_frozen_task_input() -> None: - taskset = Taskset(config={"taskset_id": "configured"}) - original = Task({"prompt": [], "taskset_id": "original"}).freeze() - - task = taskset.to_task(original) - - assert task is not original - assert task["taskset_id"] == "configured" - assert task.frozen - assert original.frozen - - -def test_user_config_rejects_string_refs() -> None: - with pytest.raises(ValidationError): - TasksetConfig(user=ref("config_user")) - - with pytest.raises(ValidationError): - HarnessConfig(user=ref("config_user")) - - -def test_config_refs_resolve_from_config_module( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "v1_relative_config_refs" - module = types.ModuleType(module_name) - monkeypatch.setitem(sys.modules, module_name, module) - exec( - """ -import verifiers as vf - - -def load_tasks(split: vf.TaskSplit = "train") -> vf.Tasks: - answer = "eval" if split == "eval" else "train" - return [{"prompt": [], "answer": answer}] - - -@vf.reward -async def exact_answer(task: vf.Task, state: vf.State) -> float: - return 1.0 - - -@vf.metric -async def metric_fn(task: vf.Task, state: vf.State) -> float: - return 1.0 - - -def local_tool() -> str: - return "ok" - - -def load_toolset(prefix: str) -> vf.Toolset: - _ = prefix - return vf.Toolset(tools=[local_tool]) - - -class LocalUserConfig(vf.UserConfig): - pass - - -class LocalUser(vf.User[LocalUserConfig]): - async def get_response(self, task: vf.Task, state: vf.State) -> list[dict[str, str]]: - return [] - - -async def program_fn(task: vf.Task, state: vf.State) -> vf.State: - state["program"] = "ok" - return state - - -def load_system_prompt() -> vf.SystemPrompt: - return "loaded system prompt" - - -class LocalTasksetConfig(vf.TasksetConfig): - user: LocalUserConfig | None = LocalUserConfig() - rewards: list[str] = ["exact_answer"] - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate({"loader": "load_tasks"}) - toolsets: dict[str, dict[str, object]] = { - "local": {"fn": "load_toolset", "prefix": "cfg"} - } - - -class LocalTaskset(vf.Taskset[LocalTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -class LocalHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="program_fn") - metrics: list[str] = ["metric_fn"] - - -def load_taskset(config: LocalTasksetConfig) -> vf.Taskset: - return LocalTaskset(config=config) - - -def load_harness(config: LocalHarnessConfig) -> vf.Harness: - return vf.Harness(config=config) -""", - module.__dict__, - ) - - taskset = vf.load_taskset(module_name, config={}) - harness = vf.load_harness(module_name, config={}) - - assert taskset.get_dataset()[0]["answer"] == "train" - assert taskset.get_eval_dataset()[0]["answer"] == "eval" - assert getattr(taskset.rewards[0], "__name__") == "exact_answer" - assert taskset.user is not None - assert isinstance(taskset.user, module.LocalUser) - assert taskset.objects["loader"] is module.load_tasks - assert taskset.named_toolsets["local"].tools == (module.local_tool,) - assert harness.config.program.data() == {"fn": "program_fn"} - assert getattr(harness.metrics[-1], "__name__") == "metric_fn" - assert callable(harness.program) - - -@pytest.mark.asyncio -async def test_sandbox_program_ref_uses_config_module( - monkeypatch: pytest.MonkeyPatch, -) -> None: - captured: dict[str, object] = {} - - async def fake_run_sandbox_python_program( - *, - program: dict[str, object], - sandbox_config: dict[str, object], - task: Task, - state: State, - runtime: object, - mode: str, - fn_ref: str | None, - max_turns: int, - ) -> State: - _ = program, sandbox_config, task, runtime, mode, max_turns - captured["fn_ref"] = fn_ref - return state - - monkeypatch.setattr( - "verifiers.v1.harness.run_sandbox_python_program", - fake_run_sandbox_python_program, - ) - - class SandboxHarnessConfig(HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="program_fn", sandbox=True) - sandbox: vf.SandboxConfig = vf.SandboxConfig() - - harness = Harness(config=SandboxHarnessConfig()) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - result = await harness.program(task, state) - - assert result is state - assert captured["fn_ref"] == f"{__name__}:program_fn" - - -def test_taskset_get_eval_dataset_uses_load_tasks_eval_split() -> None: - taskset = make_taskset() - - assert taskset.get_dataset()[0]["answer"] == "ok" - assert taskset.get_eval_dataset()[0]["answer"] == "eval ok" - - -def test_taskset_load_tasks_accepts_train_and_eval_splits() -> None: - class SplitTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - answer = "eval" if split == "eval" else "train" - return [{"prompt": [], "answer": answer}] - - taskset = SplitTaskset() - - assert taskset.load_tasks(split="train")[0]["answer"] == "train" - assert taskset.load_tasks(split="eval")[0]["answer"] == "eval" - - -def test_taskset_get_eval_dataset_returns_empty_dataset_for_empty_eval_split() -> None: - class TrainOnlyTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "eval": - return [] - return [{"prompt": [], "answer": "train"}] - - taskset = TrainOnlyTaskset() - - assert taskset.get_dataset()[0]["answer"] == "train" - assert len(taskset.get_eval_dataset()) == 0 - - -def test_empty_eval_split_uses_environment_train_fallback(monkeypatch) -> None: - class EmptyEvalTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "eval": - return [] - return [{"prompt": [], "answer": "train"}] - - env = Env(taskset=EmptyEvalTaskset()) - warnings: list[str] = [] - monkeypatch.setattr(env.logger, "warning", warnings.append) - - assert len(env.taskset.get_eval_dataset()) == 0 - eval_rows = env.get_eval_dataset() - - assert eval_rows[0]["answer"] == "train" - assert warnings == [EVAL_FALLBACK_WARNING] - - -def test_empty_train_split_uses_environment_eval_only_dataset() -> None: - class EmptyTrainTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "train": - return [] - return [{"prompt": [], "answer": "eval"}] - - env = Env(taskset=EmptyTrainTaskset()) - - assert env.get_eval_dataset()[0]["answer"] == "eval" - with pytest.raises(ValueError, match="dataset is not set"): - env.get_dataset() - - -def test_empty_train_and_eval_splits_match_environment_missing_dataset_error( - monkeypatch, -) -> None: - env = Env(taskset=Taskset()) - warnings: list[str] = [] - monkeypatch.setattr(env.logger, "warning", warnings.append) - - with pytest.raises(ValueError, match="dataset is not set"): - env.get_dataset() - with pytest.raises(ValueError, match="dataset is not set"): - env.get_eval_dataset() - assert warnings == [EVAL_FALLBACK_WARNING] - - -def test_empty_taskset_splits_are_checked_once_by_environment() -> None: - calls = {"train": 0, "eval": 0} - - class EmptyTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "train": - calls["train"] += 1 - else: - calls["eval"] += 1 - return [] - - env = Env(taskset=EmptyTaskset()) - - for _ in range(2): - with pytest.raises(ValueError, match="dataset is not set"): - env.get_dataset() - for _ in range(2): - with pytest.raises(ValueError, match="dataset is not set"): - env.get_eval_dataset() - - assert calls == {"train": 1, "eval": 1} - - -def test_env_passes_taskset_eval_dataset_to_environment() -> None: - env = Env( - taskset=make_taskset(), - harness=make_harness(program={"fn": ref("config_program")}), - ) - - assert env.get_dataset()[0]["answer"] == "ok" - assert env.get_eval_dataset()[0]["answer"] == "eval ok" - - -def test_env_defaults_to_base_harness() -> None: - taskset = make_taskset() - env = Env(taskset=taskset) - - assert isinstance(env.harness, Harness) - assert env.harness.taskset is taskset - assert env.get_dataset()[0]["answer"] == "ok" - - -def test_env_capabilities_follow_v1_group_runtime_signals() -> None: - rollout_env = Env( - taskset=make_taskset(rewards=[ref("config_reward")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - group_metric_env = Env( - taskset=make_taskset(metrics=[ref("group_config_metric")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - group_reward_env = Env( - taskset=make_taskset(rewards=[ref("group_config_reward")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - advantage_env = Env( - taskset=make_taskset(advantages=[ref("config_advantage")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - - assert not rollout_env.requires_group_rollouts - assert not rollout_env.provides_advantages - assert group_metric_env.requires_group_rollouts - assert not group_metric_env.provides_advantages - assert group_reward_env.requires_group_rollouts - assert not group_reward_env.provides_advantages - assert advantage_env.requires_group_rollouts - assert advantage_env.provides_advantages - - -def test_env_capabilities_follow_group_lifecycle_handlers() -> None: - group_update_env = Env( - taskset=make_taskset(updates=[ref("config_group_update")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - group_cleanup_env = Env( - taskset=make_taskset(cleanups=[ref("config_group_cleanup")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - - assert group_update_env.requires_group_rollouts - assert not group_update_env.provides_advantages - assert group_cleanup_env.requires_group_rollouts - assert not group_cleanup_env.provides_advantages - - -@pytest.mark.asyncio -async def test_group_lifecycle_handlers_require_bound_extra_args() -> None: - env = Env( - taskset=make_taskset(updates=[ref("bad_group_update")]), - harness=make_harness(program={"fn": ref("config_program")}), - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - with pytest.raises(TypeError, match="extra"): - await env.harness.runtime.update_group([task], [state]) - - -def test_env_capabilities_follow_custom_taskset_init_group() -> None: - class GroupSetupTaskset(Taskset): - async def init_group( - self, task: Task, num_rollouts: int - ) -> tuple[list[Task], list[State]]: - return await super().init_group(task, num_rollouts) - - env = Env( - taskset=GroupSetupTaskset(config=TasksetConfig()), - harness=make_harness(program={"fn": ref("config_program")}), - ) - - assert env.requires_group_rollouts - assert not env.provides_advantages - - -def test_harness_config_extends_constructor_surface() -> None: - direct_toolset = Toolset(tools=[direct_tool]) - harness = Harness( - config={ - "program": {"fn": ref("config_program")}, - "metrics": [ref("config_metric")], - "rewards": [ref("config_reward")], - "advantages": [ref("config_advantage")], - "setups": [ref("config_setup")], - "cleanups": [ref("config_cleanup")], - "teardowns": [ref("config_harness_teardown")], - "toolsets": [ - { - "tools": [ref("config_tool")], - "hide": ["config_tool"], - } - ], - "user": ConfigUserConfig(), - "max_turns": 3, - }, - ) - harness.add_toolset(direct_toolset) - - assert harness.config.program.data() == {"fn": ref("config_program")} - assert harness.config.max_turns == 3 - assert [metric.__name__ for metric in harness.metrics] == ["config_metric"] - assert "num_turns" in [signal["name"] for signal in harness.runtime.rollout_signals] - assert harness.rewards == [config_reward] - assert harness.advantages == [config_advantage] - assert harness.setups == [config_setup] - assert harness.cleanups == [config_cleanup] - assert harness.teardowns == [config_harness_teardown] - assert harness.user is not None - assert len(harness.toolsets) == 2 - assert harness.toolsets[0].hide == ("config_tool",) - assert harness.toolsets[1] is direct_toolset - - -def test_harness_owns_default_render_completion_update() -> None: - harness = make_harness(program={"fn": ref("config_program")}) - - assert any( - getattr(handler, "__self__", None) is harness - and getattr(handler, "__name__", "") == "render_completion" - for handler in harness.runtime.rollout_update - ) - - -def test_harness_owns_default_num_turns_metric() -> None: - harness = make_harness(program={"fn": ref("config_program")}) - - assert any( - signal["name"] == "num_turns" for signal in harness.runtime.rollout_signals - ) - - -@pytest.mark.asyncio -async def test_update_config_runs_before_rollout_scoring() -> None: - harness = make_harness( - program={"fn": ref("config_program")}, - config={ - "updates": [{"fn": ref("config_update"), "priority": 5}], - "rewards": [{"fn": ref("updated_reward"), "weight": 0.75}], - }, - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - - state = await harness.run(task) - - assert state["updated"] is True - assert state["reward"] == 0.75 - assert getattr(harness.updates[0], "__name__") == "config_update" - - -@pytest.mark.asyncio -async def test_scoring_config_entries_feed_runtime_as_dicts() -> None: - taskset = make_taskset( - rewards=[ref("config_reward")], - scoring={"config_reward": vf.SignalConfig(weight=0.5)}, - ) - harness = make_harness(program={"fn": ref("config_program")}) - Env(taskset=taskset, harness=harness) - - task = next(iter(taskset)) - state = await harness.run(task) - - assert taskset.config.scoring == {"config_reward": {"weight": 0.5}} - assert state["reward"] == 0.5 - - -@pytest.mark.asyncio -async def test_harness_scoring_config_entries_feed_runtime_as_dicts() -> None: - harness = make_harness( - program={"fn": ref("config_program")}, - config={ - "rewards": [ref("config_reward")], - "scoring": {"config_reward": {"weight": 0.5}}, - }, - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - - state = await harness.run(task) - - assert harness.config.scoring == {"config_reward": {"weight": 0.5}} - assert state["reward"] == 0.5 - - -@pytest.mark.asyncio -async def test_setup_config_runs_before_program() -> None: - harness = make_harness( - config={ - "program": {"fn": ref("setup_aware_program")}, - "setups": [{"fn": ref("config_setup"), "priority": 20}], - }, - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - - state = await harness.run(task) - - assert state["program_saw_setup"] is True - assert state["setup_order"] == ["config_setup"] - assert getattr(harness.setups[0], "__name__") == "config_setup" - - -@pytest.mark.asyncio -async def test_taskset_setup_runs_before_program() -> None: - taskset = make_taskset(setups=[ref("config_setup")]) - harness = make_harness(program={"fn": ref("setup_aware_program")}) - Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - - state = await harness.run(task) - - assert state["program_saw_setup"] is True - assert state["setup_order"] == ["config_setup"] - - -@pytest.mark.asyncio -async def test_configured_owner_teardowns_run() -> None: - TEARDOWN_EVENTS.clear() - taskset = make_taskset( - teardowns=[ref("config_taskset_teardown")], - ) - harness = make_harness(teardowns=[ref("config_harness_teardown")]) - Env(taskset=taskset, harness=harness) - - await harness.teardown() - - assert set(TEARDOWN_EVENTS) == {"taskset", "harness"} - - -@pytest.mark.asyncio -async def test_group_update_config_runs_before_group_scoring() -> None: - harness = make_harness( - config={ - "updates": [{"fn": ref("config_group_update"), "stage": "group"}], - "rewards": [{"fn": ref("group_updated_reward"), "stage": "group"}], - }, - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - await harness.score_group([task], [state]) - - assert state["group_updated"] is True - assert state["reward"] == 1.0 - - -def test_lifecycle_fields_are_framework_managed() -> None: - assert vf.State is State - - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - assert state.uses_v1_contract is True - error = vf.Error("boom") - - for key, value in { - "is_completed": True, - "stop_condition": "done", - "is_truncated": True, - "error": error, - }.items(): - assert State({key: value})[key] == value - with pytest.raises(RuntimeError, match="framework-managed"): - state[key] = value - with pytest.raises(RuntimeError, match="framework-managed"): - state.update({key: value}) - with pytest.raises(RuntimeError, match="framework-managed"): - state.setdefault(key, value) - with pytest.raises(RuntimeError, match="framework-managed"): - state.pop(key) - state["user_field"] = "ok" - assert state.popitem() == ("user_field", "ok") - - protected_only = State()._enable_v1_contract() - protected_only._set_completed(False) - protected_only._set_stop_condition(None, overwrite=True) - protected_only._set_truncated(False, overwrite=True) - protected_only._set_error(None) - with pytest.raises(RuntimeError, match="framework-managed"): - protected_only.popitem() - with pytest.raises(RuntimeError, match="framework-managed"): - state.clear() - - state._set_completed(True) - state._set_stop_condition("done") - state._set_truncated(True) - with pytest.raises(TypeError, match="vf.Error"): - state._set_error(cast(Any, {"message": "boom"})) - state._set_error(error) - - assert state["is_completed"] is True - assert state["stop_condition"] == "done" - assert state["is_truncated"] is True - assert state["error"] is error - - -def test_toolsets_config_accepts_addressable_map_and_fn_tables() -> None: - taskset = make_taskset( - config={ - "toolsets": { - "direct": {"tools": [ref("direct_tool")]}, - "configured": { - "fn": ref("config_toolset"), - "prefix": "configured", - }, - } - }, - ) - - assert set(taskset.named_toolsets) == {"direct", "configured"} - assert taskset.toolsets[0].tools == (direct_tool,) - prefix = taskset.toolsets[1].bindings["config_tool.prefix"] - assert callable(prefix) - assert prefix() == "configured" - - -def test_taskset_load_toolsets_adds_class_owned_toolsets() -> None: - class ToolsetTaskset(Taskset): - def load_toolsets(self, config: TasksetConfig) -> vf.Toolsets: - _ = config - return {"direct": Toolset(tools=[direct_tool])} - - taskset = ToolsetTaskset() - - assert set(taskset.named_toolsets) == {"direct"} - assert taskset.named_toolsets["direct"].tools == (direct_tool,) - - -def test_taskset_config_toolsets_collects_class_and_config_toolsets() -> None: - class ToolsetTaskset(Taskset): - def load_toolsets(self, config: TasksetConfig) -> vf.Toolsets: - _ = config - return {"direct": Toolset(tools=[direct_tool])} - - taskset = ToolsetTaskset( - config={ - "toolsets": {"configured": {"tools": [ref("config_tool")]}}, - } - ) - - assert set(taskset.named_toolsets) == {"direct", "configured"} - assert taskset.named_toolsets["direct"].tools == (direct_tool,) - assert taskset.named_toolsets["configured"].tools == (config_tool,) - - -def test_taskset_config_rejects_none_toolsets() -> None: - class ToolsetTaskset(Taskset): - def load_toolsets(self, config: TasksetConfig) -> vf.Toolsets: - _ = config - return {"direct": Toolset(tools=[direct_tool])} - - with pytest.raises(ValidationError): - ToolsetTaskset(config={"toolsets": None}) - - -def test_harness_config_rejects_none_toolsets() -> None: - class ToolsetHarness(Harness): - def load_toolsets(self, config: HarnessConfig) -> vf.Toolsets: - _ = config - return {"direct": Toolset(tools=[direct_tool])} - - with pytest.raises(ValidationError): - ToolsetHarness(config={"toolsets": None}) - - -def test_taskset_duplicate_toolset_names_raise_between_class_and_config() -> None: - class ToolsetTaskset(Taskset): - def load_toolsets(self, config: TasksetConfig) -> vf.Toolsets: - _ = config - return {"direct": Toolset(tools=[direct_tool])} - - with pytest.raises(ValueError, match="Toolsets are defined twice"): - ToolsetTaskset( - config={ - "toolsets": {"direct": {"tools": [ref("config_tool")]}}, - } - ) - - -@pytest.mark.asyncio -async def test_task_toolsets_show_hide_selects_named_defaults() -> None: - harness = make_harness( - toolsets={ - "direct": Toolset(tools=[direct_tool]), - "hidden": Toolset(tools=[hidden_tool]), - } - ) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "toolsets": {"show": ["direct"]}, - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert state["tools"] == ["direct_tool"] - assert list(harness.runtime.tool_calls(task, state)) == ["direct_tool"] - - -@pytest.mark.asyncio -async def test_state_can_add_rollout_local_tools() -> None: - async def provision_tool(state: State) -> None: - state.add_tool("local", config_tool) - - harness = make_harness( - toolsets={ - "local": Toolset( - scope="rollout", - bindings={"config_tool.prefix": "task.answer"}, - ) - } - ) - harness.add_setup(provision_tool) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "answer": "ok", - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - await harness.runtime.setup_rollout(task, state) - - assert state["tools"] == ["config_tool"] - assert await state.get_tools()["config_tool"](query="q") == "ok:q" - - -@pytest.mark.asyncio -async def test_state_add_tool_rejects_toolsets() -> None: - harness = make_harness(toolsets={"local": Toolset(scope="rollout")}) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - with pytest.raises(TypeError, match="tool, not a Toolset"): - state.add_tool("local", Toolset()) - - -@pytest.mark.asyncio -async def test_state_can_add_dynamic_schema_backed_tools() -> None: - async def provision_tool(task: Task, state: State) -> None: - state.add_tool("dynamic", dynamic_tool(task)) - - harness = make_harness( - toolsets={"dynamic": Toolset(scope="rollout", handler=dynamic_tool_handler)} - ) - harness.add_setup(provision_tool) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "dynamic_tool": { - "name": "lookup_city", - "description": "Look up one city.", - "parameters": { - "type": "object", - "properties": {"city": {"type": "string"}}, - "required": ["city"], - }, - }, - "tools": {"dynamic": {"show": ["lookup_city"]}}, - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - await harness.runtime.setup_rollout(task, state) - - tool_defs = harness.runtime.tool_defs(state) - assert tool_defs is not None - assert state["tools"] == ["lookup_city"] - assert tool_defs[0].name == "lookup_city" - assert tool_defs[0].parameters["properties"] == {"city": {"type": "string"}} - assert await state.get_tools()["lookup_city"](city="Paris") == "recorded" - assert state["dynamic_tool_calls"] == [{"lookup_city": {"city": "Paris"}}] - - -@pytest.mark.asyncio -async def test_tool_definition_provider_hides_bound_args() -> None: - class ProviderTool: - name = "provided_tool" - tool_def = vf.Tool( - name="provided_tool", - description="Provided schema tool.", - parameters={ - "type": "object", - "properties": { - "value": {"type": "string"}, - "prefix": {"type": "string"}, - }, - "required": ["value", "prefix"], - }, - ) - - async def __call__(self, value: str, prefix: str) -> str: - return f"{prefix}:{value}" - - harness = make_harness( - toolsets={ - "provided": Toolset( - tools=[ProviderTool()], - bindings={"provided_tool.prefix": "task.answer"}, - ) - } - ) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "answer": "ok", - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - tool_defs = harness.runtime.tool_defs(state) - assert tool_defs is not None - assert tool_defs[0].parameters["properties"] == {"value": {"type": "string"}} - assert tool_defs[0].parameters["required"] == ["value"] - assert await state.get_tools()["provided_tool"](value="done") == "ok:done" - - -@pytest.mark.asyncio -async def test_tool_bindings_inject_owner_private_objects() -> None: - harness = make_harness( - toolsets=[ - Toolset( - tools=[object_tool], - objects=vf.ObjectsConfig.model_validate( - {"box": ref("load_object_box")} - ), - bindings={"object_tool.box": "objects.box"}, - write=True, - ) - ] - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert await state.get_tools()["object_tool"](value="alpha") == "alpha" - - -@pytest.mark.asyncio -async def test_toolset_object_factory_accepts_bound_arguments() -> None: - harness = make_harness( - toolsets=[ - Toolset( - tools=[object_prefix_tool], - objects=vf.ObjectsConfig.model_validate( - {"box": ref("load_prefixed_object_box")} - ), - bindings={ - "box.prefix": "task.prefix", - "object_prefix_tool.box": "objects.box", - }, - write=True, - ) - ] - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "prefix": "bound"} - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert await state.get_tools()["object_prefix_tool"](value="alpha") == "bound:alpha" - - -@pytest.mark.asyncio -async def test_toolset_objects_require_active_owner() -> None: - active_toolset = Toolset( - tools=[direct_tool], - objects=vf.ObjectsConfig.model_validate({"box": ref("load_object_box")}), - ) - detached_toolset = Toolset( - objects=vf.ObjectsConfig.model_validate({"box": ref("load_object_box")}) - ) - harness = make_harness(toolsets={"active": active_toolset}) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - harness.runtime.prepare_state(task, state) - - with pytest.raises(RuntimeError, match="not active"): - await harness.runtime.resolve_toolset_object( - detached_toolset, "box", task, state - ) - - -@pytest.mark.asyncio -async def test_runtime_teardown_closes_scoped_toolset_objects() -> None: - closed_objects.clear() - toolset = Toolset( - scope="rollout", - objects=vf.ObjectsConfig.model_validate( - {"box": ref("load_rollout_closable_object")} - ), - ) - harness = make_harness(toolsets={"owned": toolset}) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - harness.runtime.prepare_state(task, state) - - await harness.runtime.resolve_toolset_object(toolset, "box", task, state) - await harness.teardown() - - assert closed_objects == ["rollout"] - - -def test_binding_strings_must_be_framework_paths() -> None: - with pytest.raises(ValueError, match="Binding string sources"): - Toolset(tools=[config_tool], bindings={"config_tool.prefix": "literal"}) - - -def test_binding_sources_reject_direct_objects() -> None: - with pytest.raises(TypeError, match="framework path or callable"): - Toolset(tools=[config_tool], bindings={"config_tool.prefix": object()}) - - -def test_toolset_binding_keys_must_target_callable_args() -> None: - with pytest.raises(ValueError, match="callable.arg"): - Toolset(tools=[config_tool], bindings={"prefix": "task.answer"}) - - -@pytest.mark.asyncio -async def test_rollout_handlers_receive_bound_hidden_args() -> None: - harness = make_harness( - toolsets=[ - Toolset( - updates=[update_from_binding], - bindings={"update_from_binding.expected": "task.answer"}, - ) - ] - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - await harness.runtime.update_rollout(task, state) - - assert state["expected"] == "ok" - - -@pytest.mark.asyncio -async def test_harness_handlers_receive_bound_hidden_args() -> None: - harness = make_harness( - updates=[ref("update_from_binding")], - bindings={"update_from_binding.expected": "task.answer"}, - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - await harness.runtime.update_rollout(task, state) - - assert state["expected"] == "ok" - - -@pytest.mark.asyncio -async def test_taskset_handlers_receive_bound_hidden_args() -> None: - taskset = make_taskset( - updates=[ref("update_from_binding")], - bindings={"update_from_binding.expected": "task.answer"}, - ) - harness = make_harness() - env = Env(taskset=taskset, harness=harness) - harness = env.harness - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - await harness.runtime.update_rollout(task, state) - - assert state["expected"] == "ok" - - -@pytest.mark.asyncio -async def test_group_handlers_receive_bound_hidden_args() -> None: - harness = make_harness( - updates=[ref("group_update_from_binding")], - bindings={"group_update_from_binding.expected": "tasks.0.answer"}, - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - state = State.for_task(task) - - await harness.runtime.update_group([task], [state]) - - assert state["group_expected"] == "ok" - - -@pytest.mark.asyncio -async def test_signals_receive_bound_hidden_args() -> None: - harness = make_harness( - rewards=[ref("reward_from_binding")], - bindings={"reward_from_binding.expected": "task.answer"}, - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - await harness.runtime.score_rollout(task, state) - - assert state["reward"] == 1.0 - assert state["metrics"]["reward_from_binding"] == 1.0 - - -@pytest.mark.asyncio -async def test_group_signals_receive_bound_hidden_args() -> None: - harness = make_harness( - rewards=[ref("group_reward_from_binding")], - bindings={"group_reward_from_binding.expected": "states.0.answer"}, - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - state["answer"] = "ok" - - await harness.runtime.score_group([task], [state]) - - assert state["reward"] == 1.0 - assert state["metrics"]["group_reward_from_binding"] == 1.0 - - -@pytest.mark.asyncio -async def test_object_bindings_are_private_to_callable_tools() -> None: - harness = make_harness( - toolsets=[ - Toolset( - updates=[update_from_binding], - objects=vf.ObjectsConfig.model_validate( - {"box": ref("load_object_box")} - ), - bindings={"update_from_binding.expected": "objects.box"}, - ) - ] - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - with pytest.raises(ValueError, match="objects"): - await harness.setup_state(task, State.for_task(task)) - - -@pytest.mark.asyncio -async def test_bindings_must_match_declared_callable_args() -> None: - harness = make_harness( - toolsets=[ - Toolset( - tools=[object_tool], - objects=vf.ObjectsConfig.model_validate( - {"box": ref("load_object_box")} - ), - bindings={"object_tool.missing": "objects.box"}, - ) - ] - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - with pytest.raises(TypeError, match="missing"): - await harness.setup_state(task, State.for_task(task)) - - -@pytest.mark.asyncio -async def test_tool_bindings_do_not_leak_to_same_named_handlers() -> None: - harness = make_harness( - updates=[ref("colliding_update")], - toolsets=[ - Toolset( - tools=[colliding_tool], - objects=vf.ObjectsConfig.model_validate( - {"token": ref("load_object_box")} - ), - bindings={"colliding_tool.token": "objects.token"}, - ) - ], - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert await state.get_tools()["colliding_tool"](value="x") == "{'values': []}:x" - with pytest.raises(TypeError, match="token"): - await harness.runtime.update_rollout(task, state) - - -def test_harness_constructor_accepts_model_shortcuts_only() -> None: - harness = Harness(model="configured-model", sampling_args={"temperature": 0.2}) - assert harness.config.model.name == "configured-model" - assert harness.config.model.sampling_args == {"temperature": 0.2} - - with pytest.raises(TypeError): - Harness(max_turns=9) - with pytest.raises(TypeError): - Taskset(taskset_id="configured") - with pytest.raises(TypeError): - OpenCode(max_turns=9) - with pytest.raises(TypeError): - Harness(config=HarnessConfig(), model="configured-model") - - -def test_task_prompt_rejects_system_messages() -> None: - with pytest.raises(ValueError, match="Use system_prompt instead"): - Task({"prompt": [{"role": "system", "content": "sys"}]}).freeze() - - -def test_task_system_prompt_is_normalized() -> None: - task = Task( - { - "system_prompt": "sys", - "prompt": [{"role": "user", "content": "hi"}], - } - ).freeze() - - assert task["system_prompt"] == [{"role": "system", "content": "sys"}] - assert task["prompt"] == [{"role": "user", "content": "hi"}] - - -@pytest.mark.asyncio -async def test_harness_resolves_taskset_system_prompt() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness(program={"fn": ref("config_program")}) - Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [{"role": "system", "content": "taskset sys"}] - assert state["prompt"] == [{"role": "user", "content": "Say ok."}] - - -def test_taskset_load_system_prompt_method_owns_prompt_loading() -> None: - class PromptTaskset(Taskset): - def load_system_prompt(self, config: TasksetConfig) -> vf.SystemPrompt: - _ = config - return load_system_prompt() - - taskset = PromptTaskset(config=TasksetConfig()) - - assert taskset.system_prompt == [ - {"role": "system", "content": "loaded system prompt"} - ] - - -def test_system_prompt_bare_string_is_literal() -> None: - class PromptTasksetConfig(TasksetConfig): - system_prompt: str = "load_system_prompt" - - taskset = Taskset(config=PromptTasksetConfig()) - - assert taskset.system_prompt == [ - {"role": "system", "content": "load_system_prompt"} - ] - - -def test_system_prompt_accepts_path(tmp_path) -> None: - prompt_path = tmp_path / "system_prompt.txt" - prompt_path.write_text("path system prompt", encoding="utf-8") - - taskset = make_taskset(system_prompt=vf.SystemPromptConfig(path=str(prompt_path))) - - assert taskset.system_prompt == [ - {"role": "system", "content": "path system prompt"} - ] - - -def test_system_prompt_direct_string_can_contain_colon() -> None: - taskset = make_taskset(system_prompt="Answer:yes") - - assert taskset.system_prompt == [{"role": "system", "content": "Answer:yes"}] - - -@pytest.mark.asyncio -async def test_harness_concats_multiple_system_prompt_sources_by_default() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness( - program={"fn": ref("config_program")}, system_prompt="harness sys" - ) - Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [ - {"role": "system", "content": "harness sys"}, - {"role": "system", "content": "taskset sys"}, - ] - - -@pytest.mark.asyncio -async def test_task_system_prompt_overrides_taskset_side_at_runtime() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness(program={"fn": ref("config_program")}) - Env(taskset=taskset, harness=harness) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "system_prompt": "task sys", - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [{"role": "system", "content": "task sys"}] - - -@pytest.mark.asyncio -async def test_task_override_is_resolved_before_harness_concat() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness( - program={"fn": ref("config_program")}, system_prompt="harness sys" - ) - Env(taskset=taskset, harness=harness) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "system_prompt": "task sys", - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [ - {"role": "system", "content": "harness sys"}, - {"role": "system", "content": "task sys"}, - ] - - -@pytest.mark.asyncio -async def test_system_prompt_strategy_can_concat_taskset_side_first() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness( - program={"fn": ref("config_program")}, - system_prompt="harness sys", - system_prompt_strategy="TH", - ) - Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [ - {"role": "system", "content": "taskset sys"}, - {"role": "system", "content": "harness sys"}, - ] - - -@pytest.mark.asyncio -async def test_harness_can_reject_multiple_system_prompt_sides() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness( - program={"fn": ref("config_program")}, - system_prompt="harness sys", - system_prompt_strategy="REJECT", - ) - Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - - with pytest.raises(ValueError, match="Multiple system_prompt sides"): - await harness.setup_state(task, State.for_task(task)) - - -@pytest.mark.asyncio -async def test_system_prompt_side_selection_uses_resolved_taskset_side() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness( - program={"fn": ref("config_program")}, - system_prompt="harness sys", - system_prompt_strategy="T_OR_H", - ) - Env(taskset=taskset, harness=harness) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "system_prompt": "task sys", - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [{"role": "system", "content": "task sys"}] - - -@pytest.mark.asyncio -async def test_system_prompt_side_selection_can_prefer_harness() -> None: - taskset = make_taskset(system_prompt="taskset sys") - harness = make_harness( - program={"fn": ref("config_program")}, - system_prompt="harness sys", - system_prompt_strategy="H_OR_T", - ) - Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - state = await harness.setup_state(task, State.for_task(task)) - - assert state["system_prompt"] == [{"role": "system", "content": "harness sys"}] - - -@pytest.mark.asyncio -async def test_system_prompt_strategy_can_select_exact_sides() -> None: - taskset = make_taskset(system_prompt="taskset sys") - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - harness_t = make_harness( - program={"fn": ref("config_program")}, - system_prompt="harness sys", - system_prompt_strategy="T", - ) - Env(taskset=taskset, harness=harness_t) - state_t = await harness_t.setup_state(task, State.for_task(task)) - - harness_h = make_harness( - program={"fn": ref("config_program")}, - system_prompt="harness sys", - system_prompt_strategy="H", - ) - Env(taskset=taskset, harness=harness_h) - state_h = await harness_h.setup_state(task, State.for_task(task)) - - assert state_t["system_prompt"] == [{"role": "system", "content": "taskset sys"}] - assert state_h["system_prompt"] == [{"role": "system", "content": "harness sys"}] - - -@pytest.mark.asyncio -async def test_task_max_turns_overrides_harness_default() -> None: - harness = make_harness(max_turns=9) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "max_turns": 3, - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert state.get_max_turns(harness.config.max_turns) == 3 - - -@pytest.mark.asyncio -async def test_explicit_state_runtime_max_turns_overrides_task_controls() -> None: - harness = make_harness(max_turns=9) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "max_turns": 3, - } - ).freeze() - state = State.for_task(task) - state["runtime"] = {"max_turns": 2} - state = await harness.setup_state(task, state) - - assert state.get_max_turns(harness.config.max_turns) == 2 - - -def test_task_runtime_is_not_public_task_schema() -> None: - with pytest.raises(TypeError, match="task.runtime"): - Task({"runtime": {"unknown": True}}).freeze() - - -def test_task_runtime_rejects_legacy_max_turns() -> None: - with pytest.raises(TypeError, match="task.runtime"): - Task({"runtime": {"max_turns": "3"}}).freeze() - - -def test_task_rejects_non_integer_max_turns() -> None: - with pytest.raises(TypeError): - Task({"max_turns": "3"}).freeze() - - -def test_task_sandbox_must_be_mapping() -> None: - with pytest.raises(TypeError, match="task.sandbox"): - Task({"prompt": [], "sandbox": "rollout"}).freeze() - - -def test_option_only_program_requires_sandbox_placement() -> None: - with pytest.raises(ValueError, match="require sandbox placement"): - make_harness(program={"sandbox": False}) - - make_harness(program={"sandbox": True}, sandbox={"image": "python:3.11-slim"}) - - -def test_harness_config_sandbox_values_live_only_in_config() -> None: - harness = make_harness( - config={"sandbox": {"image": "configured", "memory_gb": 8, "scope": "group"}}, - ) - - assert harness.sandbox is not None - assert harness.sandbox.image == "configured" - assert harness.sandbox.memory_gb == 8 - assert harness.sandbox.scope == "group" - - -@pytest.mark.asyncio -async def test_user_config_supports_scope_bindings_and_objects() -> None: - class ConfigUserHarnessConfig(HarnessConfig): - user: ConfigUserWithBindingsConfig = ConfigUserWithBindingsConfig( - scope="group", - bindings={"token": "objects.token"}, - objects={"token": ref("token_factory")}, - ) - - class ConfigUserHarness(Harness[ConfigUserHarnessConfig]): - pass - - harness = ConfigUserHarness(config=ConfigUserHarnessConfig()) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - messages = await harness.runtime.user_messages( - task, state, transcript=[{"role": "assistant", "content": "hello"}] - ) - - assert harness.user is not None - assert harness.user.scope == "group" - assert state["token_seen"] == "secret-token" - assert state["messages_len"] == 1 - assert messages == [{"role": "user", "content": "secret-token"}] - - -@pytest.mark.asyncio -async def test_user_objects_require_active_owner() -> None: - class ConfigUserHarnessConfig(HarnessConfig): - user: ConfigUserWithBindingsConfig = ConfigUserWithBindingsConfig( - objects={"token": ref("token_factory")} - ) - - class ConfigUserHarness(Harness[ConfigUserHarnessConfig]): - pass - - harness = ConfigUserHarness(config=ConfigUserHarnessConfig()) - detached_user = ConfigUserWithBindings( - config=ConfigUserWithBindingsConfig(objects={"token": ref("token_factory")}) - ) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - with pytest.raises(RuntimeError, match="not attached"): - await harness.runtime.resolve_user_object(detached_user, "token", task, state) - - -@pytest.mark.asyncio -async def test_user_binding_can_use_taskset_runtime_object() -> None: - class OwnerObjectUserTasksetConfig(TasksetConfig): - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"token": ref("token_factory")} - ) - user: ConfigUserWithBindingsConfig = ConfigUserWithBindingsConfig( - bindings=vf.BindingsConfig.model_validate( - {"token": "taskset.objects.token"} - ) - ) - - class OwnerObjectUserTaskset(Taskset[OwnerObjectUserTasksetConfig]): - pass - - taskset = OwnerObjectUserTaskset(config=OwnerObjectUserTasksetConfig()) - env = Env(taskset=taskset, harness=Harness()) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await env.harness.setup_state(task, State.for_task(task)) - - messages = await env.harness.runtime.user_messages(task, state) - - assert state["token_seen"] == "secret-token" - assert messages == [{"role": "user", "content": "secret-token"}] - - -def test_taskset_config_default_user_is_active() -> None: - class ConfigDefaultUserTasksetConfig(TasksetConfig): - user: ConfigUserConfig = ConfigUserConfig() - - class ConfigDefaultUserTaskset(Taskset[ConfigDefaultUserTasksetConfig]): - pass - - taskset = ConfigDefaultUserTaskset(config=ConfigDefaultUserTasksetConfig()) - - assert taskset.user is not None - - -def test_harness_config_default_user_is_active() -> None: - class ConfigDefaultUserHarnessConfig(HarnessConfig): - user: ConfigUserConfig = ConfigUserConfig() - - class ConfigDefaultUserHarness(Harness[ConfigDefaultUserHarnessConfig]): - pass - - harness = ConfigDefaultUserHarness(config=ConfigDefaultUserHarnessConfig()) - - assert harness.user is not None - - -@pytest.mark.asyncio -async def test_user_config_receives_default_messages_binding() -> None: - class DirectUserHarnessConfig(HarnessConfig): - user: DirectUserWithMessagesConfig = DirectUserWithMessagesConfig() - - class DirectUserHarness(Harness[DirectUserHarnessConfig]): - pass - - harness = DirectUserHarness(config=DirectUserHarnessConfig()) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - messages = await harness.runtime.user_messages( - task, state, transcript=[{"role": "assistant", "content": "hello"}] - ) - - assert state["direct_messages_len"] == 1 - assert messages == [{"role": "user", "content": "continue"}] - - -@pytest.mark.asyncio -async def test_user_config_can_request_scoped_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - sandbox = object() - - class SandboxUserHarnessConfig(HarnessConfig): - user: SandboxUserConfig = SandboxUserConfig( - sandbox={"image": "python:3.11-slim", "scope": "group"} - ) - - class SandboxUserHarness(Harness[SandboxUserHarnessConfig]): - pass - - harness = SandboxUserHarness(config=SandboxUserHarnessConfig()) - task = Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = State.for_task(task) - - async def resolve_user_sandbox(*args: Any, **kwargs: Any) -> object: - _ = args, kwargs - return sandbox - - monkeypatch.setattr(harness.runtime, "resolve_user_sandbox", resolve_user_sandbox) - - messages = await harness.runtime.user_messages(task, state) - - assert harness.user is not None - assert harness.user.sandbox is not None - assert harness.user.sandbox.image == "python:3.11-slim" - assert harness.user.sandbox.scope == "group" - assert state["sandbox_seen"] is sandbox - assert messages == [{"role": "user", "content": "sandbox ok"}] - - -@pytest.mark.asyncio -async def test_configured_program_scores_and_cleans_rollout() -> None: - taskset = make_taskset() - harness = make_harness( - config={ - "program": {"fn": ref("config_program")}, - "rewards": [ref("config_reward")], - "cleanups": [ref("config_cleanup")], - } - ) - task = next(iter(taskset)) - state = await harness.run(task) - - assert state["program"] == "ran" - assert state["answer"] == "ok" - assert state["reward"] == 0.25 - assert state["cleaned"] is True - assert state["is_completed"] is True - - -@pytest.mark.asyncio -async def test_harness_run_releases_group_scope_when_no_group_boundary() -> None: - harness = make_harness( - program={"fn": ref("config_program")}, cleanups=[ref("config_group_cleanup")] - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - - state = await harness.run(task) - - assert state["group_cleaned"] is True - - -@pytest.mark.asyncio -async def test_harness_run_defers_group_cleanup_when_group_boundary_exists() -> None: - harness = make_harness( - program={"fn": ref("config_program")}, cleanups=[ref("config_group_cleanup")] - ) - task = Task( - {"prompt": [{"role": "user", "content": "hi"}], "answer": "ok"} - ).freeze() - state = State.for_task(task) - state["runtime"]["group_key"] = "group" - - state = await harness.run(task, state) - - assert "group_cleaned" not in state - await harness.cleanup_group([task], [state]) - assert state["group_cleaned"] is True - - -def test_taskset_and_harness_preserve_explicit_config_subtypes() -> None: - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalTaskset(Taskset[LocalTasksetConfig]): - config: LocalTasksetConfig - - pass - - class LocalHarnessConfig(HarnessConfig): - mode: str = "default" - - class LocalHarness(Harness[LocalHarnessConfig]): - config: LocalHarnessConfig - - taskset = LocalTaskset(config=LocalTasksetConfig(split="test")) - harness = LocalHarness(config=LocalHarnessConfig(mode="custom")) - env = Env(taskset=taskset, harness=harness) - - assert env.taskset is taskset - assert env.harness is harness - assert isinstance(taskset.config, LocalTasksetConfig) - assert taskset.config.split == "test" - assert isinstance(harness.config, LocalHarnessConfig) - assert harness.config.mode == "custom" - - -def test_env_constructor_requires_required_child_configs() -> None: - class RequiredTasksetConfig(TasksetConfig): - dataset: str - - class RequiredTaskset(Taskset): - pass - - class RequiredHarnessConfig(HarnessConfig): - endpoint: str - - class RequiredHarness(Harness): - pass - - with pytest.raises(ValidationError, match="dataset"): - RequiredTasksetConfig.model_validate({}) - - prebuilt_env = Env( - taskset=RequiredTaskset(config=RequiredTasksetConfig(dataset="prebuilt-train")), - harness=RequiredHarness(config=RequiredHarnessConfig(endpoint="prebuilt")), - ) - assert prebuilt_env.taskset.config.dataset == "prebuilt-train" - assert prebuilt_env.harness.config.endpoint == "prebuilt" - - with pytest.raises(TypeError, match="Env taskset must be a Taskset"): - Env( - taskset=RequiredTasksetConfig(dataset="train"), - harness=RequiredHarnessConfig(endpoint="local"), - ) - - -def test_env_requires_taskset() -> None: - with pytest.raises(TypeError, match="requires a taskset"): - Env() - - -def test_env_config_tracks_prebuilt_children() -> None: - taskset = Taskset(config=TasksetConfig(taskset_id="actual")) - harness = Harness(config=HarnessConfig(max_turns=3)) - env = Env(taskset=taskset, harness=harness) - - assert env.config.taskset is taskset.config - assert env.config.harness is harness.config - assert env.config.taskset.taskset_id == "actual" - assert env.config.harness.max_turns == 3 - - -def test_taskset_and_harness_configs_accept_id_shorthand() -> None: - class CustomTasksetConfig(TasksetConfig): - taskset_id: str | None = "default-taskset" - - class CustomHarnessConfig(HarnessConfig): - harness_id: str | None = "default-harness" - - taskset_config = CustomTasksetConfig.model_validate({"id": "taskset-short"}) - harness_config = CustomHarnessConfig.model_validate({"id": "harness-short"}) - - assert taskset_config.taskset_id == "taskset-short" - assert harness_config.harness_id == "harness-short" - assert explicit_config_data(taskset_config) == {"taskset_id": "taskset-short"} - assert explicit_config_data(harness_config) == {"harness_id": "harness-short"} - - taskset = Taskset(config={"id": "taskset-short"}) - harness = Harness(config={"id": "harness-short"}) - - assert taskset.taskset_id == "taskset-short" - assert harness.harness_id == "harness-short" - - -def test_env_rejects_taskset_builders() -> None: - def load_taskset() -> Taskset: - return Taskset(config=TasksetConfig()) - - with pytest.raises(TypeError, match="Env taskset must be a Taskset"): - Env(taskset=load_taskset) - - -def test_env_rejects_harness_builders() -> None: - taskset = make_taskset() - - def load_harness(config: HarnessConfig | None = None) -> Harness: - return Harness(config=config) - - with pytest.raises(TypeError, match="Env harness must be a Harness"): - Env(taskset=taskset, harness=load_harness) - - -def test_package_harness_requires_package_config_subtype() -> None: - from harnesses.opencode import OpenCode - from harnesses.opencode import OpenCodeConfig - - config = OpenCode( - config=OpenCodeConfig(model=vf.ModelConfig(name="configured-model")) - ).config - - assert config.model.name == "configured-model" - assert config.max_turns == OpenCodeConfig().max_turns - base_config = OpenCode( - config=HarnessConfig(model=vf.ModelConfig(name="configured-model")) - ).config - - assert isinstance(base_config, OpenCodeConfig) - assert base_config.model.name == "configured-model" - - -def test_taskset_config_defaults_are_used_until_config_overrides() -> None: - class LocalTasksetConfig(TasksetConfig): - dataset: str = "default" - rewards: list[str] = [ref("config_reward")] - - class LocalTaskset(Taskset[LocalTasksetConfig]): - config: LocalTasksetConfig - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if self.config.dataset == "other": - return load_other_tasks(split) - return load_tasks(split) - - taskset = LocalTaskset(config=LocalTasksetConfig()) - configured = LocalTaskset( - config=LocalTasksetConfig( - dataset="other", - rewards=[ref("updated_reward")], - ) - ) - disabled = LocalTaskset(config=LocalTasksetConfig(rewards=[])) - - assert taskset.get_dataset()[0]["answer"] == "ok" - assert taskset.rewards == [config_reward] - assert configured.get_dataset()[0]["answer"] == "other ok" - assert configured.rewards[0].__name__ == "updated_reward" - assert disabled.rewards == [] - - -def test_taskset_generic_sets_subclass_config_type() -> None: - class RegisteredTasksetConfig(TasksetConfig): - dataset_name: str = "registered" - dataset_split: str = "train" - system_prompt: str | None = "default prompt" - - class RegisteredTaskset(Taskset[RegisteredTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return [ - { - "prompt": [], - "answer": f"{self.config.dataset_name}:{self.config.dataset_split}", - } - ] - - def load_system_prompt( - self, config: RegisteredTasksetConfig - ) -> vf.SystemPrompt: - _ = config - return "registered prompt" - - taskset = RegisteredTaskset(config=RegisteredTasksetConfig(dataset_split="eval")) - - assert isinstance(taskset.config, RegisteredTasksetConfig) - assert taskset.get_dataset()[0]["answer"] == "registered:eval" - assert taskset.system_prompt == [{"role": "system", "content": "registered prompt"}] - - -def test_taskset_config_annotation_registers_config_type_at_runtime() -> None: - class AnnotatedTasksetConfig(TasksetConfig): - dataset_name: str = "annotated" - - class AnnotatedTaskset(Taskset): - config: AnnotatedTasksetConfig - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return [{"prompt": [], "answer": self.config.dataset_name}] - - taskset = AnnotatedTaskset(config=AnnotatedTasksetConfig()) - - assert isinstance(taskset.config, AnnotatedTasksetConfig) - assert taskset.get_dataset()[0]["answer"] == "annotated" - - -def test_taskset_subclasses_inherit_registered_config_type() -> None: - class BaseTasksetConfig(TasksetConfig): - pass - - class BaseTaskset(Taskset[BaseTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - class ChildTaskset(BaseTaskset): - pass - - taskset = ChildTaskset(config=BaseTasksetConfig()) - - assert isinstance(taskset.config, BaseTasksetConfig) - assert taskset.get_dataset()[0]["answer"] == "ok" - - -def test_taskset_class_loader_owns_split_loading() -> None: - class LoaderTasksetConfig(TasksetConfig): - system_prompt: vf.SystemPrompt = "class prompt" - - class LoaderTaskset(Taskset[LoaderTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - answer = "class eval" if split == "eval" else "class tasks" - return [{"prompt": [], "answer": answer}] - - def load_system_prompt(self, config: LoaderTasksetConfig) -> vf.SystemPrompt: - return config.system_prompt - - defaulted = LoaderTaskset(config=LoaderTasksetConfig()) - configured = LoaderTaskset( - config=LoaderTasksetConfig( - system_prompt=ref("load_system_prompt"), - ) - ) - disabled_prompt = LoaderTaskset(config=LoaderTasksetConfig(system_prompt=None)) - - assert defaulted.get_dataset()[0]["answer"] == "class tasks" - assert defaulted.get_eval_dataset()[0]["answer"] == "class eval" - assert defaulted.system_prompt == [{"role": "system", "content": "class prompt"}] - assert configured.get_dataset()[0]["answer"] == "class tasks" - assert configured.get_eval_dataset()[0]["answer"] == "class eval" - assert configured.system_prompt == [ - {"role": "system", "content": ref("load_system_prompt")} - ] - assert disabled_prompt.system_prompt == [] - - -def test_system_prompt_alias_accepts_config_data(tmp_path) -> None: - prompt_path = tmp_path / "system_prompt.txt" - prompt_path.write_text("alias path system prompt", encoding="utf-8") - - class PromptTasksetConfig(TasksetConfig): - system_prompt: vf.SystemPrompt = None - - config = PromptTasksetConfig.model_validate( - {"system_prompt": {"path": str(prompt_path)}} - ) - assert isinstance(config.system_prompt, vf.SystemPromptConfig) - - taskset = Taskset(config=config) - - assert taskset.system_prompt == [ - {"role": "system", "content": "alias path system prompt"} - ] - - -def test_taskset_load_tasks_can_return_empty_dataset() -> None: - class LocalTasksetConfig(TasksetConfig): - enabled: bool = True - - class LocalTaskset(Taskset[LocalTasksetConfig]): - config: LocalTasksetConfig - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if not self.config.enabled: - return [] - return load_tasks(split) - - taskset = LocalTaskset(config=LocalTasksetConfig()) - disabled = LocalTaskset(config=LocalTasksetConfig(enabled=False)) - - assert taskset.get_dataset()[0]["answer"] == "ok" - assert len(disabled.get_dataset()) == 0 - - -def test_config_schema_is_visible_from_primary_types() -> None: - taskset_schema = TasksetConfig.schema_text() - assert "- toolsets:" in taskset_schema - assert "- toolsets:" in HarnessConfig.schema_text() - assert "- tasks:" not in taskset_schema - assert "- program:" in HarnessConfig.schema_text() - assert "- image:" in vf.SandboxConfig.schema_text() - assert "- bindings:" in vf.ToolsetConfig.schema_text() - - -def test_config_annotation_only_nested_config_defaults_recursively() -> None: - class LeafConfig(Config): - value: int = 1 - - class ChildConfig(Config): - leaf: LeafConfig = LeafConfig() - - class ParentConfig(Config): - child: ChildConfig = ChildConfig() - - first = ParentConfig() - second = ParentConfig() - configured = ParentConfig.model_validate({"child": {"leaf": {"value": 3}}}) - - assert isinstance(first.child, ChildConfig) - assert isinstance(first.child.leaf, LeafConfig) - assert first.child.leaf.value == 1 - assert first.child is not second.child - assert first.child.leaf is not second.child.leaf - assert configured.child.leaf.value == 3 - assert "child: ChildConfig = ChildConfig" in ParentConfig.schema_text() - - -def test_env_config_normalizes_mapping_config_to_attributes() -> None: - config = EnvConfig.model_validate( - { - "taskset": {"taskset_id": "dict"}, - "harness": {"model": {"name": "configured-model"}}, - } - ) - - assert isinstance(config.taskset, TasksetConfig) - assert isinstance(config.harness, HarnessConfig) - assert config.taskset.taskset_id == "dict" - assert config.harness.model.name == "configured-model" - - -def test_env_config_defaults_taskset_and_harness_to_base_configs() -> None: - config = EnvConfig() - - assert isinstance(config.taskset, TasksetConfig) - assert isinstance(config.harness, HarnessConfig) - - -def test_env_config_rejects_unknown_top_level_sections() -> None: - with pytest.raises(ValueError): - EnvConfig.model_validate({"taskset": {}, "math": {"taskset": {}}}) - - -def test_env_config_requires_child_sections_to_be_configs() -> None: - with pytest.raises(ValueError): - EnvConfig.model_validate({"taskset": 1}) - with pytest.raises(ValueError): - EnvConfig.model_validate({"taskset": None}) - with pytest.raises(ValueError): - EnvConfig(harness=None) - - -def test_env_config_child_config_objects_must_match_domain() -> None: - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalHarnessConfig(HarnessConfig): - mode: str = "default" - - config = EnvConfig( - taskset=LocalTasksetConfig(split="test"), - harness=LocalHarnessConfig(mode="custom"), - ) - - assert isinstance(config.taskset, LocalTasksetConfig) - assert isinstance(config.harness, LocalHarnessConfig) - - class LocalConfig(Config): - split: str = "train" - - with pytest.raises(ValueError): - EnvConfig(taskset=LocalConfig()) - with pytest.raises(ValueError): - EnvConfig(harness=LocalConfig()) - - -def test_env_config_validates_nested_sections_into_annotated_child_types() -> None: - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalEnvConfig(EnvConfig): - taskset: LocalTasksetConfig = LocalTasksetConfig() - harness: HarnessConfig = HarnessConfig(max_turns=10) - - config = LocalEnvConfig.model_validate( - {"taskset": {"split": "nested"}, "harness": {"max_turns": 3}} - ) - default_config = LocalEnvConfig() - - assert isinstance(config.taskset, LocalTasksetConfig) - assert config.taskset.split == "nested" - assert isinstance(config.harness, HarnessConfig) - assert config.harness.max_turns == 3 - assert isinstance(default_config.taskset, LocalTasksetConfig) - assert default_config.taskset.split == "train" - assert default_config.harness.max_turns == 10 - - -def test_config_model_validate_keeps_serializable_nested_values() -> None: - config = HarnessConfig.model_validate( - { - "model": { - "sampling_args": { - "temperature": 0.7, - "extra_body": { - "top_p": None, - "min_p": 0.05, - }, - "stop": [None, "DONE"], - }, - } - } - ) - - assert config.model.sampling_args == { - "temperature": 0.7, - "extra_body": { - "top_p": None, - "min_p": 0.05, - }, - "stop": [None, "DONE"], - } - - -def test_config_rejects_live_python_objects() -> None: - with pytest.raises(ValueError): - HarnessConfig( - program=config_program, - ) - with pytest.raises(ValueError): - HarnessConfig( - program={"env": {"DYNAMIC_VALUE": {"fn": config_program}}}, - ) - with pytest.raises(TypeError): - TasksetConfig( - objects={"loader": load_tasks}, - ) - - -def test_config_json_round_trip_preserves_values() -> None: - config = HarnessConfig( - model=vf.ModelConfig( - sampling_args={ - "extra_body": { - "min_p": 0.05, - }, - "stop": [None, "DONE"], - } - ), - program=vf.ProgramConfig(fn=ref("config_program")), - ) - - dump = config.model_dump(mode="json", exclude_none=True) - - assert HarnessConfig.model_validate(dump) == config - assert dump["program"]["fn"] == ref("config_program") - assert dump["program"]["files"] == {} - assert dump["program"]["setup"] == [] - - -def test_env_config_subclasses_cannot_define_root_fields() -> None: - with pytest.raises(TypeError, match="unsupported root env config fields"): - - class LocalEnvConfig(EnvConfig): - split: str = "train" - - -def test_env_config_subclasses_must_use_domain_child_configs() -> None: - class LocalConfig(Config): - split: str = "train" - - with pytest.raises(TypeError, match="taskset must be typed"): - - class LocalEnvConfig(EnvConfig): - taskset: LocalConfig - - -def test_env_config_rejects_legacy_config_ref_merging() -> None: - with pytest.raises(ValueError): - EnvConfig.model_validate( - { - "harness": { - "config": ref("load_another_harness_config"), - "rewards": [{"fn": ref("updated_reward"), "weight": 0}], - } - } - ) - - -def test_harness_config_normalizes_program_mapping() -> None: - config = HarnessConfig( - program={ - "command": ["echo", "ok"], - "sandbox": {"packages": "numpy"}, - "channels": {"mcp": True}, - } - ) - - assert isinstance(config.program, vf.ProgramConfig) - assert config.program.command == ["echo", "ok"] - assert isinstance(config.program.sandbox, vf.SandboxConfig) - assert config.program.sandbox.packages == ["numpy"] - assert config.program.channels == {"mcp": True} - - -def test_harness_config_rejects_unknown_program_tool_interface() -> None: - with pytest.raises(ValueError, match="unknown channel"): - HarnessConfig(program={"command": ["echo"], "channels": {"ptc": True}}) - - -def test_load_environment_validates_typed_env_config_arg( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "typed_env_config" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - def load_environment(split: str = "train", *, config: EnvConfig) -> Env: - seen["split"] = split - seen["config"] = config - return Env( - taskset=make_taskset(config=config.taskset), - harness=make_harness(config=config.harness), - ) - - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "typed-env-config", - split="test", - config={ - "taskset": {"taskset_id": "typed"}, - "harness": {"model": {"name": "typed-model"}}, - }, - ) - - assert seen["split"] == "test" - assert isinstance(seen["config"], EnvConfig) - assert env.taskset.config.taskset_id == "typed" - assert env.harness.config.model.name == "typed-model" - assert env.env_args == { - "split": "test", - "config": { - "taskset": {"taskset_id": "typed"}, - "harness": {"model": {"name": "typed-model"}}, - }, - } - - -def test_load_environment_validates_env_config_subclass_sections( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "typed_env_config_subclass" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalHarnessConfig(HarnessConfig): - mode: str = "default" - - class LocalEnvConfig(EnvConfig): - taskset: LocalTasksetConfig - harness: LocalHarnessConfig - - class LocalTaskset(Taskset[LocalTasksetConfig]): - config: LocalTasksetConfig - - class LocalHarness(Harness): - pass - - def load_environment(config: LocalEnvConfig) -> Env: - seen["config"] = config - return Env( - taskset=LocalTaskset(config=config.taskset), - harness=LocalHarness(config=config.harness), - ) - - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "typed-env-config-subclass", - config={ - "taskset": {"taskset_id": "typed", "split": "test"}, - "harness": {"mode": "custom"}, - }, - ) - config = seen["config"] - - assert isinstance(config, LocalEnvConfig) - assert isinstance(config.taskset, LocalTasksetConfig) - assert isinstance(config.harness, LocalHarnessConfig) - assert env.taskset.config.taskset_id == "typed" - assert env.taskset.config.split == "test" - assert env.harness.config.mode == "custom" - - -def test_load_environment_uses_factory_annotations_for_child_config_types( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "factory_typed_child_config" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalHarnessConfig(HarnessConfig): - mode: str = "default" - - class LocalTaskset(Taskset): - config: LocalTasksetConfig - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - class LocalHarness(Harness): - config: LocalHarnessConfig - - pass - - def load_taskset(config: LocalTasksetConfig) -> LocalTaskset: - seen["taskset_config"] = config - return LocalTaskset(config=config) - - def load_harness(config: LocalHarnessConfig) -> LocalHarness: - seen["harness_config"] = config - return LocalHarness(config=config) - - def load_environment(config: EnvConfig) -> Env: - taskset_config = config.taskset - harness_config = config.harness - assert isinstance(taskset_config, LocalTasksetConfig) - assert isinstance(harness_config, LocalHarnessConfig) - return Env( - taskset=load_taskset(taskset_config), - harness=load_harness(harness_config), - ) - - module.load_taskset = load_taskset - module.load_harness = load_harness - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "factory-typed-child-config", - config={ - "taskset": {"taskset_id": "typed", "split": "test"}, - "harness": {"model": {"name": "typed-model"}, "mode": "custom"}, - }, - ) - taskset_config = seen["taskset_config"] - harness_config = seen["harness_config"] - - assert isinstance(taskset_config, LocalTasksetConfig) - assert isinstance(harness_config, LocalHarnessConfig) - assert env.taskset.config.taskset_id == "typed" - assert env.taskset.config.split == "test" - assert env.harness.config.model.name == "typed-model" - assert env.harness.config.mode == "custom" - - -def test_load_environment_keeps_environment_loader_authoritative( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "explicit_env_loader_with_components" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalTaskset(Taskset[LocalTasksetConfig]): - config: LocalTasksetConfig - - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - def load_taskset(config: LocalTasksetConfig) -> Taskset: - raise AssertionError("load_environment should decide when to load components") - - def load_environment(config: EnvConfig) -> Env: - seen["config"] = config - assert isinstance(config.taskset, LocalTasksetConfig) - return Env(taskset=LocalTaskset(config=config.taskset)) - - module.load_taskset = load_taskset - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "explicit-env-loader-with-components", - config={"taskset": {"split": "test"}}, - ) - - assert isinstance(seen["config"], EnvConfig) - assert env.taskset.get_dataset()[0]["answer"] == "ok" - assert env.taskset.config.split == "test" - - -def test_public_component_loaders_coerce_factory_config_annotations( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "component_loader_config" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - class LocalHarnessConfig(HarnessConfig): - mode: str = "default" - - class LocalHarness(Harness): - config: LocalHarnessConfig - - def load_taskset(config: LocalTasksetConfig) -> Taskset: - seen["taskset_config"] = config - return Taskset(config=config) - - def load_harness(config: LocalHarnessConfig) -> LocalHarness: - seen["harness_config"] = config - return LocalHarness(config=config) - - module.load_taskset = load_taskset - module.load_harness = load_harness - monkeypatch.setitem(sys.modules, module_name, module) - - mapped = vf.load_taskset( - "component-loader-config", - config={"taskset_id": "mapped", "split": "test"}, - ) - base = vf.load_taskset( - "component-loader-config", - config=TasksetConfig(taskset_id="base"), - ) - concrete = vf.load_taskset( - "component-loader-config", - config=LocalTasksetConfig(taskset_id="concrete", split="dev"), - ) - harness = vf.load_harness( - "component-loader-config", - config={"model": {"name": "configured-model"}, "mode": "custom"}, - ) - - assert isinstance(seen["taskset_config"], LocalTasksetConfig) - assert mapped.config.taskset_id == "mapped" - assert mapped.config.split == "test" - assert base.config.taskset_id == "base" - assert base.config.split == "train" - assert concrete.config.taskset_id == "concrete" - assert concrete.config.split == "dev" - assert isinstance(seen["harness_config"], LocalHarnessConfig) - assert harness.config.model.name == "configured-model" - assert harness.config.mode == "custom" - - -def test_public_component_loaders_default_to_caller_module( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "current_module_component_loader" - module = types.ModuleType(module_name) - exec( - """ -import verifiers as vf - - -def load_tasks(split: vf.TaskSplit = "train") -> vf.Tasks: - _ = split - return [{"prompt": [], "answer": "current"}] - - -class LocalTasksetConfig(vf.TasksetConfig): - split: str = "train" - - -class LocalHarnessConfig(vf.HarnessConfig): - mode: str = "default" - - -class LocalHarness(vf.Harness): - config: LocalHarnessConfig - - -class LocalTaskset(vf.Taskset[LocalTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - -def load_taskset(config: LocalTasksetConfig) -> vf.Taskset: - return LocalTaskset(config=config) - - -def load_harness(config: LocalHarnessConfig) -> LocalHarness: - return LocalHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -""", - module.__dict__, - ) - monkeypatch.setitem(sys.modules, module_name, module) - - env = module.load_environment( - config=EnvConfig( - taskset=TasksetConfig(taskset_id="current"), - harness=HarnessConfig(model=vf.ModelConfig(name="configured-model")), - ) - ) - - assert env.taskset.config.taskset_id == "current" - assert env.taskset.config.split == "train" - assert env.taskset.get_dataset()[0]["answer"] == "current" - assert env.harness.config.model.name == "configured-model" - assert env.harness.config.mode == "default" - - -def test_load_environment_taskset_loader_uses_registered_taskset_class( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "registered_taskset_component_loader" - module = types.ModuleType(module_name) - exec( - """ -import verifiers as vf - - -def load_tasks(split: vf.TaskSplit = "train") -> vf.Tasks: - _ = split - return [{"prompt": [], "answer": "module"}] - - -class LocalTasksetConfig(vf.TasksetConfig): - dataset_name: str = "configured" - dataset_split: str = "train" - - -class LocalTaskset(vf.Taskset[LocalTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return [ - { - "prompt": [], - "answer": f"{self.config.dataset_name}:{self.config.dataset_split}", - } - ] - - -def load_taskset(config: LocalTasksetConfig) -> vf.Taskset: - return LocalTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env(taskset=vf.load_taskset(config=config.taskset)) -""", - module.__dict__, - ) - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "registered-taskset-component-loader", - config={"taskset": {"dataset_split": "eval"}}, - ) - configured = vf.load_environment( - "registered-taskset-component-loader", - config={"taskset": {}}, - ) - - assert type(env.taskset).__name__ == "LocalTaskset" - assert env.taskset.get_dataset()[0]["answer"] == "configured:eval" - assert configured.taskset.get_dataset()[0]["answer"] == "configured:train" - - -def test_load_environment_composes_component_package_without_root_loader( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "component_only_taskset_package" - module = types.ModuleType(module_name) - - class LocalTasksetConfig(TasksetConfig): - answer: str = "configured" - - class LocalTaskset(Taskset[LocalTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return [{"prompt": [], "answer": f"{split}:{self.config.answer}"}] - - def load_taskset(config: LocalTasksetConfig) -> LocalTaskset: - return LocalTaskset(config=config) - - module.load_taskset = load_taskset - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "component-only-taskset-package", - config={ - "taskset": {"answer": "composed"}, - "harness": {"max_turns": 3}, - }, - ) - - assert env.taskset.get_dataset()[0]["answer"] == "train:composed" - assert type(env.harness) is Harness - assert env.harness.config.max_turns == 3 - - -def test_load_environment_delegates_missing_child_loaders_by_config_id( - monkeypatch: pytest.MonkeyPatch, -) -> None: - env_module = types.ModuleType("thin_env_package") - exec( - """ -import verifiers as vf - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -""", - env_module.__dict__, - ) - taskset_module = types.ModuleType("external_taskset_pkg") - exec( - """ -import verifiers as vf - - -class ExternalTasksetConfig(vf.TasksetConfig): - answer: str = "external" - - -class ExternalTaskset(vf.Taskset[ExternalTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return [{"prompt": [], "answer": f"{split}:{self.config.answer}"}] - - -def load_taskset(config: ExternalTasksetConfig) -> ExternalTaskset: - return ExternalTaskset(config=config) -""", - taskset_module.__dict__, - ) - harness_module = types.ModuleType("external_harness_pkg") - exec( - """ -import verifiers as vf - - -class ExternalHarnessConfig(vf.HarnessConfig): - mode: str = "default" - - -class ExternalHarness(vf.Harness[ExternalHarnessConfig]): - pass - - -def load_harness(config: ExternalHarnessConfig) -> ExternalHarness: - return ExternalHarness(config=config) -""", - harness_module.__dict__, - ) - monkeypatch.setitem(sys.modules, "thin_env_package", env_module) - monkeypatch.setitem( - sys.modules, "empty_env_package", types.ModuleType("empty_env_package") - ) - monkeypatch.setitem(sys.modules, "external_taskset_pkg", taskset_module) - monkeypatch.setitem(sys.modules, "external_harness_pkg", harness_module) - - config = { - "taskset": {"id": "external-taskset-pkg", "answer": "delegated"}, - "harness": {"id": "external-harness-pkg", "mode": "custom"}, - } - for env_id in ("thin-env-package", "empty-env-package"): - env = vf.load_environment(env_id, config=config) - - assert env.taskset.get_dataset()[0]["answer"] == "train:delegated" - assert type(env.taskset).__name__ == "ExternalTaskset" - assert type(env.harness).__name__ == "ExternalHarness" - assert env.harness.config.mode == "custom" - - -def test_load_environment_coerces_base_env_config_with_factory_annotations( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "factory_typed_child_config_object" - module = types.ModuleType(module_name) - - class LocalTasksetConfig(TasksetConfig): - split: str = "train" - - def load_taskset(config: LocalTasksetConfig) -> Taskset: - return Taskset(config=config) - - def load_environment(config: EnvConfig) -> Env: - return Env(taskset=vf.load_taskset(module_name, config=config.taskset)) - - module.load_taskset = load_taskset - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment( - "factory-typed-child-config-object", - config=EnvConfig(taskset=TasksetConfig(taskset_id="typed")), - ) - - assert isinstance(env.taskset.config, LocalTasksetConfig) - assert env.taskset.config.taskset_id == "typed" - - -def test_load_environment_supplies_default_typed_env_config( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "default_typed_env_config" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - def load_environment(config: EnvConfig) -> Env: - seen["config"] = config - return Env( - taskset=make_taskset(config=config.taskset), - harness=make_harness(config=config.harness), - ) - - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - env = vf.load_environment("default-typed-env-config") - - assert isinstance(seen["config"], EnvConfig) - assert env.env_args == {} - - -def test_load_environment_rejects_none_typed_env_config( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "none_typed_env_config" - module = types.ModuleType(module_name) - - def load_environment(config: EnvConfig) -> Env: - return Env( - taskset=make_taskset(config=config.taskset), - harness=make_harness(config=config.harness), - ) - - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - with pytest.raises(RuntimeError, match="concrete EnvConfig object"): - vf.load_environment("none-typed-env-config", config=None) - - -def test_load_environment_leaves_untyped_config_arg_as_kwargs( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "untyped_env_config" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - def load_environment(split: str = "train", config=None) -> Env: - seen["split"] = split - seen["config"] = config - return Env(taskset=make_taskset()) - - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - vf.load_environment( - "untyped-env-config", - split="test", - config={"taskset": {"taskset_id": "raw"}}, - ) - - assert seen["split"] == "test" - assert seen["config"] == {"taskset": {"taskset_id": "raw"}} - - -def test_load_environment_leaves_non_v1_config_annotation_as_kwargs( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module_name = "dict_typed_env_config" - module = types.ModuleType(module_name) - seen: dict[str, object] = {} - - def load_environment(config: dict[str, object]) -> Env: - seen["config"] = config - return Env(taskset=make_taskset()) - - module.load_environment = load_environment - monkeypatch.setitem(sys.modules, module_name, module) - - vf.load_environment( - "dict-typed-env-config", - config={"taskset": {"taskset_id": "raw"}}, - ) - - assert seen["config"] == {"taskset": {"taskset_id": "raw"}} - - -def test_config_objects_are_strict_when_projected_to_base_config_fields() -> None: - class LocalHarnessConfig(HarnessConfig): - toolset: dict[str, object] | None = None - - config = LocalHarnessConfig( - model=vf.ModelConfig(name="parent"), - toolset={"show": ["search"]}, - ) - - assert config.toolset == {"show": ["search"]} - with pytest.raises(ValueError): - HarnessConfig.model_validate(config.model_dump()) - - -def test_unset_base_config_defaults_do_not_override_child_defaults() -> None: - class LocalHarnessConfig(HarnessConfig): - max_turns: int = 4 - - default_config = LocalHarnessConfig.model_validate({}) - explicit_config = LocalHarnessConfig.model_validate({"max_turns": 10}) - - assert default_config.max_turns == 4 - assert explicit_config.max_turns == 10 - - -def test_config_field_name_is_ordinary_serializable_data() -> None: - class LocalTasksetConfig(TasksetConfig): - config: dict[str, object] | None = None - - config = LocalTasksetConfig.model_validate({"config": {"mode": "loaded"}}) - - assert config.config == {"mode": "loaded"} - - -@pytest.mark.parametrize( - "module_name", - [ - "environments.dspy_flights.dspy_flights", - "environments.hello_group_reward_v1.hello_group_reward_v1", - "environments.hello_parallel_sandbox_v1.hello_parallel_sandbox_v1", - "environments.hello_rlm_v1.hello_rlm_v1", - "environments.hello_self_judge_v1.hello_self_judge_v1", - "environments.hello_subagent_v1.hello_subagent_v1", - "environments.nested_harness_v1.nested_harness_v1", - ], -) -def test_reference_v1_loaders_preserve_mapping_config_sections( - module_name: str, - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = importlib.import_module(module_name) - env_id = module_name.rsplit(".", 1)[-1] - monkeypatch.setitem(sys.modules, env_id, module) - - env = vf.load_environment( - env_id, - config={ - "taskset": {"taskset_id": "from-env-args"}, - "harness": {"model": {"name": "configured-model"}}, - }, - ) - - assert env.taskset.config.taskset_id == "from-env-args" - assert env.harness.config.model.name == "configured-model" - - -def test_reference_v1_harness_loaders_preserve_child_defaults() -> None: - group_reward = importlib.import_module( - "environments.hello_group_reward_v1.hello_group_reward_v1" - ) - parallel_sandbox = importlib.import_module( - "environments.hello_parallel_sandbox_v1.hello_parallel_sandbox_v1" - ) - self_judge = importlib.import_module( - "environments.hello_self_judge_v1.hello_self_judge_v1" - ) - - assert ( - group_reward.GroupRewardHarness( - config=group_reward.GroupRewardHarnessConfig() - ).config.max_turns - == 1 - ) - assert ( - parallel_sandbox.ParallelSandboxHarness( - config=parallel_sandbox.ParallelSandboxHarnessConfig() - ).config.max_turns - == 4 - ) - assert vf.Harness(config=self_judge.SelfJudgeHarnessConfig()).config.max_turns == 8 - - -def test_math_python_v1_wrapper_rejects_unsupported_sandbox_kwargs() -> None: - module = importlib.import_module("environments.math_python.math_python") - - with pytest.raises(TypeError, match="max_startup_wait_seconds"): - module.load_environment(v1=True, max_startup_wait_seconds=10) - with pytest.raises(TypeError, match="sandbox_client_max_workers"): - module.load_environment(v1=True, sandbox_client_max_workers=2) - - -def test_math_python_v1_prompt_tracks_harness_packages() -> None: - module = importlib.import_module("environments.math_python.math_python_v1") - - default_env = module.load_environment(config=module.MathPythonEnvConfig()) - assert "numpy sympy scipy" in default_env.taskset.config.system_prompt - - env = module.load_environment( - config=module.MathPythonEnvConfig( - harness=module.MathPythonHarnessConfig(pip_install_packages="numpy pandas") - ) - ) - - prompt = env.taskset.config.system_prompt - assert "numpy pandas" in prompt - assert "numpy sympy scipy" not in prompt - - -def test_math_python_v1_explicit_prompt_wins() -> None: - module = importlib.import_module("environments.math_python.math_python_v1") - - env = module.load_environment( - config=module.MathPythonEnvConfig( - taskset=module.MathPythonTasksetConfig(system_prompt="custom prompt"), - harness=module.MathPythonHarnessConfig(pip_install_packages="numpy pandas"), - ) - ) - - assert env.taskset.config.system_prompt == "custom prompt" - - -def test_bfcl_loader_preserves_mapping_config_sections( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = importlib.import_module("environments.bfcl_v3.bfcl_v3") - seen: dict[str, object] = {} - - def fake_harness(config: object = None, **kwargs: object) -> Harness: - _ = kwargs - seen["harness_config"] = config - return make_harness(config=config) - - monkeypatch.setattr(module, "load_harness", fake_harness) - - env = module.load_environment( - config=module.BFCLEnvConfig( - taskset={"taskset_id": "bfcl-env-args"}, - harness={"model": {"name": "bfcl-model"}}, - ) - ) - - assert env.taskset.config.taskset_id == "bfcl-env-args" - assert env.harness.config.model.name == "bfcl-model" - assert isinstance(env.taskset.config, module.BFCLTasksetConfig) - assert isinstance(seen["harness_config"], module.BFCLHarnessConfig) - - -def test_self_judge_loader_projects_shortcuts_to_child_configs() -> None: - module = importlib.import_module( - "environments.hello_self_judge_v1.hello_self_judge_v1" - ) - - taskset = module.SelfJudgeTaskset( - config=module.SelfJudgeTasksetConfig(num_examples=2) - ) - harness = vf.Harness(config=module.SelfJudgeHarnessConfig(max_turns=3)) - shortcut_env = module.load_environment( - config=module.SelfJudgeEnvConfig( - taskset={"num_examples": 2}, harness={"max_turns": 3} - ), - ) - override_env = module.load_environment( - config=module.SelfJudgeEnvConfig( - taskset={"num_examples": 1}, harness={"max_turns": 5} - ), - ) - - assert len(taskset.get_dataset()) == 2 - assert harness.config.max_turns == 3 - assert len(shortcut_env.taskset.get_dataset()) == 2 - assert shortcut_env.harness.config.max_turns == 3 - assert len(override_env.taskset.get_dataset()) == 1 - assert override_env.harness.config.max_turns == 5 - - -def test_subagent_loader_keeps_child_harness_internal( - monkeypatch: pytest.MonkeyPatch, -) -> None: - module = importlib.import_module("environments.hello_subagent_v1.hello_subagent_v1") - monkeypatch.setitem(sys.modules, "hello_subagent_v1", module) - - env = vf.load_environment( - "hello-subagent-v1", config={"harness": {"model": {"name": "parent"}}} - ) - - assert env.harness.config.model.name == "parent" - toolset = env.harness.toolsets[0] - assert toolset.bindings == {} - assert toolset.tools == (module.ask_subagent,) - assert not hasattr(module, "load_child_harness") - - -def test_nested_configs_validate_and_feed_runtime_objects() -> None: - sandbox = vf.SandboxConfig( - image="python:3.12-slim", - packages="numpy", - setup_commands="echo ready", - scope="group", - create_concurrency=7, - create_rate_per_second=1.5, - delete_concurrency=3, - delete_rate_per_second=2.5, - ) - harness = make_harness(program={"sandbox": True}, sandbox=sandbox) - - assert harness.sandbox is not None - assert harness.sandbox.image == "python:3.12-slim" - assert harness.sandbox.packages == ["numpy"] - assert harness.sandbox.setup_commands == ["echo ready"] - assert harness.sandbox.scope == "group" - assert harness.sandbox.create_concurrency == 7 - assert harness.sandbox.create_rate_per_second == 1.5 - assert harness.sandbox.delete_concurrency == 3 - assert harness.sandbox.delete_rate_per_second == 2.5 - - toolset = Toolset( - config=vf.ToolsetConfig( - tools=[ref("hidden_tool")], - show=["hidden_tool"], - sandbox=vf.SandboxConfig(prefer="program"), - write=True, - ) - ) - - assert toolset.tools == (hidden_tool,) - assert toolset.show == ("hidden_tool",) - assert toolset.write is True - assert isinstance(toolset.sandbox, vf.SandboxConfig) - assert toolset.sandbox.prefer == "program" - - -def test_nested_configs_reject_unknown_fields() -> None: - with pytest.raises(ValueError): - vf.SandboxConfig.model_validate({"image": "python:3.11", "unknown": True}) - - with pytest.raises(ValueError): - HarnessConfig.model_validate({"sandbox_create_concurrency": 2}) - - with pytest.raises(ValueError): - vf.ToolsetConfig.model_validate({"tools": [], "show": ["a"], "hide": ["b"]}) - - -def test_configs_validate_toml_sections(tmp_path) -> None: - config_path = tmp_path / "env.toml" - config_path.write_text( - "\n".join( - [ - "[env.taskset]", - "[[env.taskset.rewards]]", - f'fn = "{ref("config_reward")}"', - "weight = 0.5", - "", - "[env.taskset.toolsets.configured]", - f'fn = "{ref("config_toolset")}"', - 'prefix = "toml"', - "", - "[env.harness]", - "max_turns = 7", - "", - "[env.harness.program]", - f'fn = "{ref("config_program")}"', - ] - ) - ) - - with config_path.open("rb") as f: - data = load_toml(f)["env"] - taskset_config = TasksetConfig.model_validate(data["taskset"]) - harness_config = HarnessConfig.model_validate(data["harness"]) - - class TomlTaskset(Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_tasks(split) - - taskset = TomlTaskset(config=taskset_config) - harness = make_harness(config=harness_config) - - assert taskset.get_dataset()[0]["answer"] == "ok" - assert getattr(taskset.rewards[0], "__name__") == "config_reward" - assert getattr(taskset.rewards[0], "reward_weight") == 0.5 - prefix = taskset.named_toolsets["configured"].bindings["config_tool.prefix"] - assert callable(prefix) - assert prefix() == "toml" - assert harness.config.program.data() == {"fn": ref("config_program")} - assert callable(harness.program) - assert harness.config.max_turns == 7 - - -@pytest.mark.asyncio -async def test_task_tools_filter_exposed_tools() -> None: - harness = make_harness(toolsets={"main": Toolset(tools=[direct_tool, hidden_tool])}) - task = Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "tools": {"main": {"show": ["direct_tool"]}}, - } - ).freeze() - state = await harness.setup_state(task, State.for_task(task)) - - assert state["tools"] == ["direct_tool"] - assert [tool.name for tool in harness.runtime.tool_defs(state) or []] == [ - "direct_tool" - ] - assert list(harness.runtime.tool_calls(task, state)) == ["direct_tool"] - - -def test_toolset_config_is_load_bearing() -> None: - toolset = Toolset( - config={ - "tools": [ref("direct_tool"), ref("hidden_tool")], - "objects": {"task_loader": ref("load_tasks")}, - "bindings": {"hidden_tool.prefix": "task.answer"}, - "write": True, - "scope": "group", - "cleanups": [ref("config_cleanup")], - }, - ) - - assert toolset.tools == (direct_tool, hidden_tool) - assert toolset.bindings == {"hidden_tool.prefix": "task.answer"} - assert toolset.objects == {"task_loader": load_tasks} - assert toolset.write is True - assert toolset.scope == "group" - assert toolset.cleanups == (config_cleanup,) - - -def test_inline_toolset_object_refs_resolve() -> None: - toolset = normalize_toolset( - { - "tools": [ref("object_tool")], - "objects": {"box": ref("load_object_box")}, - "bindings": {"object_tool.box": "objects.box"}, - } - ) - - assert toolset.objects == {"box": load_object_box} - - -def test_toolset_rejects_mixed_config_and_constructor_fields() -> None: - with pytest.raises(ValueError, match="either config or constructor fields"): - Toolset(write=False, config={"write": True}) - - -def test_toolset_sandbox_prefer_requires_program() -> None: - with pytest.raises(ValueError, match="Input should be 'program'"): - vf.SandboxConfig(prefer="other") - - -def test_toolset_sandbox_requires_config_object() -> None: - with pytest.raises(TypeError, match="SandboxConfig"): - Toolset(sandbox={"prefer": "other"}) - - -def test_toolset_config_accepts_mcp_tool_specs() -> None: - toolset = Toolset( - config={ - "tools": [ - { - "command": "uvx", - "args": ["mcp-server-fetch"], - "env": {"API_KEY": "test"}, - "cwd": "/tmp", - } - ], - } - ) - - assert isinstance(toolset.tools[0], vf.MCPTool) - assert toolset.tools[0].command == "uvx" - assert toolset.tools[0].args == ("mcp-server-fetch",) - assert toolset.tools[0].env == {"API_KEY": "test"} - assert toolset.tools[0].cwd == "/tmp" - - -def test_add_toolset_accepts_same_shapes_as_constructor() -> None: - taskset = make_taskset() - harness = make_harness() - - taskset.add_toolset({"direct": Toolset(tools=[direct_tool])}) - harness.add_toolset({"configured": config_toolset}) - - assert taskset.named_toolsets["direct"].tools == (direct_tool,) - assert harness.named_toolsets["configured"].tools == (config_tool,) - - -def test_taskset_extension_is_available_when_env_binds_runtime() -> None: - taskset = make_taskset() - taskset.add_toolset({"direct": Toolset(tools=[direct_tool])}) - harness = Env(taskset=taskset, harness=make_harness()).harness - - assert "direct" in harness.runtime.named_toolsets - - -def test_taskset_extension_refreshes_bound_env_runtime() -> None: - taskset = make_taskset() - harness = Env(taskset=taskset, harness=make_harness()).harness - - taskset.add_toolset({"direct": Toolset(tools=[direct_tool])}) - - assert "direct" in harness.runtime.named_toolsets diff --git a/tests/test_v1_core.py b/tests/test_v1_core.py new file mode 100644 index 0000000000..e013da54bc --- /dev/null +++ b/tests/test_v1_core.py @@ -0,0 +1,2664 @@ +from __future__ import annotations + +import asyncio +import base64 +from contextlib import asynccontextmanager +import io +import importlib +import json +import os +import sys +import tarfile +from types import ModuleType + +from aiohttp import ClientSession, web +import pytest + +import verifiers.v1 as vf +from verifiers.errors import InfraError, ToolError +from verifiers.types import ClientConfig, Response +from verifiers.v1.loaders import ( + load_environment_from_components, + load_harness_from_module, + load_taskset_from_module, +) +from verifiers.v1.mcp import BoundUpdate, MCPToolRegistry, split_result +from verifiers.v1.protocols import parse_anthropic_user_messages +from verifiers.v1.toolset import ToolBinding + + +def attach_mock_model( + env: vf.Env, mock_client, model: str = "test-model" +) -> vf.ModelConfig: + config = vf.ModelConfig(client=ClientConfig(), model=model) + env.harness.load_model_client = lambda _: vf.ModelClient( + config=config, client=mock_client + ) + + async def close_model_client(_: vf.ModelClient) -> None: + return None + + env.harness.close_model_client = close_model_client + return config + + +class ExactMatchTask(vf.Task): + answer: str + + +class ExactMatchTaskset(vf.Taskset): + task_type = ExactMatchTask + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [ + ExactMatchTask( + row_id=0, + prompt=[{"role": "user", "content": "say ok"}], + answer="ok", + max_turns=1, + ) + ] + + @vf.setup + async def setup_extras(self, state: vf.State) -> None: + state.extras["seen_setup"] = True + + @vf.reward + async def exact(self, task: ExactMatchTask, state: vf.State) -> float: + completion = state.completion + text = str(completion[-1].content if completion else "") + return float(text.strip() == task.answer) + + +class GroupVariantTask(ExactMatchTask): + variant: int + + +class GroupTaskset(ExactMatchTaskset): + async def init_group( + self, task: ExactMatchTask, num_rollouts: int + ) -> tuple[list[GroupVariantTask], list[vf.State]]: + tasks = [ + GroupVariantTask.model_validate( + { + **task.model_dump( + mode="json", exclude_none=True, exclude_defaults=True + ), + "variant": index, + } + ) + for index in range(num_rollouts) + ] + return tasks, [vf.State(task_id=task.task_id) for task in tasks] + + @vf.reward(stage="group") + async def relative( + self, tasks: list[GroupVariantTask], states: list[vf.State] + ) -> list[float]: + _ = states + return [float(task.variant == 0) for task in tasks] + + +@vf.advantage +def custom_env_advantage(tasks: list[vf.Task], states: list[vf.State]) -> None: + _ = tasks + for index, state in enumerate(states): + value = float(index + 10) + for turn in state.transcript: + if turn.tokens is not None: + turn.tokens.prompt_advantages = [0.0 for _ in turn.tokens.prompt_ids] + turn.tokens.completion_advantages = [ + value for _ in turn.tokens.completion_ids + ] + + +class EmptyPromptUserConfig(vf.UserConfig): + pass + + +class EmptyPromptUser(vf.User[EmptyPromptUserConfig]): + @vf.user + def respond(self) -> dict: + return {"messages": [{"role": "user", "content": "server prompt"}]} + + +class EmptyPromptTasksetConfig(vf.TasksetConfig): + user: vf.UserConfig | None = EmptyPromptUserConfig() + + +class EmptyPromptTaskset(vf.Taskset[EmptyPromptTasksetConfig]): + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [{"example_id": 0, "prompt": [], "max_turns": 1}] + + +class PatchOnlyUserConfig(vf.UserConfig): + pass + + +class PatchOnlyUser(vf.User[PatchOnlyUserConfig]): + @vf.user( + sets={ + "done": "state.extras.done", + "stop_condition": "state.stop_condition", + } + ) + def respond(self) -> dict: + return {"done": True, "stop_condition": "user_bootstrap_done"} + + +class PatchOnlyUserTasksetConfig(vf.TasksetConfig): + user: vf.UserConfig | None = PatchOnlyUserConfig() + + +class PatchOnlyUserTaskset(vf.Taskset[PatchOnlyUserTasksetConfig]): + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [{"example_id": 0, "prompt": [], "max_turns": 1}] + + +class ToolUserConfig(vf.UserConfig): + pass + + +class ToolUser(vf.User[ToolUserConfig]): + @vf.user + def respond(self) -> dict: + return {"messages": [{"role": "user", "content": "server prompt"}]} + + @vf.tool + def ping(self) -> str: + return "pong" + + +class TaskSetupToolsetConfig(vf.ToolsetConfig): + scope: vf.Scope = "env" + + +class TaskSetupToolset(vf.Toolset[TaskSetupToolsetConfig]): + count = 0 + + @vf.tool( + args={"task": "task"}, + sets={ + "grader_case": "state.metadata.grader_case", + "setup_count": "state.extras.setup_count", + "setup": "state.artifacts.setup", + }, + ) + def materialize(self, task: dict) -> dict: + type(self).count += 1 + count = type(self).count + return { + "content": "", + "grader_case": f"{task['task_id']}:{count}", + "setup_count": count, + "setup": {"count": count}, + } + + +class ServerSetupTasksetConfig(vf.TasksetConfig): + toolsets: vf.ToolsetConfigs = { + "setup": TaskSetupToolsetConfig(hide=["materialize"]) + } + + +class ServerSetupTaskset(vf.Taskset[ServerSetupTasksetConfig]): + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [ + {"example_id": 0, "prompt": "say ok", "max_turns": 1}, + {"example_id": 1, "prompt": "say ok", "max_turns": 1}, + ] + + @vf.setup + async def materialize_task( + self, + task: vf.Task, + state: vf.State, + harness: vf.Harness, + runtime: vf.Runtime, + toolsets: MCPToolRegistry, + ) -> None: + _ = task + result = await toolsets.call_hidden("materialize", {}) + harness.apply_tool_result(state, result) + case = state.metadata["grader_case"] + assert isinstance(case, str) + await runtime.write("grader_case.txt", case.encode()) + + +class DemoToolsetConfig(vf.ToolsetConfig): + pass + + +class DemoToolset(vf.Toolset[DemoToolsetConfig]): + @vf.tool + def echo(self, text: str) -> str: + return text.upper() + + @vf.tool + def suffix(self, text: str) -> str: + return text + "!" + + +class BoundToolsetConfig(vf.ToolsetConfig): + pass + + +class BoundToolset(vf.Toolset[BoundToolsetConfig]): + @vf.tool( + args={"task_name": "task.name"}, + extends={"events": "state.extras.events"}, + sets={ + "profile": "state.extras.profile", + }, + ) + def record(self, name: str, task_name: str) -> dict: + return { + "content": f"{task_name}:{name}", + "events": [{"task": task_name, "name": name}], + "profile": {"task": task_name}, + } + + +class DynamicToolsetConfig(vf.ToolsetConfig): + pass + + +class DynamicToolset(vf.Toolset[DynamicToolsetConfig]): + @vf.tool( + hidden=True, + args={"task_name": "task.name"}, + sets={"ready": "state.extras.dynamic_ready"}, + ) + def setup(self, task_name: str) -> dict: + return { + "ready": True, + "tools": [ + { + "name": "dynamic_echo", + "description": f"Echo through {task_name}", + "parameters": { + "type": "object", + "properties": {"text": {"type": "string"}}, + "required": ["text"], + }, + } + ], + } + + @vf.tool(hidden=True, sets={"called": "state.extras.dynamic_called"}) + def call_tool(self, name: str, input: vf.JsonData) -> dict: + return {"content": f"{name}:{input['text']}", "called": True} + + +class DynamicUserConfig(vf.UserConfig): + pass + + +class DynamicUser(vf.User[DynamicUserConfig]): + @vf.tool( + hidden=True, + sets={"ready": "state.extras.user_dynamic_ready"}, + ) + def setup(self) -> dict: + return { + "ready": True, + "tools": [ + { + "name": "user_echo", + "description": "Echo through user server", + "parameters": { + "type": "object", + "properties": {"text": {"type": "string"}}, + "required": ["text"], + }, + } + ], + } + + @vf.tool(hidden=True, sets={"called": "state.extras.user_dynamic_called"}) + def call_tool(self, name: str, input: vf.JsonData) -> dict: + return {"content": f"{name}:{input['text']}", "called": True} + + +class EnvScopedToolsetConfig(vf.ToolsetConfig): + scope: vf.Scope = "env" + + +class EnvScopedUserConfig(vf.UserConfig): + scope: vf.Scope = "env" + + +class EnvServerTasksetConfig(vf.TasksetConfig): + toolsets: vf.ToolsetConfigs = {"demo": EnvScopedToolsetConfig()} + user: vf.UserConfig | None = EnvScopedUserConfig() + + +class EnvServerTaskset(vf.Taskset[EnvServerTasksetConfig]): + pass + + +class ConfiguredToolsetTasksetConfig(vf.TasksetConfig): + toolsets: vf.ToolsetConfigs = {"demo": DemoToolsetConfig()} + + +class ConfiguredToolsetTaskset(vf.Taskset[ConfiguredToolsetTasksetConfig]): + pass + + +def test_v1_state_is_pydantic_and_extras_owned() -> None: + task = ExactMatchTask(prompt="hello", answer="world") + state = vf.State(task_id=task.task_id) + + state.extras["x"] = 1 + state.reward += 0.5 + state.assert_serializable() + + assert state.extras == {"x": 1} + assert state.reward == 0.5 + with pytest.raises(TypeError): + state["x"] = 2 + assert state.to_output(task)["transcript"] == [] + + +def test_v1_state_to_output_preserves_empty_turn_prompt() -> None: + task = vf.Task(prompt="fallback") + state = vf.State(transcript=[vf.Turn(prompt=[], completion=[])]) + + assert state.prompt == [] + assert state.to_output(task)["prompt"] == [] + + +def test_v1_state_to_output_state_columns_preserve_task_prompt_fallback() -> None: + task = vf.Task(prompt="fallback") + state = vf.State() + + assert state.to_output(task, state_columns=["prompt"])["prompt"] == [ + {"role": "user", "content": "fallback"} + ] + + +def test_v1_task_id_is_deterministic_from_task_contents() -> None: + first = ExactMatchTask(prompt="hello", answer="world") + second = ExactMatchTask(prompt="hello", answer="world") + changed = ExactMatchTask(prompt="hello", answer="there") + explicit = ExactMatchTask(task_id="chosen", prompt="hello", answer="world") + + assert first.task_id == second.task_id + assert first.task_id != changed.task_id + assert explicit.task_id == "chosen" + + +def test_v1_state_messages_uses_latest_prompt_once() -> None: + first_prompt = [vf.UserMessage(content="first")] + first_completion = [vf.AssistantMessage(content="one")] + second_prompt = [*first_prompt, *first_completion, vf.UserMessage(content="second")] + second_completion = [vf.AssistantMessage(content="two")] + state = vf.State( + transcript=[ + vf.Turn(prompt=first_prompt, completion=first_completion), + vf.Turn(prompt=second_prompt, completion=second_completion), + ] + ) + + assert [message.content for message in state.messages] == [ + "first", + "one", + "second", + "two", + ] + + +def test_task_user_defaults_to_auto_and_omits_from_json() -> None: + task = vf.Task(prompt="hello") + + assert task.user is None + assert "user" not in task.model_dump(mode="json", exclude_none=True) + assert ( + vf.Task(user=False).model_dump(mode="json", exclude_none=True)["user"] is False + ) + + +def test_v1_loader_does_not_recurse_to_same_package_config_id() -> None: + module = ModuleType("same_env.taskset") + + with pytest.raises(AttributeError, match="does not expose load_taskset"): + load_taskset_from_module(module, config={"id": "same-env"}) + + with pytest.raises(AttributeError, match="does not expose load_harness"): + load_harness_from_module(module, config={"id": "same-env"}) + + +def test_v1_harness_allows_metrics_but_rejects_rewards() -> None: + class MetricHarness(vf.Harness): + @vf.metric + async def command_calls(self) -> float: + return 1.0 + + assert MetricHarness().signals[0]["name"] == "command_calls" + + class RewardHarness(vf.Harness): + @vf.reward + async def execution_reward(self) -> float: + return 1.0 + + with pytest.raises(ValueError, match="Harness signals must be metrics"): + RewardHarness() + + +class TaskExtras(vf.Extras): + task_flag: bool = True + + +class HarnessExtras(vf.Extras): + harness_count: int = 2 + + +class ExtrasTasksetConfig(vf.TasksetConfig): + extras: TaskExtras = TaskExtras() + + +class ExtrasHarnessConfig(vf.HarnessConfig): + extras: HarnessExtras = HarnessExtras() + max_turns: int = 1 + + +class ExtrasTaskset(vf.Taskset[ExtrasTasksetConfig]): + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + if split == "eval": + return [] + return [{"example_id": 0, "prompt": "say ok"}] + + +def test_v1_extras_config_schemas_realize_and_reject_conflicts() -> None: + env = vf.Env(taskset=ExtrasTaskset(), harness=vf.Harness(ExtrasHarnessConfig())) + state = vf.State() + env.harness.initialize_extras(state) + + assert state.extras == {"task_flag": True, "harness_count": 2} + env.harness.validate_extras(state) + + class TaskConflict(vf.Extras): + shared: int = 1 + + class HarnessConflict(vf.Extras): + shared: str = "x" + + class ConflictTasksetConfig(vf.TasksetConfig): + extras: TaskConflict = TaskConflict() + + class ConflictHarnessConfig(vf.HarnessConfig): + extras: HarnessConflict = HarnessConflict() + + class ConflictTaskset(vf.Taskset[ConflictTasksetConfig]): + pass + + with pytest.raises(ValueError, match="defined by both"): + vf.Env( + taskset=ConflictTaskset(), + harness=vf.Harness(ConflictHarnessConfig()), + ) + + +def test_toolset_config_sources_and_enabled_flags_resolve_directly() -> None: + taskset = ConfiguredToolsetTaskset( + config={"toolsets": {"demo": {"enabled": False}}} + ) + assert taskset.toolsets == {} + + taskset = ConfiguredToolsetTaskset( + config={ + "toolsets": { + "custom": { + "source": f"{__name__}:DemoToolsetConfig", + } + } + } + ) + assert list(taskset.toolsets) == ["demo", "custom"] + + with pytest.raises(ValueError, match="cannot be disabled"): + ConfiguredToolsetTaskset(config={"toolsets": {"custom": {"enabled": False}}}) + + with pytest.raises(ValueError, match="set source"): + ConfiguredToolsetTaskset(config={"toolsets": {"custom": {}}}) + + with pytest.raises(TypeError, match="source must match"): + ConfiguredToolsetTaskset( + config={ + "toolsets": { + "demo": { + "source": f"{__name__}:BoundToolsetConfig", + } + } + } + ) + + +@pytest.mark.asyncio +async def test_v1_model_client_uses_serialized_transcript_record(mock_client) -> None: + config = vf.ModelConfig(client=ClientConfig(), model="test-model") + state = vf.State(task_id="task-1") + state.transcript.append( + vf.Turn( + prompt=[vf.UserMessage(content="first")], + completion=[vf.AssistantMessage(content="done")], + tokens=vf.TurnTokens( + prompt_ids=[1], + prompt_mask=[0], + completion_ids=[2], + completion_mask=[1], + completion_logprobs=[-0.1], + ), + ) + ) + + assert [message.content for message in state.messages] == ["first", "done"] + + await vf.ModelClient(config=config, client=mock_client).get_response( + prompt=[vf.UserMessage(content="next")], + state=state, + ) + + client_state = mock_client.last_call_kwargs["state"] + assert "trajectory" not in client_state + assert client_state["task_id"] == "task-1" + assert client_state["transcript"][0]["prompt"][0]["content"] == "first" + assert client_state["transcript"][0]["tokens"]["completion_ids"] == [2] + + +def test_v1_model_client_renderer_handle_is_live_only(mock_client) -> None: + config = vf.ModelConfig(client=ClientConfig(), model="test-model") + model = vf.ModelClient(config=config, client=mock_client) + assert model.get_renderer() is None + + renderer = object() + model = vf.ModelClient(config=config, client=mock_client, renderer=renderer) + assert model.get_renderer() is renderer + + +@pytest.mark.asyncio +async def test_v1_standalone_env_rollout_scores_from_transcript(mock_client) -> None: + mock_client.set_default_response("ok") + env = vf.Env(taskset=ExactMatchTaskset(), harness=vf.Harness()) + model = attach_mock_model(env, mock_client) + row = env.get_dataset()[0] + task = env.taskset.to_task(row) + + state = await env.run_rollout(row, model=model) + output = state.to_output(task) + + assert output["reward"] == 1.0 + assert output["metrics"]["exact"] == 1.0 + assert output["metrics"]["num_turns"] == 1.0 + assert output["transcript"][0]["completion"][0]["content"] == "ok" + + +@pytest.mark.asyncio +async def test_harness_run_accepts_string_task_and_model(mock_client) -> None: + mock_client.set_default_response("ok") + harness = vf.Harness() + configs: list[vf.ModelConfig] = [] + + def load_model_client(config: vf.ModelConfig) -> vf.ModelClient: + configs.append(config) + return vf.ModelClient(config=config, client=mock_client) + + async def close_model_client(_: vf.ModelClient) -> None: + return None + + harness.load_model_client = load_model_client + harness.close_model_client = close_model_client + + state = await harness.run(task="hello world", model="openai/gpt-5") + + assert state.task_id is not None + assert state.prompt[0].content == "hello world" + assert state.completion[0].content == "ok" + assert configs[0].model == "openai/gpt-5" + assert configs[0].client.api_key_var == "PRIME_API_KEY" + assert configs[0].client.api_base_url == "https://api.pinference.ai/api/v1" + assert mock_client.last_call_kwargs["model"] == "openai/gpt-5" + + +@pytest.mark.asyncio +async def test_harness_run_requires_user_when_task_user_is_true(mock_client) -> None: + mock_client.set_default_response("ok") + env = vf.Env(taskset=ExactMatchTaskset(), harness=vf.Harness()) + model = attach_mock_model(env, mock_client) + + state = await env.run_rollout( + {"prompt": "say ok", "answer": "ok", "user": True, "max_turns": 1}, + model=model, + ) + + assert state.stop_condition == "has_error" + assert state.error is not None + assert "requires a user server" in state.error["message"] + + +@pytest.mark.asyncio +async def test_harness_run_rejects_nested_scoring_context(mock_client) -> None: + task = vf.Task(prompt="judge") + state = vf.State(task_id=task.task_id) + model = vf.ModelConfig(client=ClientConfig(), model="test-model") + context = vf.Context( + task=task, + state=state, + model_client=vf.ModelClient(config=model, client=mock_client), + scoring=True, + ) + + with pytest.raises(RuntimeError, match="Nested scored harness runs"): + await vf.Harness().run(task="nested judge", context=context, score=True) + + +@pytest.mark.asyncio +async def test_harness_run_with_context_preserves_parent_state_task_id( + mock_client, +) -> None: + mock_client.set_default_response("child") + harness = vf.Harness() + model = vf.ModelConfig(model="test-model") + + def load_model_client(config: vf.ModelConfig) -> vf.ModelClient: + return vf.ModelClient(config=config, client=mock_client) + + async def close_model_client(_: vf.ModelClient) -> None: + return None + + harness.load_model_client = load_model_client + harness.close_model_client = close_model_client + parent_task = vf.Task(task_id="parent-task", prompt="parent") + state = vf.State(task_id="parent-task") + + async with harness.open_context( + task=parent_task, + state=state, + model=model, + ) as context: + await harness.run("child prompt", context=context) + + assert state.task_id == "parent-task" + assert state.transcript[-1].prompt[-1].content == "child prompt" + + +@pytest.mark.asyncio +async def test_v1_group_rewards_and_advantages_apply_to_turns(mock_client) -> None: + mock_client.set_default_response("ok") + env = vf.Env(taskset=GroupTaskset(), harness=vf.Harness(), advantage="grpo") + model = attach_mock_model(env, mock_client) + row = env.get_dataset()[0] + base_task = env.taskset.to_task(row) + tasks, states = await env.taskset.init_group(base_task, 2) + + states = await asyncio.gather( + *[ + env.run_rollout(task, model=model, state=state) + for task, state in zip(tasks, states, strict=True) + ] + ) + states = await env.score_group(tasks, states) + outputs = [state.to_output(task) for task, state in zip(tasks, states, strict=True)] + + assert [output["reward"] for output in outputs] == [2.0, 1.0] + assert "advantage" not in outputs[0] + assert "advantage" not in outputs[0]["transcript"][0] + + +@pytest.mark.asyncio +async def test_v1_empty_prompt_bootstraps_from_user_server(mock_client) -> None: + mock_client.set_default_response("ok") + env = vf.Env(taskset=EmptyPromptTaskset(), harness=vf.Harness()) + model = attach_mock_model(env, mock_client) + row = env.get_dataset()[0] + task = env.taskset.to_task(row) + + state = await env.run_rollout(row, model=model) + output = state.to_output(task) + + prompt = mock_client.last_call_kwargs["prompt"] + assert prompt[-1].role == "user" + assert prompt[-1].content == "server prompt" + assert output["completion"][0]["content"] == "ok" + assert output["transcript"][0]["prompt"][-1]["content"] == "server prompt" + + +@pytest.mark.asyncio +async def test_user_server_can_return_only_bound_state_updates(mock_client) -> None: + mock_client.set_default_response("should not run") + env = vf.Env(taskset=PatchOnlyUserTaskset(), harness=vf.Harness()) + model = attach_mock_model(env, mock_client) + + state = await env.run_rollout(env.get_dataset()[0], model=model) + + assert state.extras["done"] is True + assert state.stop_condition == "user_bootstrap_done" + assert state.transcript == [] + assert mock_client.last_call_kwargs == {} + await env.close() + + +@pytest.mark.asyncio +async def test_user_server_hides_respond_and_exposes_user_tools() -> None: + async with MCPToolRegistry({"user": ToolUserConfig()}) as registry: + names = [tool.name for tool in registry.tools() or []] + prompt = await registry.call_hidden("respond", {}) + result = await registry.call("user_ping", {}) + with pytest.raises(ToolError, match="disabled"): + await registry.call("user_respond", {}) + + assert names == ["user_ping"] + assert prompt.response.messages[0].content == "server prompt" + assert result.response.content == "pong" + + +@pytest.mark.asyncio +async def test_v1_setup_can_materialize_task_local_state_from_env_server( + mock_client, +) -> None: + mock_client.set_default_response("ok") + env = vf.Env(taskset=ServerSetupTaskset(), harness=vf.Harness()) + model = attach_mock_model(env, mock_client) + rows = list(env.get_dataset()) + + async with env.run() as env_run: + states = [await env_run.run_rollout(row, model=model) for row in rows] + + assert [state.extras["setup_count"] for state in states] == [1, 2] + assert [state.artifacts["setup"]["count"] for state in states] == [1, 2] + assert states[0].metadata["grader_case"].endswith(":1") + assert states[1].metadata["grader_case"].endswith(":2") + assert "setup_materialize" not in [ + tool.name for tool in mock_client.last_call_kwargs["tools"] or [] + ] + await env.close() + + +@pytest.mark.asyncio +async def test_migrated_hello_rlm_v1_runs_without_old_runtime(mock_client) -> None: + from environments.hello_rlm_v1.hello_rlm_v1 import taskset as module + + env = load_environment_from_components( + module, + { + "config": { + "harness": { + "command": [ + sys.executable, + "-c", + "import os; print(os.environ['VF_PROMPT'].split('exactly ', 1)[1].rstrip('.'))", + ] + } + } + }, + ) + model = attach_mock_model(env, mock_client, "unused-model") + row = env.get_dataset()[0] + task = env.taskset.to_task(row) + + state = await env.run_rollout(row, model=model) + output = state.to_output(task) + + assert output["reward"] == 1.0 + assert output["metrics"]["exact_answer"] == 1.0 + assert output["transcript"][0]["completion"][0]["content"] == output["answer"] + + +@pytest.mark.asyncio +async def test_bfcl_multi_turn_respects_completed_state( + monkeypatch, mock_client +) -> None: + from environments.bfcl_v3_v1.bfcl_v3_v1 import taskset as bfcl + + for package_name in [ + "bfcl_eval", + "bfcl_eval.constants", + "bfcl_eval.eval_checker", + "bfcl_eval.eval_checker.multi_turn_eval", + "bfcl_eval.model_handler", + ]: + package = ModuleType(package_name) + package.__path__ = [] + monkeypatch.setitem(sys.modules, package_name, package) + + default_prompts = ModuleType("bfcl_eval.constants.default_prompts") + default_prompts.DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC = "add tools" + monkeypatch.setitem( + sys.modules, + "bfcl_eval.constants.default_prompts", + default_prompts, + ) + + simulator_calls: list[list[str]] = [] + multi_turn_utils = ModuleType( + "bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils" + ) + + def execute_multi_turn_func_call( + func_call_list: list[str], + initial_config: vf.JsonData, + involved_classes: list[str], + model_name: str, + test_entry_id: str, + *, + long_context: bool, + ) -> tuple[list[str], None]: + _ = initial_config, involved_classes, model_name, test_entry_id, long_context + simulator_calls.append(func_call_list) + return [], None + + multi_turn_utils.execute_multi_turn_func_call = execute_multi_turn_func_call + monkeypatch.setitem( + sys.modules, + "bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils", + multi_turn_utils, + ) + + base_handler = ModuleType("bfcl_eval.model_handler.base_handler") + base_handler.is_empty_execute_response = lambda value: not value + monkeypatch.setitem( + sys.modules, + "bfcl_eval.model_handler.base_handler", + base_handler, + ) + monkeypatch.setattr(bfcl, "bfcl_tool_defs", lambda _functions: []) + monkeypatch.setattr(bfcl, "bfcl_missed_function", lambda _task: {}) + monkeypatch.setattr(bfcl, "bfcl_involved_classes", lambda _task: []) + + task = bfcl.BFCLTask( + row_id=0, + prompt="first", + category="multi_turn_base", + question=[ + [{"role": "user", "content": "first"}], + [{"role": "user", "content": "second"}], + ], + function=[], + initial_config={}, + involved_classes=[], + max_steps_per_turn=1, + max_turns=2, + ) + state = vf.State(task_id=task.task_id) + state.stop("precompleted") + model = vf.ModelConfig(client=ClientConfig(), model="test-model") + context = vf.Context( + task=task, + state=state, + model_client=vf.ModelClient(config=model, client=mock_client), + ) + + await bfcl.BFCLHarness(config=bfcl.BFCLHarnessConfig()).run_multi_turn( + context, task, state + ) + + assert simulator_calls == [[]] + assert mock_client.last_call_kwargs == {} + assert state.stop_condition == "precompleted" + + bounded_task = task.model_copy(update={"task_id": "bounded", "max_turns": 1}) + bounded_state = vf.State(task_id=bounded_task.task_id) + bounded_context = vf.Context( + task=bounded_task, + state=bounded_state, + model_client=vf.ModelClient(config=model, client=mock_client), + ) + + mock_client.set_default_response("done") + await bfcl.BFCLHarness(config=bfcl.BFCLHarnessConfig()).run_multi_turn( + bounded_context, bounded_task, bounded_state + ) + + assert mock_client.call_count == 1 + assert len(bounded_state.transcript) == 1 + assert bounded_state.transcript[0].timing.start > 0.0 + assert ( + bounded_state.transcript[0].timing.end + >= bounded_state.transcript[0].timing.start + ) + assert bounded_state.stop_condition == "max_turns" + + +@pytest.mark.asyncio +async def test_mcp_toolset_exposes_multiple_server_tools() -> None: + toolsets = {"demo": DemoToolsetConfig()} + + async with MCPToolRegistry(toolsets) as registry: + names = [tool.name for tool in registry.tools() or []] + result = await registry.call("demo_echo", {"text": "ok"}) + + assert names == ["demo_echo", "demo_suffix"] + assert result.response.content == "OK" + + +@pytest.mark.asyncio +async def test_mcp_owned_runtime_stops_when_server_start_fails(monkeypatch) -> None: + mcp_module = importlib.import_module("verifiers.v1.mcp") + events: list[str] = [] + + class FailingRuntime(vf.Runtime): + async def start(self) -> None: + events.append("start") + + async def stop(self) -> None: + events.append("stop") + + async def expose(self, port: int) -> str: + return f"http://127.0.0.1:{port}" + + async def run( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> vf.CommandResult: + _ = command, cwd, env, timeout + return vf.CommandResult(returncode=0) + + async def read(self, path: str) -> bytes: + _ = path + return b"" + + async def write(self, path: str, data: bytes) -> None: + _ = path, data + + async def run_background( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + log: str | None = None, + ) -> None: + _ = command, cwd, env, log + events.append("run_background") + raise RuntimeError("server failed") + + class FailingProvider(vf.RuntimeProvider): + def create_runtime(self) -> vf.Runtime: + return FailingRuntime() + + monkeypatch.setattr( + mcp_module, + "make_runtime_provider", + lambda _: FailingProvider(), + ) + registry = MCPToolRegistry({}) + + with pytest.raises(RuntimeError, match="server failed"): + async with registry.open_server("demo", DemoToolsetConfig()): + pass + + assert events == ["start", "run_background", "stop"] + + +@pytest.mark.asyncio +async def test_subprocess_runtime_stop_gracefully_interrupts_background_process( + tmp_path, +) -> None: + marker = tmp_path / "cleanup.txt" + runtime = vf.SubprocessRuntime(vf.SubprocessRuntimeConfig()) + await runtime.start() + await runtime.write( + "worker.py", + b""" +import sys +import time + +try: + while True: + time.sleep(0.05) +except KeyboardInterrupt: + pass +finally: + with open(sys.argv[1], "w", encoding="utf-8") as f: + f.write("clean") +""", + ) + await runtime.run_background([sys.executable, "worker.py", str(marker)]) + await asyncio.sleep(0.2) + + await runtime.stop() + + assert marker.read_text() == "clean" + + +@pytest.mark.asyncio +async def test_mcp_registry_closes_partial_stack_when_enter_fails(monkeypatch) -> None: + import mcp.client.session as session_module + + events: list[str] = [] + + class FailingClientSession: + def __init__(self, read: object, write: object) -> None: + _ = read, write + + async def __aenter__(self) -> "FailingClientSession": + events.append("session_enter") + return self + + async def __aexit__(self, *exc: object) -> None: + _ = exc + events.append("session_exit") + + async def initialize(self) -> None: + events.append("initialize") + raise RuntimeError("initialize failed") + + @asynccontextmanager + async def open_server(_name: str, _server: vf.ServerConfig): + events.append("server_enter") + try: + yield object(), object() + finally: + events.append("server_exit") + + monkeypatch.setattr(session_module, "ClientSession", FailingClientSession) + registry = MCPToolRegistry({"demo": DemoToolsetConfig()}) + monkeypatch.setattr(registry, "open_server", open_server) + + with pytest.raises(RuntimeError, match="initialize failed"): + async with registry: + pass + + assert events == [ + "server_enter", + "session_enter", + "initialize", + "session_exit", + "server_exit", + ] + assert registry.tools() is None + + +@pytest.mark.asyncio +async def test_mcp_tool_registry_applies_task_visibility() -> None: + toolsets = {"demo": DemoToolsetConfig()} + + async with MCPToolRegistry(toolsets) as registry: + registry.set_visibility( + toolsets=vf.TaskVisibility(show=["demo"]), + tools=vf.TaskVisibility(hide=["demo_suffix"]), + ) + names = [tool.name for tool in registry.tools() or []] + result = await registry.call("demo_echo", {"text": "ok"}) + with pytest.raises(ToolError, match="disabled"): + await registry.call("demo_suffix", {"text": "ok"}) + + assert names == ["demo_echo"] + assert result.response.content == "OK" + + +@pytest.mark.asyncio +async def test_mcp_bindings_hide_args_and_bind_returns() -> None: + toolsets = {"bound": BoundToolsetConfig()} + state = vf.State() + task = vf.Task(name="demo", prompt="say ok") + harness = vf.Harness() + + async with MCPToolRegistry(toolsets) as registry: + registry.set_context(harness.binding_context(task, state)) + tools = registry.tools() or [] + result = await registry.call("bound_record", {"name": "alpha"}) + + properties = tools[0].parameters["properties"] + assert "name" in properties + assert "task_name" not in properties + assert result.response.content == "demo:alpha" + + harness.apply_bound_updates(state, list(result.updates)) + assert state.extras == { + "events": [{"task": "demo", "name": "alpha"}], + "profile": {"task": "demo"}, + } + + +@pytest.mark.asyncio +async def test_mcp_toolset_setup_registers_dynamic_tools() -> None: + toolsets = {"dynamic": DynamicToolsetConfig()} + state = vf.State() + task = vf.Task(name="demo", prompt="say ok") + harness = vf.Harness() + + async with MCPToolRegistry(toolsets) as registry: + await registry.resolve( + context=harness.binding_context(task, state), + resolution_key=f"{state.id}:{task.task_id}", + apply_updates=lambda updates: harness.apply_bound_updates(state, updates), + ) + names = [tool.name for tool in registry.tools() or []] + assert registry.has_hidden("setup") + assert registry.has_hidden("call_tool") + result = await registry.call("dynamic_echo", {"text": "ok"}) + + harness.apply_tool_result(state, result) + + assert names == ["dynamic_echo"] + assert result.response.content == "dynamic_echo:ok" + assert state.extras["dynamic_ready"] is True + assert state.extras["dynamic_called"] is True + + +@pytest.mark.asyncio +async def test_dynamic_tools_obey_toolset_config_visibility() -> None: + toolsets = {"dynamic": DynamicToolsetConfig(hide=["dynamic_echo"])} + state = vf.State() + task = vf.Task(name="demo", prompt="say ok") + harness = vf.Harness() + + async with MCPToolRegistry(toolsets) as registry: + await registry.resolve( + context=harness.binding_context(task, state), + resolution_key=f"{state.id}:{task.task_id}", + apply_updates=lambda updates: harness.apply_bound_updates(state, updates), + ) + assert registry.tools() is None + with pytest.raises(ToolError, match="Unknown MCP tool"): + await registry.call("dynamic_echo", {"text": "ok"}) + + +@pytest.mark.asyncio +async def test_user_setup_registers_dynamic_model_tools() -> None: + state = vf.State() + task = vf.Task(name="demo", prompt="say ok") + harness = vf.Harness() + + async with MCPToolRegistry({"user": DynamicUserConfig()}) as user_registry: + async with MCPToolRegistry({}, parents=[user_registry]) as registry: + await registry.resolve( + context=harness.binding_context(task, state), + resolution_key=f"{state.id}:{task.task_id}", + apply_updates=lambda updates: harness.apply_bound_updates( + state, updates + ), + ) + tools = registry.tools() or [] + result = await registry.call("user_echo", {"text": "ok"}) + + harness.apply_tool_result(state, result) + + assert [tool.name for tool in tools] == ["user_echo"] + assert result.response.content == "user_echo:ok" + assert state.extras["user_dynamic_ready"] is True + assert state.extras["user_dynamic_called"] is True + + +@pytest.mark.asyncio +async def test_env_scope_servers_start_once_under_concurrent_rollouts( + monkeypatch, +) -> None: + harness_module = importlib.import_module("verifiers.v1.harness") + events: list[str] = [] + + class FakeRegistry: + def __init__(self, servers, **_: object) -> None: + self.kind = "user" if "user" in servers else "toolsets" + + async def __aenter__(self) -> "FakeRegistry": + events.append(f"enter:{self.kind}") + await asyncio.sleep(0.01) + return self + + async def __aexit__(self, *_: object) -> None: + events.append(f"exit:{self.kind}") + + monkeypatch.setattr(harness_module, "MCPToolRegistry", FakeRegistry) + harness = vf.Harness() + harness.bind(taskset=EnvServerTaskset()) + + await asyncio.gather(*(harness.start_env_scope() for _ in range(8))) + await harness.close() + + assert events == [ + "enter:toolsets", + "enter:user", + "exit:user", + "exit:toolsets", + ] + + +@pytest.mark.asyncio +async def test_env_user_startup_failure_closes_started_env_toolsets( + monkeypatch, +) -> None: + harness_module = importlib.import_module("verifiers.v1.harness") + events: list[str] = [] + + class FakeRegistry: + def __init__(self, servers, **_: object) -> None: + self.kind = "user" if "user" in servers else "toolsets" + + async def __aenter__(self) -> "FakeRegistry": + events.append(f"enter:{self.kind}") + if self.kind == "user": + raise RuntimeError("user failed") + return self + + async def __aexit__(self, *_: object) -> None: + events.append(f"exit:{self.kind}") + + monkeypatch.setattr(harness_module, "MCPToolRegistry", FakeRegistry) + harness = vf.Harness() + harness.bind(taskset=EnvServerTaskset()) + + with pytest.raises(RuntimeError, match="user failed"): + await harness.start_env_scope() + + assert events == ["enter:toolsets", "enter:user", "exit:toolsets"] + assert harness._env_toolsets is None + assert harness._env_user is None + + +@pytest.mark.asyncio +async def test_env_scope_startup_failure_closes_entered_env_user( + monkeypatch, +) -> None: + harness_module = importlib.import_module("verifiers.v1.harness") + events: list[str] = [] + + class FakeRegistry: + def __init__(self, servers, **_: object) -> None: + self.kind = "user" if "user" in servers else "toolsets" + + async def __aenter__(self) -> "FakeRegistry": + events.append(f"enter:{self.kind}") + return self + + async def __aexit__(self, *_: object) -> None: + events.append(f"exit:{self.kind}") + + class FailingScopeCountHarness(vf.Harness): + def __init__(self) -> None: + self.fail_entered_scope = False + self.scope_count_value = 0 + super().__init__() + + @property + def _env_scope_count(self) -> int: + return self.scope_count_value + + @_env_scope_count.setter + def _env_scope_count(self, value: int) -> None: + if value == 1 and self.fail_entered_scope: + raise RuntimeError("scope count failed") + self.scope_count_value = value + + monkeypatch.setattr(harness_module, "MCPToolRegistry", FakeRegistry) + harness = FailingScopeCountHarness() + harness.bind(taskset=EnvServerTaskset()) + harness.fail_entered_scope = True + + with pytest.raises(RuntimeError, match="scope count failed"): + await harness.start_env_scope() + + assert events == ["enter:toolsets", "enter:user", "exit:user", "exit:toolsets"] + assert harness._env_toolsets is None + assert harness._env_user is None + + +def test_nemo_gym_task_row_preserves_explicit_task_fields() -> None: + from tasksets.nemo_gym import normalize_nemo_gym_task_row + + row: vf.JsonData = { + "responses_create_params": { + "input": [ + {"role": "system", "content": "source system"}, + {"role": "user", "content": "source prompt"}, + ] + }, + "prompt": [{"role": "user", "content": "override prompt"}], + "system_prompt": [{"role": "system", "content": "override system"}], + "info": {"source": "override"}, + "row_id": 99, + "custom_root_field": "not a task field", + } + + task_row = normalize_nemo_gym_task_row(row, index=3, agent_name="agent") + + assert task_row["prompt"] == row["prompt"] + assert task_row["system_prompt"] == row["system_prompt"] + assert task_row["info"] == { + "source": "override", + "nemo_gym": {"agent_name": "agent"}, + } + assert task_row["row_id"] == 99 + assert "custom_root_field" not in task_row + assert task_row["nemo_gym_row"]["agent_ref"] == { + "type": "responses_api_agents", + "name": "agent", + } + + +def test_bound_state_updates_allow_extends_and_reject_set_conflicts() -> None: + harness = vf.Harness() + state = vf.State() + + harness.apply_bound_updates( + state, + [ + BoundUpdate("extras.events", [{"name": "a"}], "extend"), + BoundUpdate("extras.events", [{"name": "b"}], "extend"), + ], + ) + assert sorted(event["name"] for event in state.extras["events"]) == ["a", "b"] + + with pytest.raises(ValueError, match="Conflicting bound state updates"): + harness.apply_bound_updates( + state, + [ + BoundUpdate("state.extras.profile", {"name": "a"}), + BoundUpdate("state.extras.profile.name", "b"), + ], + ) + + with pytest.raises(ValueError, match="state.advantage"): + harness.apply_bound_updates(state, [BoundUpdate("state.advantage", 1.0)]) + + +def test_bound_state_updates_write_latest_turn_reward() -> None: + harness = vf.Harness() + state = vf.State( + transcript=[ + vf.Turn( + prompt=[vf.UserMessage(content="go")], + completion=[vf.AssistantMessage(content="done")], + ) + ] + ) + + harness.apply_bound_updates( + state, + [BoundUpdate("state.transcript.last.reward", 0.75)], + ) + + assert state.reward == 0.0 + assert state.transcript[-1].reward == 0.75 + + +def test_state_to_output_rejects_metric_reserved_field_collision() -> None: + state = vf.State(metrics={"reward": 1.0}) + task = vf.Task(prompt="hello") + + with pytest.raises(ValueError, match="Metric name 'reward' conflicts"): + state.to_output(task) + + +def test_state_to_output_rejects_state_column_metric_collision() -> None: + state = vf.State(metrics={"custom": 1.0}) + task = vf.Task(prompt="hello") + + with pytest.raises(ValueError, match="State column 'custom' conflicts"): + state.to_output(task, state_columns=["custom"]) + + +def test_split_result_copies_unbound_content() -> None: + content: vf.JsonData = {"messages": [{"role": "user", "content": "before"}]} + result = split_result(content, ToolBinding()) + + messages = content["messages"] + assert isinstance(messages, list) + first = messages[0] + assert isinstance(first, dict) + first["content"] = "after" + + assert result.value == {"messages": [{"role": "user", "content": "before"}]} + + +def test_parse_anthropic_user_messages_preserves_structured_content() -> None: + messages = parse_anthropic_user_messages( + [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "abc", + }, + } + ] + ) + + assert len(messages) == 1 + message = messages[0] + assert isinstance(message, vf.UserMessage) + assert message.content == [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "abc", + }, + } + ] + + +def test_harbor_taskset_maps_task_image_and_resources( + tmp_path, monkeypatch: pytest.MonkeyPatch +) -> None: + from tasksets.harbor import HarborTask, HarborTaskset, HarborTasksetConfig + + root = tmp_path / "tasks" + task_dir = root / "image-task" + task_dir.mkdir(parents=True) + (task_dir / "instruction.md").write_text("fix it\n") + (task_dir / "task.toml").write_text( + "[task]\n" + 'name = "Image task"\n' + 'description = "Check resource parsing."\n' + 'keywords = ["image", "resources"]\n' + "[[task.authors]]\n" + 'name = "Ada"\n' + 'email = "ada@example.com"\n' + "[metadata]\n" + 'difficulty = "easy"\n' + 'category = "smoke"\n' + 'tags = ["typed"]\n' + "[agent]\n" + "timeout_sec = 7\n" + "[verifier]\n" + "timeout_sec = 11\n" + "[environment]\n" + 'docker_image = "owner/task:latest"\n' + "cpus = 2\n" + "memory_mb = 4096\n" + "storage_mb = 10240\n" + "gpus = 1\n" + ) + taskset = HarborTaskset(config=HarborTasksetConfig()) + + task = HarborTask.from_dir(task_dir, require_image=False) + + assert task.image == "owner/task:latest" + assert task.name == "Image task" + assert task.description == "Check resource parsing." + assert task.agent_timeout == 7.0 + assert task.scoring_timeout == 11.0 + assert task.keywords == ["image", "resources"] + assert task.authors[0].name == "Ada" + assert task.authors[0].email == "ada@example.com" + assert task.difficulty == "easy" + assert task.category == "smoke" + assert task.tags == ["typed"] + assert task.resources == vf.Resources( + cpu_cores=2.0, + memory_gb=4.0, + gpu_count=1, + disk_gb=10.0, + ) + task_json = task.model_dump(mode="json") + assert "task_dir" not in task_json + assert "task_toml" not in task_json + assert "runtime_config" not in task_json + assert "program" not in task_json + monkeypatch.setattr(taskset, "task_root", lambda: root) + rehydrated = taskset.to_task(task_json) + + assert isinstance(rehydrated, task.__class__) + assert rehydrated.task_dir == str(task_dir) + assert rehydrated.image == task.image + + +def test_harbor_taskset_rejects_dockerfile_only_tasks(tmp_path) -> None: + from tasksets.harbor import HarborTask + + task_dir = tmp_path / "dockerfile-only" + (task_dir / "environment").mkdir(parents=True) + (task_dir / "environment" / "Dockerfile").write_text("FROM python:3.11\n") + (task_dir / "instruction.md").write_text("fix it\n") + (task_dir / "task.toml").write_text("[environment]\n") + with pytest.raises(ValueError, match="Dockerfile"): + HarborTask.from_dir(task_dir, require_image=False) + + +def test_harbor_taskset_rejects_falsy_non_mapping_sections(tmp_path) -> None: + from tasksets.harbor import HarborTask + + task_dir = tmp_path / "bad-section" + task_dir.mkdir() + (task_dir / "instruction.md").write_text("fix it\n") + (task_dir / "task.toml").write_text("task = false\n") + + with pytest.raises(TypeError, match=r"\[task\] must be a mapping"): + HarborTask.from_dir(task_dir, require_image=False) + + +@pytest.mark.asyncio +async def test_harbor_reward_runs_verifier_in_live_runtime(tmp_path) -> None: + from tasksets.harbor import ( + HarborTask, + HarborTaskset, + HarborTasksetConfig, + ) + + task_dir = tmp_path / "task" + tests_dir = task_dir / "tests" + tests_dir.mkdir(parents=True) + (tests_dir / "test.sh").write_text("echo 1 > /logs/verifier/reward.txt\n") + task = HarborTask( + task_name="task", + instruction="do it", + task_dir=str(task_dir), + prompt=[vf.UserMessage(content="do it")], + image="python:3.11", + scoring_timeout=12.0, + ) + + class FakeRuntime(vf.Runtime): + def __init__(self) -> None: + self.writes: list[tuple[str, bytes]] = [] + self.runs: list[tuple[list[str], float | None]] = [] + + async def start(self) -> None: + return None + + async def stop(self) -> None: + return None + + async def expose(self, port: int) -> str: + return f"http://127.0.0.1:{port}" + + async def run( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> vf.CommandResult: + _ = cwd, env + self.runs.append((command, timeout)) + return vf.CommandResult(returncode=0) + + async def read(self, path: str) -> bytes: + assert path == "/logs/verifier/reward.txt" + return b"1" + + async def write(self, path: str, data: bytes) -> None: + self.writes.append((path, data)) + + runtime = FakeRuntime() + reward = await HarborTaskset(HarborTasksetConfig()).harbor_reward(task, runtime) + + assert reward == 1.0 + assert runtime.writes[0][0] == "/tmp/tests.tgz" + with tarfile.open(fileobj=io.BytesIO(runtime.writes[0][1]), mode="r:gz") as tar: + assert tar.getnames() == ["test.sh"] + assert runtime.runs == [ + ( + [ + "sh", + "-c", + "mkdir -p /logs/verifier /tests && tar -xzf /tmp/tests.tgz -C /tests", + ], + None, + ), + (["sh", "-c", "cd /tests && bash test.sh"], 12.0), + ] + + +@pytest.mark.asyncio +async def test_harbor_reward_returns_zero_when_tests_missing(tmp_path) -> None: + from tasksets.harbor import HarborTask, HarborTaskset, HarborTasksetConfig + + task_dir = tmp_path / "task" + task_dir.mkdir() + task = HarborTask( + task_name="task", + instruction="do it", + task_dir=str(task_dir), + prompt=[vf.UserMessage(content="do it")], + image="python:3.11", + ) + + class FakeRuntime(vf.Runtime): + async def start(self) -> None: + return None + + async def stop(self) -> None: + return None + + async def expose(self, port: int) -> str: + return f"http://127.0.0.1:{port}" + + async def run( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> vf.CommandResult: + _ = command, cwd, env, timeout + raise AssertionError("missing tests should not execute commands") + + async def read(self, path: str) -> bytes: + _ = path + raise AssertionError("missing tests should not read verifier output") + + async def write(self, path: str, data: bytes) -> None: + _ = path, data + raise AssertionError("missing tests should not write an archive") + + reward = await HarborTaskset(HarborTasksetConfig()).harbor_reward( + task, FakeRuntime() + ) + + assert reward == 0.0 + + +@pytest.mark.asyncio +async def test_score_group_empty_group_returns_empty_states() -> None: + env = vf.Env(taskset=ExactMatchTaskset()) + + assert await env.score_group([], []) == [] + + +@pytest.mark.asyncio +async def test_run_rollout_retries_copy_supplied_state() -> None: + class RetryHarness(vf.Harness): + attempts: int = 0 + + async def run_with_context(self, context: vf.Context) -> None: + self.attempts += 1 + context.state.transcript.append(vf.Turn(prompt=[], completion=[])) + if self.attempts == 1: + raise InfraError("retry") + + initial_state = vf.State() + env = vf.Env(taskset=ExactMatchTaskset(), harness=RetryHarness()) + + result = await env.run_rollout( + ExactMatchTask(prompt="hello", answer="ok"), + model=vf.ModelConfig(model="student"), + state=initial_state, + max_retries=1, + ) + + assert len(result.transcript) == 1 + assert initial_state.transcript == [] + + +@pytest.mark.asyncio +async def test_score_group_closes_model_client_if_teacher_load_fails( + mock_client, +) -> None: + env = vf.Env(taskset=ExactMatchTaskset()) + task = ExactMatchTask(prompt="hello", answer="ok") + state = vf.State(task_id=task.task_id) + model = vf.ModelConfig(model="student") + teacher = vf.ModelConfig(model="teacher") + closed: list[str] = [] + + def load_model_client(config: vf.ModelConfig) -> vf.ModelClient: + if config.model == "teacher": + raise RuntimeError("teacher failed") + return vf.ModelClient(config=config, client=mock_client) + + async def close_model_client(model_client: vf.ModelClient) -> None: + closed.append(model_client.config.model) + + env.harness.load_model_client = load_model_client + env.harness.close_model_client = close_model_client + + with pytest.raises(RuntimeError, match="teacher failed"): + await env.score_group([task], [state], model=model, teacher=teacher) + + assert closed == ["student"] + + +@pytest.mark.asyncio +async def test_score_group_closes_model_clients_if_cleanup_fails(mock_client) -> None: + env = vf.Env(taskset=ExactMatchTaskset()) + task = ExactMatchTask(prompt="hello", answer="ok") + state = vf.State(task_id=task.task_id) + model = vf.ModelConfig(model="student") + teacher = vf.ModelConfig(model="teacher") + closed: list[str] = [] + + def load_model_client(config: vf.ModelConfig) -> vf.ModelClient: + return vf.ModelClient(config=config, client=mock_client) + + async def close_model_client(model_client: vf.ModelClient) -> None: + closed.append(model_client.config.model) + + async def run_handlers_for_group( + kind: str, *args: object, **kwargs: object + ) -> None: + if kind == "cleanup": + raise RuntimeError("cleanup failed") + + env.harness.load_model_client = load_model_client + env.harness.close_model_client = close_model_client + env.run_handlers_for_group = run_handlers_for_group + + with pytest.raises(RuntimeError, match="cleanup failed"): + await env.score_group([task], [state], model=model, teacher=teacher) + + assert closed == ["teacher", "student"] + + +@pytest.mark.asyncio +async def test_openenv_and_openreward_rewards_sum_turn_rewards() -> None: + from tasksets.openenv import OpenEnvTaskset, OpenEnvTasksetConfig + + state = vf.State( + reward=99.0, + transcript=[ + vf.Turn(prompt=[], completion=[], reward=0.25), + vf.Turn(prompt=[], completion=[], reward=0.5), + ], + ) + + assert ( + await OpenEnvTaskset(OpenEnvTasksetConfig()).openenv_reward(state) + ) == pytest.approx(0.75) + pytest.importorskip("openreward") + from tasksets.openreward import OpenRewardTaskset, OpenRewardTasksetConfig + + assert ( + await OpenRewardTaskset( + OpenRewardTasksetConfig(environment="test") + ).openreward_reward(state) + ) == pytest.approx(0.75) + + +def test_openenv_build_config_accepts_current_build_metadata() -> None: + from tasksets.openenv import OpenEnvBuildConfig + + config = OpenEnvBuildConfig.model_validate( + { + "app": "server.app:app", + "contract": "mcp", + "environment_id": "openenv-echo", + "image": "owner/openenv-echo:latest", + "image_status": "COMPLETED", + "port": 8000, + "schema_version": 1, + "start_command": "uvicorn server.app:app", + "tools": [ + {"name": "echo", "description": "", "parameters": {"type": "object"}} + ], + } + ) + + assert config.image == "owner/openenv-echo:latest" + assert config.contract == "mcp" + assert config.tools == [ + {"name": "echo", "description": "", "parameters": {"type": "object"}} + ] + + +def test_openenv_and_openreward_task_schemas_are_explicit() -> None: + from tasksets.openenv import OpenEnvTask + + openenv_task = OpenEnvTask.model_validate( + { + "prompt": [], + "openenv": {"image": "owner/env:latest"}, + "info": {"seed": 0}, + } + ) + pytest.importorskip("openreward") + from tasksets.openreward import OpenRewardVFTask + + openreward_task = OpenRewardVFTask.model_validate( + { + "prompt": [], + "openreward": {"environment": "demo"}, + } + ) + + assert openenv_task.openenv == {"image": "owner/env:latest"} + assert openreward_task.openreward == {"environment": "demo"} + + +@pytest.mark.asyncio +async def test_openenv_mcp_setup_lists_tools_without_reset(monkeypatch) -> None: + import tasksets.openenv as openenv_module + from tasksets.openenv import OpenEnvUser, OpenEnvUserConfig + + calls: list[str] = [] + + class FakeOpenEnvSession: + def __init__(self, config: object, server: object) -> None: + self.config = config + self.server = server + + async def start(self) -> None: + calls.append("start") + + async def reset(self) -> None: + calls.append("reset") + + async def tool_defs(self) -> list[vf.JsonData]: + calls.append("tool_defs") + return [{"name": "echo", "description": "", "parameters": {}}] + + async def close(self) -> None: + calls.append("close") + + monkeypatch.setattr(openenv_module, "OpenEnvSession", FakeOpenEnvSession) + + config: vf.JsonData = { + "openenv_project": "proj", + "prompt_renderer": "x:y", + "image": "owner/env:latest", + "port": 8000, + "start_command": "serve", + "contract": "mcp", + "seed": 0, + "startup_timeout_seconds": 1, + "startup_poll_interval_seconds": 0.1, + "health_request_timeout_seconds": 1.0, + "schema_request_timeout_seconds": 1.0, + "wait_for_creation_max_attempts": 1, + "max_retries": 1, + "base_delay": 0.1, + "backoff_factor": 1.0, + "max_backoff_seconds": 1.0, + "jitter": 0.0, + } + + user = OpenEnvUser(OpenEnvUserConfig()) + user.start() + + payload = await user.setup("state-1", config) + + assert calls == ["tool_defs"] + assert payload["openenv_done"] is False + assert payload["tools"] == [{"name": "echo", "description": "", "parameters": {}}] + + calls.clear() + payload = await user.setup( + "state-2", + { + **config, + "tools": [ + {"name": "build_echo", "description": "", "parameters": {}}, + ], + }, + ) + + assert calls == [] + assert payload["tools"] == [ + {"name": "build_echo", "description": "", "parameters": {}}, + ] + + +def test_openenv_user_reuses_server_across_rollout_seeds() -> None: + from tasksets.openenv import OpenEnvRuntimeConfig, OpenEnvUser, OpenEnvUserConfig + + base: vf.JsonData = { + "openenv_project": "proj", + "prompt_renderer": "x:y", + "image": "owner/env:latest", + "port": 8000, + "start_command": "serve", + "contract": "mcp", + "tools": [], + "startup_timeout_seconds": 1, + "startup_poll_interval_seconds": 0.1, + "health_request_timeout_seconds": 1.0, + "schema_request_timeout_seconds": 1.0, + "wait_for_creation_max_attempts": 1, + "max_retries": 1, + "base_delay": 0.1, + "backoff_factor": 1.0, + "max_backoff_seconds": 1.0, + "jitter": 0.0, + } + config_a = OpenEnvRuntimeConfig.model_validate({**base, "seed": 1}) + config_b = OpenEnvRuntimeConfig.model_validate({**base, "seed": 2}) + + user = OpenEnvUser(OpenEnvUserConfig()) + user.start() + + assert user.server_for(config_a) is user.server_for(config_b) + + +@pytest.mark.asyncio +async def test_openenv_user_tool_returns_bound_turn_reward_payload() -> None: + from tasksets.openenv import OpenEnvUser, OpenEnvUserConfig + + class FakeOpenEnvObservation(vf.Config): + result: vf.JsonData + reward: float + done: bool + + class FakeOpenEnvResult: + observation = FakeOpenEnvObservation( + result={"data": {"echo": "ok"}}, + reward=1.25, + done=True, + ) + reward = None + done = False + + class FakeOpenEnvSession: + def __init__(self) -> None: + self.calls: list[tuple[str, vf.JsonData]] = [] + + async def call_tool(self, name: str, input: vf.JsonData) -> FakeOpenEnvResult: + self.calls.append((name, input)) + return FakeOpenEnvResult() + + session = FakeOpenEnvSession() + user = OpenEnvUser(OpenEnvUserConfig()) + user.start() + user.sessions["state-1"] = session + + payload = await user.call_tool("state-1", "echo", {"message": "hi"}) + + assert session.calls == [("echo", {"message": "hi"})] + assert json.loads(str(payload["content"])) == {"echo": "ok"} + assert payload["openenv_done"] is True + assert payload["reward"] == pytest.approx(1.25) + assert payload["finished"] is True + assert payload["stop_condition"] == "openenv_done" + + +@pytest.mark.asyncio +async def test_openreward_user_tool_returns_bound_turn_reward_payload() -> None: + pytest.importorskip("openreward") + from tasksets.openreward import OpenRewardUser, OpenRewardUserConfig + + class FakeOpenRewardOutput: + blocks = ["raw"] + reward = 0.5 + finished = True + + class FakeOpenRewardSession: + def __init__(self) -> None: + self.calls: list[tuple[str, vf.JsonData]] = [] + + async def call_tool( + self, name: str, input: vf.JsonData + ) -> FakeOpenRewardOutput: + self.calls.append((name, input)) + return FakeOpenRewardOutput() + + def content(self, blocks: object) -> str: + assert blocks == ["raw"] + return "judged" + + session = FakeOpenRewardSession() + user = OpenRewardUser(OpenRewardUserConfig()) + user.session = session + + payload = await user.call_tool("score", {"answer": "ok"}) + + assert session.calls == [("score", {"answer": "ok"})] + assert payload == { + "content": "judged", + "reward": 0.5, + "finished": True, + "stop_condition": "openreward_finished", + } + + +@pytest.mark.asyncio +async def test_v1_subprocess_runtime_session_read_write_run() -> None: + async with vf.make_runtime_provider( + vf.SubprocessRuntimeConfig() + ).create_runtime() as runtime: + await runtime.write("payload.txt", b"hello runtime") + result = await runtime.run( + [ + sys.executable, + "-c", + "from pathlib import Path; print(Path('payload.txt').read_text())", + ] + ) + payload = await runtime.read("payload.txt") + + assert result.returncode == 0 + assert result.stdout.strip() == "hello runtime" + assert payload == b"hello runtime" + + +@pytest.mark.asyncio +async def test_v1_docker_runtime_read_preserves_binary_bytes() -> None: + payload = b"\xff\x00binary" + + class ReadOnlyDockerRuntime(vf.DockerRuntime): + async def run( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> vf.CommandResult: + _ = cwd, env, timeout + assert command == ["sh", "-c", "base64 < /payload.bin"] + return vf.CommandResult( + returncode=0, + stdout=base64.b64encode(payload).decode(), + ) + + runtime = ReadOnlyDockerRuntime(vf.DockerRuntimeConfig()) + + assert await runtime.read("/payload.bin") == payload + + +@pytest.mark.asyncio +async def test_v1_prime_runtime_write_uses_gateway_upload() -> None: + class FakePrimeClient: + def __init__(self) -> None: + self.uploads: list[tuple[str, str, bytes, str]] = [] + + async def upload_bytes( + self, sandbox_id: str, path: str, data: bytes, *, filename: str + ) -> None: + self.uploads.append((sandbox_id, path, data, filename)) + + runtime = vf.PrimeRuntime(vf.PrimeRuntimeConfig(workdir="/app")) + client = FakePrimeClient() + runtime.client = client + runtime.sandbox_id = "sandbox-id" + commands: list[list[str]] = [] + + async def run( + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> vf.CommandResult: + _ = cwd, env, timeout + commands.append(command) + return vf.CommandResult(returncode=0) + + runtime.run = run + + await runtime.write("payload.bin", b"large payload") + + assert commands == [["sh", "-c", "mkdir -p /app"]] + assert client.uploads == [ + ("sandbox-id", "/app/payload.bin", b"large payload", "payload.bin") + ] + + +@pytest.mark.asyncio +async def test_v1_prime_runtime_public_url_exposes_sandbox_port() -> None: + class ExposedPort: + url = "https://sandbox.example/mcp/" + + class FakePrimeClient: + async def expose(self, sandbox_id: str, port: int) -> ExposedPort: + assert sandbox_id == "sandbox-id" + assert port == 8765 + return ExposedPort() + + runtime = vf.PrimeRuntime(vf.PrimeRuntimeConfig()) + runtime.client = FakePrimeClient() + runtime.sandbox_id = "sandbox-id" + + assert await runtime.public_url(8765) == "https://sandbox.example/mcp" + + +@pytest.mark.asyncio +async def test_v1_prime_runtime_run_honors_timeout() -> None: + class FakePrimeClient: + async def run_background_job( + self, + sandbox_id: str, + command: str, + *, + working_dir: str, + env: dict[str, str], + ) -> None: + _ = sandbox_id, command, working_dir, env + await asyncio.sleep(1) + + runtime = vf.PrimeRuntime(vf.PrimeRuntimeConfig()) + runtime.client = FakePrimeClient() + runtime.sandbox_id = "sandbox-id" + + with pytest.raises(TimeoutError): + await runtime.run(["sleep", "10"], timeout=0.001) + + +@pytest.mark.asyncio +@pytest.mark.prime_sandbox +async def test_v1_prime_runtime_live_run_read_write() -> None: + if not os.environ.get("PRIME_API_KEY") or not os.environ.get("PRIME_TEAM_ID"): + pytest.skip("Prime sandbox credentials are required.") + + runtime = vf.PrimeRuntime( + vf.PrimeRuntimeConfig( + idle_timeout_minutes=5, + labels=["ci"], + ) + ) + try: + await runtime.start() + result = await runtime.run(["python", "--version"]) + assert result.returncode == 0 + assert "Python" in f"{result.stdout}\n{result.stderr}" + + await runtime.write("probe.txt", b"prime runtime ok") + assert await runtime.read("/app/probe.txt") == b"prime runtime ok" + finally: + await runtime.stop() + + +@pytest.mark.asyncio +async def test_v1_prime_runtime_start_labels_child_sandbox(monkeypatch) -> None: + prime_sandboxes = ModuleType("prime_sandboxes") + requests: list[dict[str, object]] = [] + + class Sandbox: + id = "sandbox-id" + + class FakeRequest: + def __init__(self, **kwargs: object) -> None: + requests.append(kwargs) + + class FakeAdvancedConfigs: + @classmethod + def model_validate(cls, data: object) -> "FakeAdvancedConfigs": + if not isinstance(data, dict): + raise TypeError + return cls(**data) + + def __init__(self, **kwargs: object) -> None: + self.data = kwargs + + class FakePrimeClient: + def __init__(self) -> None: + self.deleted: list[str] = [] + self.closed = False + + async def create(self, request: FakeRequest) -> Sandbox: + _ = request + return Sandbox() + + async def wait_for_creation(self, sandbox_id: str) -> None: + assert sandbox_id == "sandbox-id" + + async def run_background_job(self, sandbox_id: str, command: str) -> None: + assert sandbox_id == "sandbox-id" + assert command == "mkdir -p /app" + + async def delete(self, sandbox_id: str) -> None: + self.deleted.append(sandbox_id) + + async def aclose(self) -> None: + self.closed = True + + prime_sandboxes.AsyncSandboxClient = FakePrimeClient + prime_sandboxes.AdvancedConfigs = FakeAdvancedConfigs + prime_sandboxes.CreateSandboxRequest = FakeRequest + monkeypatch.setitem(sys.modules, "prime_sandboxes", prime_sandboxes) + monkeypatch.setenv("EVALUATION_ID", "eval-id") + monkeypatch.setenv("PRIME_JOB_ID", "job-id") + + runtime = vf.PrimeRuntime( + vf.PrimeRuntimeConfig( + idle_timeout_minutes=7, + labels=["custom", "vf-v1-runtime"], + ) + ) + + await runtime.start() + + assert requests == [ + { + "name": "vf-v1-runtime", + "docker_image": "python:3.11-slim", + "cpu_cores": 1.0, + "memory_gb": 2.0, + "disk_size_gb": 5.0, + "gpu_count": 0, + "timeout_minutes": 360, + "network_access": True, + "vm": False, + "guaranteed": False, + "gpu_type": None, + "region": None, + "advanced_configs": requests[0]["advanced_configs"], + "labels": [ + "vf-v1-runtime", + "custom", + "eval-eval-id", + "prime-job-job-id", + ], + } + ] + advanced_configs = requests[0]["advanced_configs"] + assert isinstance(advanced_configs, FakeAdvancedConfigs) + assert advanced_configs.data == {"idle_timeout_minutes": 7} + + client = runtime.client + await runtime.stop() + + assert client.deleted == ["sandbox-id"] + assert client.closed is True + + +@pytest.mark.asyncio +async def test_v1_prime_runtime_start_failure_deletes_child_sandbox( + monkeypatch, +) -> None: + prime_sandboxes = ModuleType("prime_sandboxes") + clients: list[FakePrimeClient] = [] + + class Sandbox: + id = "sandbox-id" + + class FakeRequest: + def __init__(self, **kwargs: object) -> None: + self.kwargs = kwargs + + class FakeAdvancedConfigs: + @classmethod + def model_validate(cls, data: object) -> "FakeAdvancedConfigs": + if not isinstance(data, dict): + raise TypeError + return cls(**data) + + def __init__(self, **kwargs: object) -> None: + self.data = kwargs + + class FakePrimeClient: + def __init__(self) -> None: + self.deleted: list[str] = [] + self.closed = False + clients.append(self) + + async def create(self, request: FakeRequest) -> Sandbox: + _ = request + return Sandbox() + + async def wait_for_creation(self, sandbox_id: str) -> None: + assert sandbox_id == "sandbox-id" + raise RuntimeError("creation failed") + + async def delete(self, sandbox_id: str) -> None: + self.deleted.append(sandbox_id) + + async def aclose(self) -> None: + self.closed = True + + prime_sandboxes.AsyncSandboxClient = FakePrimeClient + prime_sandboxes.AdvancedConfigs = FakeAdvancedConfigs + prime_sandboxes.CreateSandboxRequest = FakeRequest + monkeypatch.setitem(sys.modules, "prime_sandboxes", prime_sandboxes) + + runtime = vf.PrimeRuntime(vf.PrimeRuntimeConfig()) + + with pytest.raises(RuntimeError, match="creation failed"): + await runtime.start() + + assert runtime.client is None + assert runtime.sandbox_id is None + assert clients[0].deleted == ["sandbox-id"] + assert clients[0].closed is True + + +def test_v1_runtime_image_is_container_only() -> None: + task = vf.Task({"prompt": [], "image": "python:3.12-slim"}) + subprocess_env = vf.Env( + taskset=ExactMatchTaskset(), runtime=vf.SubprocessRuntimeConfig() + ) + docker_env = vf.Env( + taskset=ExactMatchTaskset(), + runtime=vf.DockerRuntimeConfig(image="python:3.11-slim"), + ) + + with pytest.raises(ValueError, match="declares an image"): + subprocess_env.harness.runtime_for(task) + docker_config = docker_env.harness.runtime_for(task) + + assert isinstance(docker_config, vf.DockerRuntimeConfig) + assert docker_config.image == "python:3.12-slim" + + +def test_v1_runtime_config_applies_task_resources_with_config_precedence() -> None: + task = vf.Task( + { + "prompt": [], + "resources": { + "cpu_cores": 2.0, + "memory_gb": 4.0, + "gpu_count": 1, + "disk_gb": 12.0, + }, + } + ) + default_env = vf.Env(taskset=ExactMatchTaskset(), runtime=vf.DockerRuntimeConfig()) + configured_env = vf.Env( + taskset=ExactMatchTaskset(), + runtime=vf.DockerRuntimeConfig(cpu_cores=8.0, memory_gb=16.0), + ) + prime_env = vf.Env(taskset=ExactMatchTaskset(), runtime=vf.PrimeRuntimeConfig()) + subprocess_env = vf.Env( + taskset=ExactMatchTaskset(), runtime=vf.SubprocessRuntimeConfig() + ) + + docker_config = default_env.harness.runtime_for(task) + configured_config = configured_env.harness.runtime_for(task) + prime_config = prime_env.harness.runtime_for(task) + + assert isinstance(docker_config, vf.DockerRuntimeConfig) + assert docker_config.cpu_cores == 2.0 + assert docker_config.memory_gb == 4.0 + assert docker_config.gpu_count == 1 + assert docker_config.disk_gb == 12.0 + + assert isinstance(configured_config, vf.DockerRuntimeConfig) + assert configured_config.cpu_cores == 8.0 + assert configured_config.memory_gb == 16.0 + assert configured_config.gpu_count == 1 + assert configured_config.disk_gb == 12.0 + + assert isinstance(prime_config, vf.PrimeRuntimeConfig) + assert prime_config.cpu_cores == 2.0 + assert prime_config.memory_gb == 4.0 + assert prime_config.gpu_count == 1 + assert prime_config.disk_gb == 12.0 + + with pytest.raises(ValueError, match="does not support"): + subprocess_env.harness.runtime_for(task) + + +def test_v1_default_advantages_fill_turn_tokens() -> None: + states = [ + vf.State( + transcript=[ + vf.Turn( + prompt=[{"role": "user", "content": "p"}], + completion=[{"role": "assistant", "content": "a"}], + tokens=vf.TurnTokens( + prompt_ids=[1, 2], + prompt_mask=[1, 1], + completion_ids=[3, 4, 5], + completion_mask=[1, 1, 1], + completion_logprobs=[0.1, 0.2, 0.3], + ), + ) + ], + reward=2.0, + ), + vf.State( + transcript=[ + vf.Turn( + prompt=[{"role": "user", "content": "p"}], + completion=[{"role": "assistant", "content": "b"}], + tokens=vf.TurnTokens( + prompt_ids=[6], + prompt_mask=[1], + completion_ids=[7, 8], + completion_mask=[1, 1], + completion_logprobs=[0.4, 0.5], + ), + ) + ], + reward=0.0, + ), + ] + + tasks = [vf.Task(prompt="p"), vf.Task(prompt="p")] + vf.advantages.grpo(tasks, states) + + assert states[0].transcript[0].tokens is not None + assert states[0].transcript[0].tokens.prompt_advantages == pytest.approx([0.0, 0.0]) + assert states[0].transcript[0].tokens.completion_advantages == pytest.approx( + [1.0, 1.0, 1.0] + ) + + vf.advantages.rl(tasks, states) + assert states[0].transcript[0].tokens.completion_advantages == pytest.approx( + [1.0, 1.0, 1.0] + ) + + vf.advantages.sft(tasks, states) + assert states[0].transcript[0].tokens.prompt_advantages == pytest.approx([1.0, 1.0]) + assert states[0].transcript[0].tokens.completion_advantages == pytest.approx( + [1.0, 1.0, 1.0] + ) + + +def test_env_advantage_defaults_to_rl() -> None: + env = vf.Env(taskset=ExactMatchTaskset()) + + assert env.advantage == "rl" + assert env.provides_advantages + assert env.requires_group_rollouts + + +@pytest.mark.asyncio +async def test_env_advantage_config_sets_group_default() -> None: + tasks = [vf.Task(prompt="p"), vf.Task(prompt="p")] + states = [ + vf.State( + transcript=[ + vf.Turn( + prompt=[{"role": "user", "content": "p"}], + completion=[{"role": "assistant", "content": "a"}], + tokens=vf.TurnTokens( + prompt_ids=[1], + prompt_mask=[1], + completion_ids=[2, 3], + completion_mask=[1, 1], + completion_logprobs=[0.1, 0.2], + ), + ) + ], + reward=2.0, + ), + vf.State( + transcript=[ + vf.Turn( + prompt=[{"role": "user", "content": "p"}], + completion=[{"role": "assistant", "content": "b"}], + tokens=vf.TurnTokens( + prompt_ids=[4], + prompt_mask=[1], + completion_ids=[5], + completion_mask=[1], + completion_logprobs=[0.3], + ), + ) + ], + reward=0.0, + ), + ] + env = vf.Env( + taskset=ExactMatchTaskset(), + advantage="reinforce", + ) + + await env.score_group(tasks, states) + + assert states[0].transcript[0].tokens is not None + assert states[0].transcript[0].tokens.completion_advantages == pytest.approx( + [2.0, 2.0] + ) + assert states[1].transcript[0].tokens is not None + assert states[1].transcript[0].tokens.completion_advantages == pytest.approx([0.0]) + + +@pytest.mark.asyncio +async def test_env_advantage_path_supports_user_authored_group_logic() -> None: + tasks = [vf.Task(prompt="p"), vf.Task(prompt="p")] + states = [ + vf.State( + transcript=[ + vf.Turn( + prompt=[{"role": "user", "content": "p"}], + completion=[{"role": "assistant", "content": "a"}], + tokens=vf.TurnTokens( + prompt_ids=[1], + prompt_mask=[1], + completion_ids=[2], + completion_mask=[1], + completion_logprobs=[0.1], + ), + ) + ], + reward=2.0, + ), + vf.State( + transcript=[ + vf.Turn( + prompt=[{"role": "user", "content": "p"}], + completion=[{"role": "assistant", "content": "b"}], + tokens=vf.TurnTokens( + prompt_ids=[3], + prompt_mask=[1], + completion_ids=[4], + completion_mask=[1], + completion_logprobs=[0.2], + ), + ) + ], + reward=0.0, + ), + ] + env = vf.Env( + taskset=ExactMatchTaskset(), advantage=f"{__name__}:custom_env_advantage" + ) + + await env.score_group(tasks, states) + + assert states[0].transcript[0].tokens is not None + assert states[0].transcript[0].tokens.completion_advantages == pytest.approx([10.0]) + assert states[1].transcript[0].tokens is not None + assert states[1].transcript[0].tokens.completion_advantages == pytest.approx([11.0]) + + +def test_v1_runtime_default_resolves_taskset_and_harness_fields() -> None: + class RuntimeTasksetConfig(vf.TasksetConfig): + runtime: vf.RuntimeConfig | None = vf.DockerRuntimeConfig( + image="python:3.12-slim" + ) + + class RuntimeTaskset(vf.Taskset[RuntimeTasksetConfig]): + config: RuntimeTasksetConfig + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return ExactMatchTaskset().load_tasks(split) + + class RuntimeHarnessConfig(vf.HarnessConfig): + runtime: vf.RuntimeConfig | None = vf.DockerRuntimeConfig(workdir="/workspace") + + class RuntimeHarness(vf.Harness[RuntimeHarnessConfig]): + pass + + env = vf.Env(taskset=RuntimeTaskset(), harness=RuntimeHarness()) + + assert isinstance(env.runtime_config, vf.DockerRuntimeConfig) + assert env.runtime_config.image == "python:3.12-slim" + assert env.runtime_config.workdir == "/workspace" + + +def test_v1_runtime_default_rejects_provider_conflicts() -> None: + class RuntimeTasksetConfig(vf.TasksetConfig): + runtime: vf.RuntimeConfig | None = vf.DockerRuntimeConfig() + + class RuntimeTaskset(vf.Taskset[RuntimeTasksetConfig]): + config: RuntimeTasksetConfig + + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return ExactMatchTaskset().load_tasks(split) + + class RuntimeHarnessConfig(vf.HarnessConfig): + runtime: vf.RuntimeConfig | None = vf.PrimeRuntimeConfig() + + class RuntimeHarness(vf.Harness[RuntimeHarnessConfig]): + pass + + with pytest.raises(ValueError, match="single provider type"): + vf.Env(taskset=RuntimeTaskset(), harness=RuntimeHarness()) + + +@pytest.mark.asyncio +async def test_v1_interception_supports_custom_protocols(mock_client) -> None: + class CustomProtocol(vf.EndpointProtocol): + name = "custom_json" + routes = (vf.ProtocolRoute("POST", "/custom/generate"),) + + def env(self, *, base_url: str, api_key: str, model: str) -> dict[str, str]: + return { + "CUSTOM_BASE_URL": base_url, + "CUSTOM_API_KEY": api_key, + "CUSTOM_MODEL": model, + } + + async def parse( + self, request: web.Request, body: vf.JsonData + ) -> vf.InterceptedRequest: + _ = request + prompt = body.get("input") + if not isinstance(prompt, str): + raise TypeError("input must be a string.") + sampling_args: dict[str, vf.JsonValue] = {} + temperature = body.get("temperature") + if isinstance(temperature, int | float) and not isinstance( + temperature, bool + ): + sampling_args["temperature"] = temperature + return vf.InterceptedRequest( + protocol=self.name, + prompt=[vf.UserMessage(content=prompt)], + model=body.get("model") if isinstance(body.get("model"), str) else None, + sampling_args=sampling_args, + body=body, + ) + + def serialize( + self, response: Response, request: vf.InterceptedRequest + ) -> vf.JsonData: + content = response.message.content + if not isinstance(content, str): + content = "" + return { + "protocol": request.protocol, + "text": content, + "model": response.model or request.model or "", + } + + task = vf.Task(prompt=[]) + state = vf.State(task_id=task.task_id) + model = vf.ModelConfig(client=ClientConfig(), model="fallback-model") + ctx = vf.Context( + task=task, + state=state, + model_client=vf.ModelClient(config=model, client=mock_client), + ) + + async with vf.InterceptionServer( + ctx, task, state, protocols=[CustomProtocol()] + ) as server: + url = f"http://127.0.0.1:{server.port}/custom/generate" + async with ClientSession() as session: + response = await session.post( + url, + headers={"Authorization": f"Bearer {server.secret}"}, + json={"input": "hello protocol", "model": "custom-model"}, + ) + payload = await response.json() + + assert response.status == 200 + assert payload == { + "protocol": "custom_json", + "text": "This is a test response", + "model": "test-model", + } + assert mock_client.last_call_kwargs["model"] == "custom-model" + assert mock_client.last_call_kwargs["prompt"][0].content == "hello protocol" + assert state.transcript[0].prompt[0].content == "hello protocol" + assert state.completion[-1].content == "This is a test response" diff --git a/tests/test_v1_empty_completions.py b/tests/test_v1_empty_completions.py deleted file mode 100644 index ebed809e73..0000000000 --- a/tests/test_v1_empty_completions.py +++ /dev/null @@ -1,57 +0,0 @@ -import importlib.util -from pathlib import Path -from types import ModuleType - -import pytest - - -def load_env_module(name: str, filename: str) -> ModuleType: - module_path = Path(__file__).parents[1] / "environments" / name / filename - spec = importlib.util.spec_from_file_location(f"test_{name}", module_path) - assert spec is not None - assert spec.loader is not None - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module - - -def test_dspy_rlm_empty_completion_scores_zero() -> None: - module = load_env_module("dspy_rlm", "dspy_rlm.py") - - assert module.answer_reward({"answer": "4"}, {"completion": []}) == 0.0 - - -def test_openai_agents_empty_completion_scores_zero() -> None: - module = load_env_module("openai_agents_env", "openai_agents_env.py") - - assert module.answer_reward({"answer": "4"}, {"completion": []}) == 0.0 - - -@pytest.mark.asyncio -async def test_math_python_empty_completion_scores_zero() -> None: - module = load_env_module("math_python", "math_python_v1.py") - - assert await module.correct_answer({"answer": "4"}, {"completion": []}) == 0.0 - - -@pytest.mark.asyncio -async def test_hello_subagent_missing_completion_scores_zero() -> None: - module = load_env_module("hello_subagent_v1", "hello_subagent_v1.py") - - assert ( - await module.exact_answer({"answer": "hello alice"}, {"completion": None}) - == 0.0 - ) - - -def test_hello_parallel_reward_prompt_allows_missing_completion() -> None: - module = load_env_module( - "hello_parallel_sandbox_v1", "hello_parallel_sandbox_v1.py" - ) - - prompt = module.reward_prompt( - {"instruction": "write an answer", "answer": "done"}, - {"completion": None}, - ) - - assert "Assistant final answer:\n\n" in prompt diff --git a/tests/test_v1_endpoint_protocols.py b/tests/test_v1_endpoint_protocols.py deleted file mode 100644 index 3f2682e4b5..0000000000 --- a/tests/test_v1_endpoint_protocols.py +++ /dev/null @@ -1,224 +0,0 @@ -from anthropic import Anthropic, AsyncAnthropic -from openai import AsyncOpenAI, OpenAI -import pytest - -from verifiers.clients import AnthropicMessagesClient, OpenAIResponsesClient -from verifiers.types import ClientConfig, Response, ResponseMessage, ToolCall -from verifiers.utils.interception_utils import serialize_intercept_response -from verifiers.v1.runtime import Runtime -from verifiers.v1.state import State -from verifiers.v1.utils.endpoint_utils import ( - Endpoint, - normalize_endpoint_api, - normalize_endpoint_prompt, -) - - -def test_runtime_records_client_config_protocol(): - runtime = Runtime() - state = State({"runtime": {}}) - - runtime.bind_model_client( - state, - ClientConfig( - client_type="anthropic_messages", - api_base_url="https://api.anthropic.com", - api_key_var="ANTHROPIC_API_KEY", - ), - ) - - assert state["runtime"]["client_type"] == "anthropic_messages" - assert isinstance(runtime.model_client(state), AnthropicMessagesClient) - - -def test_runtime_preserves_concrete_client_config_protocol(): - runtime = Runtime() - state = State({"runtime": {}}) - client = OpenAIResponsesClient( - ClientConfig(client_type="openai_responses", api_key_var="OPENAI_API_KEY") - ) - - runtime.bind_model_client(state, client) - - assert state["runtime"]["client_type"] == "openai_responses" - assert runtime.model_client(state) is client - - -def test_endpoint_client_protocol_accepts_explicit_api_surface(): - endpoint = Endpoint(port=9999) - root = "http://127.0.0.1:9999/rollout/test" - state = State( - { - "runtime": {"client_type": "openai_responses"}, - "endpoint_root_url": root, - "endpoint_base_url": f"{root}/v1", - } - ) - - openai = endpoint.client(state) - openai_sync = endpoint.client(state, api="chat", sync=True) - completions = endpoint.client(state, api="openai_completions") - responses = endpoint.client(state, api="openai_responses") - anthropic = endpoint.client(state, api="anthropic_messages") - anthropic_sync = endpoint.client(state, api="messages", sync=True) - - assert isinstance(openai, AsyncOpenAI) - assert isinstance(openai_sync, OpenAI) - assert isinstance(completions, AsyncOpenAI) - assert isinstance(responses, AsyncOpenAI) - assert isinstance(anthropic, AsyncAnthropic) - assert isinstance(anthropic_sync, Anthropic) - - -def test_endpoint_config_uses_endpoint_client_type_names(): - endpoint = Endpoint(port=9999, secret="test-secret") - root = "http://127.0.0.1:9999/rollout/test" - state = State( - { - "runtime": {"model": "test-model"}, - "endpoint_root_url": root, - "endpoint_base_url": f"{root}/v1", - "endpoint_api_key_var": "VF_ENDPOINT_API_KEY_ROLLOUT_TEST", - } - ) - - openai_config = endpoint.config(state, api="responses") - anthropic_config = endpoint.config(state, api="messages") - - assert openai_config.model_dump(exclude_none=True) == { - "model": "test-model", - "base_url": f"{root}/v1", - "api_key_var": "VF_ENDPOINT_API_KEY_ROLLOUT_TEST", - "api_client_type": "openai_responses", - "extra_headers": {}, - } - assert anthropic_config.model_dump(exclude_none=True) == { - "model": "test-model", - "base_url": root, - "api_key_var": "VF_ENDPOINT_API_KEY_ROLLOUT_TEST", - "api_client_type": "anthropic_messages", - "extra_headers": {}, - } - - -@pytest.mark.parametrize( - ("alias", "api"), - [ - ("chat", "chat_completions"), - ("chat_completions", "chat_completions"), - ("openai_chat_completions", "chat_completions"), - ("completions", "completions"), - ("openai_completions", "completions"), - ("responses", "responses"), - ("openai_responses", "responses"), - ("messages", "messages"), - ("anthropic_messages", "messages"), - ], -) -def test_endpoint_api_aliases_match_endpoint_config_type(alias, api): - assert normalize_endpoint_api(alias) == api - - -@pytest.mark.parametrize( - "api", - [ - "openai", - "anthropic", - "completion", - "openai_chat_completions_token", - "renderer", - "nemorl_chat_completions", - ], -) -def test_endpoint_api_rejects_unsupported_client_types(api): - with pytest.raises(ValueError): - normalize_endpoint_api(api) - - -def test_openai_completions_endpoint_prompt_normalizes_text_prompt(): - messages = normalize_endpoint_prompt( - {"protocol": "openai_completions", "prompt": "Complete this sentence"} - ) - - assert messages[0].role == "text" - assert messages[0].content == "Complete this sentence" - - -def test_anthropic_endpoint_prompt_normalizes_tool_messages(): - messages = normalize_endpoint_prompt( - { - "protocol": "anthropic_messages", - "system": "system text", - "messages": [ - {"role": "user", "content": "question"}, - { - "role": "assistant", - "content": [ - { - "type": "tool_use", - "id": "call_1", - "name": "search", - "input": {"query": "x"}, - } - ], - }, - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "call_1", - "content": "answer", - } - ], - }, - ], - } - ) - - assert messages[0]["role"] == "system" - assert messages[1]["role"] == "user" - assert messages[2]["tool_calls"][0]["name"] == "search" - assert messages[3]["role"] == "tool" - - -def test_openai_responses_serialization_includes_function_calls(): - response = Response( - id="resp_1", - created=123, - model="m", - usage=None, - message=ResponseMessage( - content=None, - finish_reason="tool_calls", - is_truncated=False, - tool_calls=[ - ToolCall(id="call_1", name="search", arguments='{"query": "x"}') - ], - ), - ) - - payload = serialize_intercept_response(response, protocol="openai_responses") - - assert payload["object"] == "response" - assert payload["output"][0]["type"] == "function_call" - assert payload["output"][0]["call_id"] == "call_1" - - -def test_openai_completions_serialization_returns_text_completion_shape(): - response = Response( - id="cmpl_1", - created=123, - model="m", - usage=None, - message=ResponseMessage( - content="done", - finish_reason="stop", - is_truncated=False, - ), - ) - - payload = serialize_intercept_response(response, protocol="openai_completions") - - assert payload["object"] == "text_completion" - assert payload["choices"][0]["text"] == "done" diff --git a/tests/test_v1_example_counts.py b/tests/test_v1_example_counts.py deleted file mode 100644 index 097e1af648..0000000000 --- a/tests/test_v1_example_counts.py +++ /dev/null @@ -1,115 +0,0 @@ -import importlib -from collections.abc import Iterable, Mapping -from pathlib import Path -from typing import Any - -try: - import tomllib -except ModuleNotFoundError: - import tomli as tomllib - - -REPO_ROOT = Path(__file__).resolve().parents[1] - - -STATIC_SOURCES = [ - ( - "environments.mcp_search_env.mcp_search_env", - "load_tasks", - "question", - ), - ( - "environments.hello_subagent_v1.hello_subagent_v1", - "load_tasks", - "prompt", - ), - ( - "environments.nested_harness_v1.nested_harness_v1", - "load_tasks", - "prompt", - ), - ( - "environments.hello_rlm_v1.hello_rlm_v1", - "load_tasks", - "question", - ), - ( - "environments.hello_parallel_sandbox_v1.hello_parallel_sandbox_v1", - "load_tasks", - "instruction", - ), - ( - "environments.hello_group_reward_v1.hello_group_reward_v1", - "load_tasks", - "question", - ), - ( - "environments.hello_self_judge_v1.hello_self_judge_v1", - "load_tasks", - "question", - ), - ( - "environments.dspy_flights.dspy_flights", - "load_tasks", - "user_request", - ), -] - -DEFAULT_EVAL_NUM_EXAMPLES = 5 -DEFAULT_EVAL_ROLLOUTS_PER_EXAMPLE = 3 - - -def test_static_v1_example_sources_have_at_least_ten_unique_problems() -> None: - for module_name, loader_name, key in STATIC_SOURCES: - module = importlib.import_module(module_name) - rows = list(getattr(module, loader_name)()) - problems = {problem_text(row, key) for row in rows} - - assert len(rows) >= 10, module_name - assert len(problems) >= 10, module_name - - -def test_mcp_search_env_bundles_at_least_ten_self_contained_records() -> None: - module = importlib.import_module("environments.mcp_search_env.mcp_server") - records = module.RECORDS - - assert len(records) >= 10 - for record in records.values(): - assert record["title"] - assert record["summary"] - - -def test_environment_eval_configs_use_shared_smoke_defaults() -> None: - pyprojects = sorted((REPO_ROOT / "environments").glob("*/pyproject.toml")) - assert pyprojects - - for pyproject in pyprojects: - config = tomllib.loads(pyproject.read_text()) - eval_config = config["tool"]["verifiers"]["eval"] - - env_name = pyproject.parent.name - assert eval_config["num_examples"] == DEFAULT_EVAL_NUM_EXAMPLES, env_name - assert ( - eval_config["rollouts_per_example"] == DEFAULT_EVAL_ROLLOUTS_PER_EXAMPLE - ), env_name - - -def problem_text(row: Mapping[str, Any], key: str) -> str: - value = row[key] - if key == "prompt": - return prompt_text(value) - return str(value) - - -def prompt_text(prompt: object) -> str: - if isinstance(prompt, str): - return prompt - if isinstance(prompt, Iterable): - parts = [] - for item in prompt: - if isinstance(item, Mapping): - parts.append(str(item.get("content", ""))) - else: - parts.append(str(item)) - return "\n".join(parts) - return str(prompt) diff --git a/tests/test_v1_group_reward_env.py b/tests/test_v1_group_reward_env.py index 3076274354..47a5fcf31a 100644 --- a/tests/test_v1_group_reward_env.py +++ b/tests/test_v1_group_reward_env.py @@ -1,40 +1,55 @@ +import asyncio from typing import cast import pytest +import verifiers.v1 as vf from verifiers.clients import Client +from verifiers.types import ClientConfig from verifiers.types import RolloutInput -from environments.hello_group_reward_v1.hello_group_reward_v1 import ( - GroupRewardEnvConfig, - load_environment, -) +from environments.hello_group_reward_v1.hello_group_reward_v1 import taskset as module +from verifiers.v1.loaders import load_environment_from_components @pytest.mark.asyncio async def test_hello_group_reward_v1_scores_full_group_lifecycle() -> None: - env = load_environment(config=GroupRewardEnvConfig(taskset={"num_examples": 1})) + env = load_environment_from_components( + module, {"config": {"taskset": {"num_examples": 1}, "advantage": "grpo"}} + ) assert env.requires_group_rollouts assert env.provides_advantages row = cast(RolloutInput, env.taskset.get_dataset()[0]) - states = await env._run_group_states( - [row, row, row, row], - cast(Client, object()), - "unused-model", - {}, + model = vf.ModelConfig(client=ClientConfig(), model="unused-model") + env.harness.load_model_client = lambda _: vf.ModelClient( + config=model, client=cast(Client, object()) + ) + + async def close_model_client(_: vf.ModelClient) -> None: + return None + + env.harness.close_model_client = close_model_client + base_task = env.taskset.to_task(row) + tasks, states = await env.taskset.init_group(base_task, 4) + states = list( + await asyncio.gather( + *[ + env.run_rollout(task, model=model, state=state) + for task, state in zip(tasks, states, strict=True) + ] + ) ) + states = await env.score_group(tasks, states) assert len(states) == 4 - by_candidate = {state["candidate_id"]: state for state in states} + by_candidate = {state.extras["candidate_id"]: state for state in states} exact = by_candidate["exact"] off_topic = by_candidate["off-topic"] - assert exact["group_summary"]["rank"] == 1 - assert exact["group_summary"]["best_candidate_id"] == "exact" - assert exact["metrics"]["relative_group_reward"] == 1.0 - assert off_topic["metrics"]["relative_group_reward"] == 0.0 - assert exact["reward"] > off_topic["reward"] - assert all(state["group_cleaned"] is True for state in states) - assert sum(float(state["advantage"]) for state in states) == pytest.approx(0.0) - assert all("runtime_id" not in state.get("runtime", {}) for state in states) + assert exact.metrics["group_rank"] == 1.0 + assert exact.metrics["relative_group_reward"] == 1.0 + assert off_topic.metrics["relative_group_reward"] == 0.0 + assert exact.reward > off_topic.reward + assert all(state.transcript for state in states) + assert all("runtime_id" not in state.metadata for state in states) diff --git a/tests/test_v1_harbor_cli.py b/tests/test_v1_harbor_cli.py deleted file mode 100644 index 348c193c94..0000000000 --- a/tests/test_v1_harbor_cli.py +++ /dev/null @@ -1,555 +0,0 @@ -import importlib -import sys -import types -from pathlib import Path -from types import ModuleType -from typing import Any, cast -from uuid import uuid4 - -import pytest - -import verifiers as vf -from harnesses import ( - MiniSWEAgent, - MiniSWEAgentConfig, - MiniSWEAgentProgramConfig, - OpenCode, - OpenCodeConfig, - OpenCodeProgramConfig, - Pi, - PiConfig, - PiProgramConfig, - RLM, - RLMConfig, - RLMProgramConfig, - Terminus2Config, - Terminus2ProgramConfig, -) -from harnesses.pi import PI_DEFAULT_VERSION -from harnesses.terminus_2 import ( - TERMINUS_2_DEFAULT_API_BASE_URL, - TERMINUS_2_DEFAULT_VERSION, - TERMINUS_2_DEFAULT_MODEL_NAME, - Terminus2, -) -from tasksets import HarborTaskset, HarborTasksetConfig -from verifiers.v1.utils.program_utils import merge_task_program, merge_task_sandbox -from verifiers.v1.utils.sandbox_python_utils import SANDBOX_PYTHON - - -def write_harbor_task(root: Path, name: str = "task-a") -> Path: - task_dir = root / name - (task_dir / "tests").mkdir(parents=True) - (task_dir / "solution").mkdir() - (task_dir / "instruction.md").write_text("Write hello to /app/hello.txt\n") - (task_dir / "task.toml").write_text( - """ -version = "1.0" - -[environment] -docker_image = "ubuntu:24.04" -cpus = 1 -memory = "2G" -storage = "8G" - -[agent] -timeout_sec = 600 - -[verifier] -timeout_sec = 300 -""".strip() - ) - (task_dir / "tests" / "test.sh").write_text("echo 1 > /logs/verifier/reward.txt") - (task_dir / "solution" / "solve.sh").write_text("echo hello > /app/hello.txt") - return task_dir - - -def write_harbor_package(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> ModuleType: - package_name = f"harbor_pkg_{uuid4().hex}" - package_dir = tmp_path / package_name - tasks_root = package_dir / "tasks" - tasks_root.mkdir(parents=True) - (package_dir / "__init__.py").write_text( - """ -import verifiers as vf -from harnesses import OpenCode, OpenCodeConfig -from tasksets import HarborTaskset, HarborTasksetConfig - - -def load_taskset(config: HarborTasksetConfig): - if config.bundle_package is None: - config = config.model_copy(update={"bundle_package": __name__}) - return HarborTaskset(config=config) - - -def load_env(): - return vf.Env(taskset=HarborTaskset(config=HarborTasksetConfig(bundle_package=__name__)), harness=OpenCode(config=OpenCodeConfig())) -""".lstrip() - ) - monkeypatch.syspath_prepend(str(tmp_path)) - importlib.invalidate_caches() - module = importlib.import_module(package_name) - setattr(module, "tasks_root", tasks_root) - return module - - -def test_harbor_taskset_loads_package_tasks_with_program_patch( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - package = write_harbor_package(tmp_path, monkeypatch) - write_harbor_task(cast(Path, getattr(package, "tasks_root"))) - - taskset = getattr(package, "load_taskset")(config=HarborTasksetConfig()) - task = next(iter(taskset)) - - assert task["taskset_id"] == "harbor" - assert task["task_name"] == "task-a" - assert task["prompt"] == [ - {"role": "user", "content": "Write hello to /app/hello.txt"} - ] - assert task["sandbox"]["image"] == "ubuntu:24.04" - assert task["sandbox"]["memory_gb"] == 2.0 - assert task["sandbox"]["disk_size_gb"] == 8.0 - assert task["sandbox"]["command_timeout"] == 600 - assert "network_access" not in task["sandbox"] - assert ( - merge_task_sandbox( - vf.SandboxConfig(network_access=False, scope="rollout"), task - ).network_access - is False - ) - assert task["harbor"]["test_timeout"] == 300.0 - assert task["program"]["files"] == { - "/task/instruction.md": {"task": "instruction"}, - "/task/task.toml": {"task": "task_toml"}, - } - assert task["program"]["env"]["HARBOR_TASK_NAME"] == "task-a" - assert task["program"]["env"]["AGENT_WORKDIR"] == "/app" - - -def test_harbor_taskset_rejects_malformed_package_task( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - package = write_harbor_package(tmp_path, monkeypatch) - bad_task = cast(Path, getattr(package, "tasks_root")) / "bad-task" - bad_task.mkdir() - (bad_task / "task.toml").write_text('version = "1.0"') - - taskset = getattr(package, "load_taskset")(config=HarborTasksetConfig()) - - with pytest.raises(ValueError, match="Malformed Harbor task"): - list(taskset) - - -@pytest.mark.parametrize("section", ["agent", "verifier"]) -def test_harbor_task_rejects_non_mapping_agent_sections( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, section: str -) -> None: - package = write_harbor_package(tmp_path, monkeypatch) - task_dir = write_harbor_task(cast(Path, getattr(package, "tasks_root"))) - (task_dir / "task.toml").write_text( - f""" -version = "1.0" -{section} = "invalid" - -[environment] -docker_image = "ubuntu:24.04" -""".strip() - ) - taskset = getattr(package, "load_taskset")(config=HarborTasksetConfig()) - - with pytest.raises(TypeError, match=rf"\[{section}\] must be a mapping"): - list(taskset) - - -def test_harbor_taskset_constructs_env_with_opencode( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - package = write_harbor_package(tmp_path, monkeypatch) - write_harbor_task(cast(Path, getattr(package, "tasks_root"))) - - env = getattr(package, "load_env")() - - task = next(iter(env.taskset)) - assert task["task_name"] == "task-a" - assert isinstance(env.harness, OpenCode) - assert "task_dir" not in cast(dict[str, object], env.harness.config.program.data()) - - -class FakeHarborCommandResult: - def __init__( - self, - *, - exit_code: int = 0, - stdout: str = "", - stderr: str = "", - ): - self.exit_code = exit_code - self.stdout = stdout - self.stderr = stderr - - -class FakeHarborSandboxClient: - instances: list["FakeHarborSandboxClient"] = [] - - def __init__(self): - self.execute_commands: list[tuple[str, int | None, str | None]] = [] - self.background_jobs: list[tuple[str, str, int | None, str | None]] = [] - type(self).instances.append(self) - - async def upload_file(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - - async def execute_command( - self, *args: object, **kwargs: object - ) -> FakeHarborCommandResult: - command = str(kwargs.get("command") or args[1]) - timeout = cast(int | None, kwargs.get("timeout")) - working_dir = cast(str | None, kwargs.get("working_dir")) - self.execute_commands.append((command, timeout, working_dir)) - if "reward.txt" in command: - return FakeHarborCommandResult(stdout="1\n") - return FakeHarborCommandResult() - - async def run_background_job( - self, *args: object, **kwargs: object - ) -> FakeHarborCommandResult: - sandbox_id = str(kwargs.get("sandbox_id") or args[0]) - command = str(kwargs.get("command") or args[1]) - timeout = cast(int | None, kwargs.get("timeout")) - working_dir = cast(str | None, kwargs.get("working_dir")) - self.background_jobs.append((sandbox_id, command, timeout, working_dir)) - return FakeHarborCommandResult(stdout="tests passed") - - async def aclose(self) -> None: - pass - - -@pytest.mark.asyncio -async def test_harbor_reward_uses_background_job_for_tests( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - task_dir = write_harbor_task(tmp_path) - fake_module = cast(Any, types.ModuleType("prime_sandboxes")) - fake_module.AsyncSandboxClient = FakeHarborSandboxClient - monkeypatch.setitem(sys.modules, "prime_sandboxes", fake_module) - FakeHarborSandboxClient.instances = [] - - taskset = HarborTaskset(config=HarborTasksetConfig(bundle_package=__name__)) - reward = await taskset.harbor_reward( - vf.Task( - {"prompt": [], "harbor": {"task_dir": str(task_dir), "test_timeout": 120}} - ).freeze(), - vf.State({"sandbox_id": "sbx-1"}), - ) - - client = FakeHarborSandboxClient.instances[0] - assert reward == 1.0 - assert client.background_jobs == [("sbx-1", "bash test.sh", 120, "/tests")] - assert ("bash test.sh", 120, "/tests") not in client.execute_commands - - -def test_packaged_harbor_and_opencode_imports_are_available_from_packages() -> None: - assert OpenCode - assert OpenCodeConfig - assert Pi - assert Terminus2 - assert HarborTaskset - - -def test_opencode_config_owns_opencode_harness_fields() -> None: - harness = OpenCode( - config=OpenCodeConfig( - system_prompt=None, - program=OpenCodeProgramConfig( - agent_workdir="/workspace", - disabled_tools=["webfetch"], - ), - max_turns=2, - ) - ) - program = cast(dict[str, object], harness.program_config.data()) - command = cast(list[object], program["command"]) - mcp_setup = cast(dict[str, object], program["channels"])["mcp"] - setup = cast(str, program["setup"]) - - assert harness.config.program.agent_workdir == "/workspace" - assert harness.config.program.disabled_tools == ["webfetch"] - assert harness.config.system_prompt is None - assert harness.config.max_turns == 2 - assert "apt-get -o Acquire::Retries=3 update" in setup - assert "apt-get -o Acquire::Retries=3 install" in setup - assert "OPENCODE_RELEASE_REPO=PrimeIntellect-ai/opencode" in setup - assert "OPENCODE_RELEASE_PATH=releases/download/v1.1.63-rl2" in setup - assert "/workspace" in cast(str, command[2]) - assert '"webfetch": false' in cast(str, mcp_setup) - assert "/opencode/system.txt" in cast(dict[str, object], program["files"]) - - -@pytest.mark.parametrize( - "version", - ["PrimeIntellect-ai/opencode@latest", " PrimeIntellect-ai/opencode "], -) -def test_opencode_latest_version_uses_latest_download_url( - version: str, -) -> None: - harness = OpenCode( - config=OpenCodeConfig( - version=version, - program=OpenCodeProgramConfig( - install_ripgrep=False, - ), - ) - ) - program = cast(dict[str, object], harness.program_config.data()) - setup = cast(str, program["setup"]) - - assert "OPENCODE_RELEASE_REPO=PrimeIntellect-ai/opencode" in setup - assert "OPENCODE_RELEASE_PATH=releases/latest/download" in setup - - -def test_opencode_custom_version_uses_versioned_release() -> None: - harness = OpenCode( - config=OpenCodeConfig( - version="Example/open-code@v2.0.0", - ) - ) - program = cast(dict[str, object], harness.program_config.data()) - setup = cast(str, program["setup"]) - - assert "OPENCODE_RELEASE_REPO=Example/open-code" in setup - assert "OPENCODE_RELEASE_PATH=releases/download/v2.0.0" in setup - - -@pytest.mark.parametrize( - ("harness_cls", "config_cls", "program_cls"), - [ - (OpenCode, OpenCodeConfig, OpenCodeProgramConfig), - (MiniSWEAgent, MiniSWEAgentConfig, MiniSWEAgentProgramConfig), - (Pi, PiConfig, PiProgramConfig), - (RLM, RLMConfig, RLMProgramConfig), - (Terminus2, Terminus2Config, Terminus2ProgramConfig), - ], -) -def test_packaged_command_harnesses_defer_partial_program_overrides( - harness_cls, config_cls, program_cls -) -> None: - override = { - "setup": "echo caller", - "env": {"CALLER": "1"}, - "args": ["--caller"], - } - harness = harness_cls(config=config_cls(program=override)) - program = cast(dict[str, object], harness.program_config.data()) - env = cast(dict[str, object], program["env"]) - setup = cast(list[object], program["setup"]) - args = cast(list[object], program["args"]) - - assert program["command"] - assert env["CALLER"] == "1" - assert setup[-1] == "echo caller" - assert args[-1] == "--caller" - assert isinstance(harness.config.program, program_cls) - assert isinstance(harness.program_config, vf.ProgramConfig) - config_args = cast(list[object], harness.config.program.args) - assert harness.program_config.command == program["command"] - assert harness.config.program.env["CALLER"] == "1" - assert harness.config.program.setup == override["setup"] - assert config_args[-1] == "--caller" - - -def test_packaged_command_harness_config_program_patch_precedence() -> None: - harness = MiniSWEAgent( - config=MiniSWEAgentConfig( - program=MiniSWEAgentProgramConfig(env={"OPENAI_MODEL": "caller-model"}) - ) - ) - program = cast(dict[str, object], harness.program_config.data()) - env = cast(dict[str, object], program["env"]) - - assert env["OPENAI_MODEL"] == "caller-model" - - -@pytest.mark.parametrize( - ("key", "value"), - [ - ("command", ["other"]), - ("channels", "mcp"), - ], -) -def test_packaged_command_harness_config_program_rejects_owned_keys( - key: str, value: object -) -> None: - with pytest.raises(ValueError, match="Command ProgramConfig can only"): - OpenCode(config=OpenCodeConfig.model_validate({"program": {key: value}})) - - -def test_pi_harness_writes_intercepted_model_and_mcp_config() -> None: - harness = Pi() - program = cast(dict[str, object], harness.program_config.data()) - setup = cast(str, program["setup"]) - channels = cast(dict[str, object], program["channels"]) - mcp_setup = cast(str, channels["mcp"]) - - assert "apt-get -o Acquire::Retries=3 update" in setup - assert "apt-get -o Acquire::Retries=3 install" in setup - assert harness.config.version == PI_DEFAULT_VERSION - assert PI_DEFAULT_VERSION == "@earendil-works/pi-coding-agent@latest" - assert f"npm install -g --ignore-scripts {PI_DEFAULT_VERSION}" in setup - assert "mariozechner" not in setup - assert '"baseUrl": "${OPENAI_BASE_URL}"' in mcp_setup - assert '"api": "openai-completions"' in mcp_setup - assert '"apiKey": "${OPENAI_API_KEY:-intercepted}"' in mcp_setup - assert '"id": "model"' in mcp_setup - assert '"name": "${OPENAI_MODEL}"' in mcp_setup - assert f'"command": "{SANDBOX_PYTHON}"' in mcp_setup - - -def test_pi_harness_preserves_scoped_npm_versions() -> None: - harness = Pi(config=PiConfig(version="@anthropic-ai/claude-code@1.2.3")) - program = cast(dict[str, object], harness.program_config.data()) - setup = cast(str, program["setup"]) - - assert "npm install -g --ignore-scripts @anthropic-ai/claude-code@1.2.3" in setup - - -def test_terminus_2_harness_builds_sandbox_program() -> None: - harness = Terminus2( - config=Terminus2Config( - system_prompt="extra system prompt", - program=Terminus2ProgramConfig( - agent_workdir="/workspace", - max_turns=7, - python_version="3.12", - ), - ) - ) - program = cast(dict[str, object], harness.config.program.data()) - command = cast(list[object], program["command"]) - setup = cast(str, program["setup"]) - files = cast(dict[str, object], program["files"]) - artifacts = cast(dict[str, object], program["artifacts"]) - env = cast(dict[str, object], program.get("env", {})) - - assert isinstance(harness, vf.Harness) - assert "/terminus_2/instruction.md" in files - assert "/terminus_2/system_prompt.txt" in files - assert "apt-get -o Acquire::Retries=3 update" in setup - assert "apt-get -o Acquire::Retries=3 install" in setup - assert "git" not in setup - assert "terminus_2_log" in artifacts - assert "OPENAI_MODEL" not in env - - run_script = cast(str, command[2]) - assert "TERMINUS_2_WORKDIR=/workspace" in run_script - assert f"--with {TERMINUS_2_DEFAULT_VERSION}" in run_script - assert "git+https://github.com" not in run_script - assert "max_turns=7" in run_script - - script = run_script.split("python - <<'PY' 2>&1 | tee -a", 1)[1] - script = script.split("\n", 1)[1].rsplit("\nPY", 1)[0] - compile(script, "terminus_2_agent.py", "exec") - assert TERMINUS_2_DEFAULT_MODEL_NAME in script - assert TERMINUS_2_DEFAULT_API_BASE_URL in script - assert "OPENAI_MODEL" not in script - assert "PRIME_API_KEY" not in script - assert "async def prepare_logs_for_host(self) -> None" in script - assert "max_turns=7" in script - - -def test_task_program_merges_into_command_program_without_collisions() -> None: - harness = vf.Harness( - config=vf.HarnessConfig( - program=vf.ProgramConfig( - command=["tool"], - sandbox=True, - files={"/harness.txt": "harness"}, - setup="echo harness", - channels={"mcp": "echo harness tools"}, - env={"HARNESS": "1"}, - artifacts=vf.ArtifactsConfig.model_validate( - {"log": {"path": "/logs/harness.log", "format": "text"}} - ), - args=["--base"], - ), - sandbox=vf.SandboxConfig(image="python:3.11-slim"), - ) - ) - task = vf.Task( - { - "prompt": [], - "program": { - "files": {"/task/instruction.md": "task"}, - "setup": "echo task", - "env": {"TASK": "1"}, - "artifacts": {"task_log": {"path": "/logs/task.log", "format": "text"}}, - "args": ["--task"], - }, - } - ).freeze() - - program = merge_task_program( - cast(vf.ConfigData, harness.config.program.data()), task, kind="command" - ) - - assert program["files"] == { - "/harness.txt": "harness", - "/task/instruction.md": "task", - } - assert program["setup"] == ["echo harness", "echo task"] - assert program["channels"] == {"mcp": "echo harness tools"} - assert program["env"] == {"HARNESS": "1", "TASK": "1"} - assert program["args"] == ["--base", "--task"] - assert program["artifacts"] == { - "log": {"path": "/logs/harness.log", "format": "text"}, - "task_log": {"path": "/logs/task.log", "format": "text"}, - } - - -def test_command_program_patch_preserves_explicit_default_values() -> None: - program = vf.ProgramConfig(setup_timeout=300).resolve_command( - command=["tool"], - setup_timeout=600, - ) - - assert program.data()["setup_timeout"] == 300 - - -def test_task_program_rejects_harness_owned_keys() -> None: - harness = vf.Harness( - config=vf.HarnessConfig( - program=vf.ProgramConfig(command=["tool"], sandbox=True), - sandbox=vf.SandboxConfig(image="python:3.11-slim"), - ) - ) - task = vf.Task({"prompt": [], "program": {"command": ["other"]}}).freeze() - - with pytest.raises(ValueError, match="task.program can only define"): - merge_task_program( - cast(vf.ConfigData, harness.config.program.data()), - task, - kind="command", - ) - - -def test_task_program_rejects_colliding_upload_paths() -> None: - harness = vf.Harness( - config=vf.HarnessConfig( - program=vf.ProgramConfig( - command=["tool"], - sandbox=True, - files={"/task/instruction.md": "harness"}, - ), - sandbox=vf.SandboxConfig(image="python:3.11-slim"), - ) - ) - task = vf.Task( - {"prompt": [], "program": {"files": {"/task/instruction.md": "task"}}} - ).freeze() - - with pytest.raises(ValueError, match="define the same keys"): - merge_task_program( - cast(vf.ConfigData, harness.config.program.data()), - task, - kind="command", - ) diff --git a/tests/test_v1_mini_swe_agent.py b/tests/test_v1_mini_swe_agent.py deleted file mode 100644 index eb2f34c660..0000000000 --- a/tests/test_v1_mini_swe_agent.py +++ /dev/null @@ -1,126 +0,0 @@ -import importlib -from pathlib import Path -from types import ModuleType -from typing import Any, cast -from uuid import uuid4 - -import pytest -import verifiers as vf -from harnesses import MiniSWEAgent, MiniSWEAgentConfig, MiniSWEAgentProgramConfig - - -def write_harbor_task(root: Path) -> Path: - task_dir = root / "task-a" - (task_dir / "tests").mkdir(parents=True) - (task_dir / "solution").mkdir() - (task_dir / "instruction.md").write_text("Fix the bug.\n") - (task_dir / "task.toml").write_text( - """ -version = "1.0" - -[environment] -docker_image = "python:3.11-slim" -cpus = 2 -memory = "4G" -storage = "8G" - -[agent] -timeout_sec = 600 - -[verifier] -timeout_sec = 300 -""".strip() - ) - (task_dir / "tests" / "test.sh").write_text("echo 1 > /logs/verifier/reward.txt") - (task_dir / "solution" / "solve.sh").write_text("true") - return task_dir - - -def write_harbor_package(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> ModuleType: - package_name = f"mini_swe_harbor_pkg_{uuid4().hex}" - package_dir = tmp_path / package_name - tasks_root = package_dir / "tasks" - tasks_root.mkdir(parents=True) - (package_dir / "__init__.py").write_text( - """ -import verifiers as vf -from harnesses import MiniSWEAgent, MiniSWEAgentConfig -from tasksets import HarborTaskset, HarborTasksetConfig - - -def load_env(): - return vf.Env(taskset=HarborTaskset(config=HarborTasksetConfig(bundle_package=__name__)), harness=MiniSWEAgent(config=MiniSWEAgentConfig())) -""".lstrip() - ) - monkeypatch.syspath_prepend(str(tmp_path)) - importlib.invalidate_caches() - module = importlib.import_module(package_name) - setattr(module, "tasks_root", tasks_root) - return module - - -def test_mini_swe_agent_builds_sandbox_program(): - harness = MiniSWEAgent( - config=MiniSWEAgentConfig( - system_prompt="Use tests.", - program=MiniSWEAgentProgramConfig( - agent_workdir="/app", - ), - ) - ) - program = cast(dict[str, Any], harness.program_config.data()) - command = cast(list[str], program["command"]) - script = command[-1] - - assert isinstance(harness, vf.Harness) - assert program["sandbox"] is not False - assert "OPENAI_MODEL" in cast(dict[str, object], program["env"]) - assert "-c mini " in script - assert "model.model_class=litellm" in script - assert "model.model_kwargs.parallel_tool_calls=true" in script - assert "apt-get -o Acquire::Retries=3 update" in cast(str, program["setup"]) - assert "apt-get -o Acquire::Retries=3 install" in cast(str, program["setup"]) - assert "mini-swe-agent==2.2.8" in cast(str, program["setup"]) - assert "/mini-swe-agent/prompt.txt" in cast(dict[str, object], program["files"]) - assert "/mini-swe-agent/system.txt" in cast(dict[str, object], program["files"]) - assert "mini_swe_agent_log" in cast(dict[str, object], program["artifacts"]) - - -@pytest.mark.parametrize("version", ["mini-swe-agent@latest", " mini-swe-agent "]) -def test_mini_swe_agent_latest_version_uses_unpinned_pip_requirement( - version: str, -): - harness = MiniSWEAgent(config=MiniSWEAgentConfig(version=version)) - program = cast(dict[str, Any], harness.program_config.data()) - setup = cast(str, program["setup"]) - - assert ( - "vf_python_install --target /opt/mini-swe-agent/prefix/site-packages mini-swe-agent" - in setup - ) - - -def test_mini_swe_agent_pinned_version_uses_pip_requirement(): - harness = MiniSWEAgent(config=MiniSWEAgentConfig(version="mini-swe-agent@2.2.7")) - program = cast(dict[str, Any], harness.program_config.data()) - setup = cast(str, program["setup"]) - - assert "mini-swe-agent==2.2.7" in setup - - -def test_mini_swe_agent_composes_with_harbor_taskset( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -): - package = write_harbor_package(tmp_path, monkeypatch) - write_harbor_task(cast(Path, getattr(package, "tasks_root"))) - - env = getattr(package, "load_env")() - task = next(iter(env.taskset)) - - assert isinstance(env.harness, MiniSWEAgent) - assert task["taskset_id"] == "harbor" - assert task["instruction"] == "Fix the bug." - - -def test_mini_swe_agent_imports_from_package(): - assert MiniSWEAgent diff --git a/tests/test_v1_nemo_gym_harness.py b/tests/test_v1_nemo_gym_harness.py deleted file mode 100644 index 4f61f5d6a4..0000000000 --- a/tests/test_v1_nemo_gym_harness.py +++ /dev/null @@ -1,427 +0,0 @@ -import asyncio -import gzip -import os - -import pytest -import verifiers as vf -from aiohttp import ClientSession, web - -from harnesses.nemo_gym import ( - NEMO_GYM_EXTERNAL_POLICY_MODEL_ENTRYPOINT, - NEMO_GYM_POLICY_MODEL_SERVER_NAME, - NEMO_GYM_POLICY_MODEL_TYPE_NAME, - NeMoGymHarness, - NeMoGymHarnessConfig, - NeMoGymModelProxy, - PersistentNeMoGymRunner, - apply_nemo_gym_result, - build_nemo_gym_global_config, - build_nemo_gym_policy_model_config, - disable_ray_uv_run_runtime_env, - nemo_gym_proxy_model_name, - set_nemo_gym_proxy_model, - skip_nemo_gym_policy_model_process, -) -from verifiers.utils.serve_utils import get_free_port - - -@pytest.mark.asyncio -async def test_nemo_gym_proxy_routes_concurrent_rollouts_by_model(): - upstream_a = await _start_upstream("a") - upstream_b = await _start_upstream("b") - proxy = NeMoGymModelProxy() - await proxy.start() - routing_model_a = nemo_gym_proxy_model_name("rollout-a") - routing_model_b = nemo_gym_proxy_model_name("rollout-b") - - try: - async with ( - upstream_a, - upstream_b, - proxy.activate( - routing_model_a, - { - "base_url": upstream_a.base_url, - "api_key": "key-a", - "model": "model-a", - }, - ), - proxy.activate( - routing_model_b, - { - "base_url": upstream_b.base_url, - "api_key": "key-b", - "model": "model-b", - }, - ), - ): - async with ClientSession() as session: - response_a, response_b = await asyncio.gather( - _proxy_response(session, proxy, routing_model_a), - _proxy_response(session, proxy, routing_model_b), - ) - - assert response_a == {"label": "a", "model": "model-a"} - assert response_b == {"label": "b", "model": "model-b"} - assert upstream_a.authorizations == ["Bearer key-a"] - assert upstream_b.authorizations == ["Bearer key-b"] - finally: - await proxy.stop() - - -@pytest.mark.asyncio -async def test_nemo_gym_proxy_rejects_unrouted_request_with_multiple_rollouts(): - upstream_a = await _start_upstream("a") - upstream_b = await _start_upstream("b") - proxy = NeMoGymModelProxy() - await proxy.start() - routing_model_a = nemo_gym_proxy_model_name("rollout-a") - routing_model_b = nemo_gym_proxy_model_name("rollout-b") - - try: - async with ( - upstream_a, - upstream_b, - proxy.activate( - routing_model_a, - { - "base_url": upstream_a.base_url, - "api_key": "key-a", - "model": "model-a", - }, - ), - proxy.activate( - routing_model_b, - { - "base_url": upstream_b.base_url, - "api_key": "key-b", - "model": "model-b", - }, - ), - ): - async with ClientSession() as session: - response = await session.post( - f"http://{proxy.host}:{proxy.port}/v1/responses", - headers={"Authorization": f"Bearer {proxy.secret}"}, - json={"model": "ignored"}, - ) - body = await response.json() - - assert response.status == 409 - assert "model" in body["error"] - finally: - await proxy.stop() - - -@pytest.mark.asyncio -async def test_nemo_gym_proxy_falls_back_when_only_one_rollout_is_active(): - upstream = await _start_upstream("single") - proxy = NeMoGymModelProxy() - await proxy.start() - - try: - async with ( - upstream, - proxy.activate( - nemo_gym_proxy_model_name("rollout"), - { - "base_url": upstream.base_url, - "api_key": "key", - "model": "real-model", - }, - ), - ): - async with ClientSession() as session: - response = await session.post( - f"http://{proxy.host}:{proxy.port}/v1/responses", - headers={"Authorization": f"Bearer {proxy.secret}"}, - json={"model": "verifiers-nemo-gym-proxy"}, - ) - body = await response.json() - - assert response.status == 200 - assert body == {"label": "single", "model": "real-model"} - finally: - await proxy.stop() - - -@pytest.mark.asyncio -async def test_nemo_gym_proxy_strips_content_encoding_after_decompression(): - async def handle_response(request: web.Request) -> web.Response: - body = await request.json() - return web.Response( - body=gzip.compress(f'{{"model":"{body["model"]}","ok":true}}'.encode()), - headers={ - "Content-Encoding": "gzip", - "Content-Type": "application/json", - }, - ) - - app = web.Application() - app.router.add_post("/v1/responses", handle_response) - runner = web.AppRunner(app) - await runner.setup() - port = get_free_port() - site = web.TCPSite(runner, "127.0.0.1", port) - await site.start() - - proxy = NeMoGymModelProxy() - await proxy.start() - routing_model = nemo_gym_proxy_model_name("rollout") - try: - async with proxy.activate( - routing_model, - { - "base_url": f"http://127.0.0.1:{port}/v1", - "api_key": "key", - "model": "real-model", - }, - ): - async with ClientSession(auto_decompress=False) as session: - response = await session.post( - f"http://{proxy.host}:{proxy.port}/v1/responses", - headers={"Authorization": f"Bearer {proxy.secret}"}, - json={"model": routing_model}, - ) - body = await response.read() - - assert response.status == 200 - assert "Content-Encoding" not in response.headers - assert body == b'{"model":"real-model","ok":true}' - finally: - await proxy.stop() - await runner.cleanup() - - -def test_nemo_gym_global_config_uses_proxy_endpoint_without_header_forwarding(): - config = build_nemo_gym_global_config( - config_paths=["agent.yaml"], - endpoint_config={ - "base_url": "http://127.0.0.1:12345/v1", - "api_key": "secret", - "model": "proxy-model", - }, - global_config={"custom": "value"}, - ) - - assert config["policy_base_url"] == "http://127.0.0.1:12345/v1" - assert config["policy_api_key"] == "secret" - assert config["policy_model_name"] == "proxy-model" - assert config["custom"] == "value" - assert "forward_request_headers" not in config - assert config[NEMO_GYM_POLICY_MODEL_SERVER_NAME] == { - "responses_api_models": { - NEMO_GYM_POLICY_MODEL_TYPE_NAME: { - "entrypoint": NEMO_GYM_EXTERNAL_POLICY_MODEL_ENTRYPOINT, - "host": "127.0.0.1", - "port": 12345, - } - } - } - - -def test_nemo_gym_harness_does_not_add_a_model_server_config_path(): - harness = NeMoGymHarness( - NeMoGymHarnessConfig( - config_paths=["agent.yaml"], - server_name="agent_server", - agent_name="agent", - ) - ) - - assert harness._config_paths() == ["agent.yaml"] - - -def test_apply_nemo_gym_result_rejects_non_numeric_string_reward(): - with pytest.raises(TypeError, match="reward must be numeric"): - apply_nemo_gym_result(vf.State(), {"reward": "high"}) - - -def test_build_nemo_gym_policy_model_config_requires_explicit_port(): - with pytest.raises(ValueError, match="host and port"): - build_nemo_gym_policy_model_config( - { - "base_url": "https://api.openai.com/v1", - "api_key": "secret", - "model": "model", - } - ) - - -def test_set_nemo_gym_proxy_model_preserves_row_without_mutating_create_params(): - create_params = {"input": [{"role": "user", "content": "hi"}], "temperature": 0.2} - row = {"responses_create_params": create_params} - - set_nemo_gym_proxy_model(row, "proxy-rollout") - - assert row["responses_create_params"] == { - "input": [{"role": "user", "content": "hi"}], - "temperature": 0.2, - "model": "proxy-rollout", - } - assert create_params == { - "input": [{"role": "user", "content": "hi"}], - "temperature": 0.2, - } - - -def test_disable_ray_uv_run_runtime_env_sets_and_restores_env(monkeypatch): - monkeypatch.delenv("RAY_ENABLE_UV_RUN_RUNTIME_ENV", raising=False) - - with disable_ray_uv_run_runtime_env(): - assert os.environ["RAY_ENABLE_UV_RUN_RUNTIME_ENV"] == "0" - - assert "RAY_ENABLE_UV_RUN_RUNTIME_ENV" not in os.environ - - -def test_skip_nemo_gym_policy_model_process_leaves_other_processes_untouched(): - class FakeNemoCliModule: - def __init__(self) -> None: - self.calls = [] - self.setup_calls = [] - - def setup_env_command( - self, dir_path: object, global_config_dict: object, prefix: str - ) -> str: - self.setup_calls.append((dir_path, global_config_dict, prefix)) - return "real-setup" - - def run_command(self, command: str, working_dir_path: object): - self.calls.append((command, working_dir_path)) - return "real-process" - - module = FakeNemoCliModule() - with skip_nemo_gym_policy_model_process(module): - assert module.setup_env_command(".", {}, "policy_model") == "true" - assert ( - module.setup_env_command(".", {}, "example_single_tool_call") - == "real-setup" - ) - proxy_process = module.run_command( - "NEMO_GYM_CONFIG_PATH=policy_model python app.py", - ".", - ) - real_process = module.run_command( - "NEMO_GYM_CONFIG_PATH=example_single_tool_call python app.py", - ".", - ) - - assert proxy_process.poll() is None - proxy_process.send_signal(2) - assert proxy_process.poll() == 0 - assert real_process == "real-process" - assert module.setup_calls == [(".", {}, "example_single_tool_call")] - assert module.calls == [ - ("NEMO_GYM_CONFIG_PATH=example_single_tool_call python app.py", ".") - ] - - -@pytest.mark.asyncio -async def test_nemo_gym_runner_uses_own_head_server_config_for_rollouts(): - head_server_config = object() - collector = FakeRolloutCollector({"reward": 1.0}) - runner = PersistentNeMoGymRunner() - runner._helper = FakeRunHelper() - runner._rollout_collector = collector - runner._proxy = NeMoGymModelProxy() - runner._head_server_config = head_server_config - - result = await runner._run_once( - { - "agent_ref": {"type": "responses_api_agents", "name": "agent"}, - "responses_create_params": {"input": "hi"}, - }, - server_name=None, - agent_name=None, - endpoint_config={ - "base_url": "http://127.0.0.1:1/v1", - "api_key": "key", - "model": "real-model", - }, - ) - - assert result == {"reward": 1.0} - assert collector.head_server_config is head_server_config - assert collector.row["responses_create_params"]["model"].startswith( - "verifiers-nemo-gym-proxy-" - ) - - -async def _proxy_response( - session: ClientSession, proxy: NeMoGymModelProxy, routing_model: str -) -> dict[str, str]: - response = await session.post( - f"http://{proxy.host}:{proxy.port}/v1/responses", - headers={"Authorization": f"Bearer {proxy.secret}"}, - json={"model": routing_model}, - ) - assert response.status == 200 - return await response.json() - - -class UpstreamServer: - def __init__( - self, - *, - label: str, - runner: web.AppRunner, - site: web.TCPSite, - base_url: str, - authorizations: list[str], - ) -> None: - self.label = label - self.runner = runner - self.site = site - self.base_url = base_url - self.authorizations = authorizations - - async def __aenter__(self) -> "UpstreamServer": - return self - - async def __aexit__(self, *args: object) -> None: - await self.runner.cleanup() - - -async def _start_upstream(label: str) -> UpstreamServer: - authorizations: list[str] = [] - - async def handle_response(request: web.Request) -> web.Response: - authorizations.append(request.headers["Authorization"]) - body = await request.json() - return web.json_response({"label": label, "model": body["model"]}) - - app = web.Application() - app.router.add_post("/v1/responses", handle_response) - runner = web.AppRunner(app) - await runner.setup() - port = get_free_port() - site = web.TCPSite(runner, "127.0.0.1", port) - await site.start() - return UpstreamServer( - label=label, - runner=runner, - site=site, - base_url=f"http://127.0.0.1:{port}/v1", - authorizations=authorizations, - ) - - -class FakeRunHelper: - def poll(self) -> None: - return None - - -class FakeRolloutCollector: - def __init__(self, result: dict[str, float]) -> None: - self.result = result - self.row: dict | None = None - self.head_server_config: object | None = None - - def run_examples( - self, rows: list[dict], *, head_server_config: object | None = None - ): - self.row = rows[0] - self.head_server_config = head_server_config - future = asyncio.get_running_loop().create_future() - future.set_result((self.row, self.result)) - return iter([future]) diff --git a/tests/test_v1_openenv_taskset.py b/tests/test_v1_openenv_taskset.py deleted file mode 100644 index c1244a3694..0000000000 --- a/tests/test_v1_openenv_taskset.py +++ /dev/null @@ -1,236 +0,0 @@ -import json -from collections.abc import Awaitable, Callable -from pathlib import Path -from typing import cast - -import pytest - -import verifiers as vf -from openenv.core.env_server.mcp_types import Tool as OpenEnvToolSpec -from tasksets import openenv - - -class OpenEnvStepResult: - def __init__( - self, observation: dict[str, object], reward: float | None, done: bool - ): - self.observation = observation - self.reward = reward - self.done = done - - -class FakeGenericEnvClient: - instances: list["FakeGenericEnvClient"] = [] - - @classmethod - async def from_docker_image(cls, *args: object, **kwargs: object): - del args, kwargs - client = cls(base_url="http://localhost:8000") - await client.connect() - return client - - def __init__(self, base_url: str): - self.base_url = base_url - self.connected = False - self.closed = False - self.reset_seeds: list[int] = [] - self.actions: list[dict[str, object]] = [] - FakeGenericEnvClient.instances.append(self) - - async def connect(self) -> None: - self.connected = True - - async def reset(self, *, seed: int) -> OpenEnvStepResult: - self.reset_seeds.append(seed) - return OpenEnvStepResult({"prompt": f"seed-{seed}"}, None, False) - - async def step(self, action: dict[str, object]) -> OpenEnvStepResult: - self.actions.append(action) - return OpenEnvStepResult({"prompt": "done"}, 1.0, True) - - async def close(self) -> None: - self.closed = True - - -class FakeMCPToolClient: - instances: list["FakeMCPToolClient"] = [] - - @classmethod - async def from_docker_image(cls, *args: object, **kwargs: object): - del args, kwargs - client = cls(base_url="http://localhost:8000") - await client.connect() - return client - - def __init__(self, base_url: str): - self.base_url = base_url - self.connected = False - self.closed = False - self.actions: list[object] = [] - FakeMCPToolClient.instances.append(self) - - async def connect(self) -> None: - self.connected = True - - async def reset(self, *, seed: int) -> OpenEnvStepResult: - return OpenEnvStepResult({"prompt": f"mcp-{seed}"}, None, False) - - async def list_tools(self) -> list[OpenEnvToolSpec]: - return [ - OpenEnvToolSpec( - name="echo", - description="Echo a message", - input_schema={ - "type": "object", - "properties": {"message": {"type": "string"}}, - "required": ["message"], - }, - ) - ] - - async def step(self, action: object) -> OpenEnvStepResult: - self.actions.append(action) - return OpenEnvStepResult({"result": {"data": "ok"}}, 0.5, True) - - async def close(self) -> None: - self.closed = True - - -class FakeOpenEnvProvider: - def __init__(self, spec: openenv.OpenEnvRuntimeConfig): - self.spec = spec - self.base_url = "http://localhost:8000" - self.stopped = False - - def stop_container(self) -> None: - self.stopped = True - - def fetch_schema(self) -> dict[str, object]: - if self.spec.contract == "mcp": - return { - "action": { - "type": "object", - "properties": {"type": {"enum": ["list_tools", "call_tool"]}}, - } - } - return { - "action": { - "type": "object", - "properties": {"command": {"type": "string"}}, - "required": ["command"], - } - } - - -def openenv_prompt_renderer(observation: object, **kwargs: object) -> list[vf.Message]: - del kwargs - assert isinstance(observation, dict) - observation_data = cast(dict[str, object], observation) - return [vf.UserMessage(content=str(observation_data["prompt"]))] - - -@pytest.fixture -def fake_openenv_runtime(monkeypatch): - FakeGenericEnvClient.instances.clear() - FakeMCPToolClient.instances.clear() - - monkeypatch.setattr(openenv, "PrimeSandboxOpenEnvProvider", FakeOpenEnvProvider) - monkeypatch.setattr(openenv, "GenericEnvClient", FakeGenericEnvClient) - monkeypatch.setattr(openenv, "MCPToolClient", FakeMCPToolClient) - - -def write_openenv_manifest(project: Path, contract: str) -> None: - project.mkdir(exist_ok=True) - (project / ".build.json").write_text( - json.dumps( - { - "image": "image", - "port": 8000, - "start_command": "run", - "contract": contract, - } - ) - ) - - -@pytest.mark.asyncio -async def test_openenv_taskset_runs_gym_rollout_boundary( - tmp_path, fake_openenv_runtime -): - write_openenv_manifest(tmp_path, "gym") - taskset = openenv.OpenEnvTaskset( - config=openenv.OpenEnvTasksetConfig( - openenv_project=str(tmp_path), - prompt_renderer="tests.test_v1_openenv_taskset:openenv_prompt_renderer", - num_train_examples=1, - num_eval_examples=0, - seed=7, - ) - ) - env = vf.Env(taskset=taskset, harness=vf.Harness()) - task = next(iter(taskset)) - state = vf.State.for_task(task) - - await env.harness.setup_state(task, state) - await env.harness.runtime.setup_rollout(task, state) - - assert state["prompt"] == [vf.UserMessage(content="seed-7")] - assert "openenv_client" not in state - assert "openenv_action_schema" not in state - client = FakeGenericEnvClient.instances[0] - assert client.connected is True - assert client.reset_seeds == [7] - - state["completion"] = [vf.AssistantMessage(content='{"command": "advance"}')] - state["trajectory"].append({"reward": None}) - messages = await env.harness.runtime.user_messages(task, state) - - assert client.actions == [{"command": "advance"}] - assert messages == [{"role": "user", "content": "done"}] - assert state["trajectory"][-1]["reward"] == 1.0 - assert state["openenv_done"] is True - - await env.harness.runtime.cleanup_rollout(task, state) - assert client.closed is True - - -@pytest.mark.asyncio -async def test_openenv_taskset_exposes_mcp_tools(tmp_path, fake_openenv_runtime): - write_openenv_manifest(tmp_path, "mcp") - taskset = openenv.OpenEnvTaskset( - config=openenv.OpenEnvTasksetConfig( - openenv_project=str(tmp_path), - prompt_renderer="tests.test_v1_openenv_taskset:openenv_prompt_renderer", - num_train_examples=1, - num_eval_examples=0, - seed=9, - ) - ) - env = vf.Env(taskset=taskset, harness=vf.Harness()) - task = next(iter(taskset)) - state = vf.State.for_task(task) - - await env.harness.setup_state(task, state) - await env.harness.runtime.setup_rollout(task, state) - - assert state["prompt"] == [vf.UserMessage(content="mcp-9")] - assert state["tools"] == ["echo"] - assert "openenv_client" not in state - assert "openenv_action_schema" not in state - client = FakeMCPToolClient.instances[0] - state["trajectory"].append({"reward": None}) - tool = cast( - Callable[..., Awaitable[object]], - env.harness.runtime.tool_calls(task, state)["echo"], - ) - result = await tool(message="hello") - - action = client.actions[0] - assert getattr(action, "tool_name") == "echo" - assert getattr(action, "arguments") == {"message": "hello"} - assert result == "ok" - assert state["trajectory"][-1]["reward"] == 0.5 - assert state["openenv_done"] is True - - await env.harness.runtime.cleanup_rollout(task, state) - assert client.closed is True diff --git a/tests/test_v1_openreward_taskset.py b/tests/test_v1_openreward_taskset.py deleted file mode 100644 index d8530b829c..0000000000 --- a/tests/test_v1_openreward_taskset.py +++ /dev/null @@ -1,243 +0,0 @@ -from collections.abc import Awaitable, Callable -from typing import cast - -import pytest - -import verifiers as vf - -pytest.importorskip("openreward") - -from openreward.api.environments.types import ( - Task as OpenRewardTask, - TextBlock as OpenRewardTextBlock, - ToolOutput as OpenRewardToolOutput, -) -from tasksets import openreward - - -class FakeOpenRewardSession: - def __init__(self, task: OpenRewardTask): - self.task = task - self.entered = False - self.exited = False - self.calls: list[tuple[str, dict[str, object]]] = [] - - def __enter__(self) -> "FakeOpenRewardSession": - self.entered = True - return self - - def __exit__(self, *exc: object) -> None: - self.exited = True - - def get_prompt(self) -> list[OpenRewardTextBlock]: - return [OpenRewardTextBlock(text="Solve the task.")] - - def list_tools(self, format: str | None = None) -> list[dict[str, object]]: - assert format == "openai" - return [ - { - "type": "function", - "name": "answer", - "description": "Submit an answer", - "parameters": { - "type": "object", - "properties": {"answer": {"type": "string"}}, - "required": ["answer"], - }, - } - ] - - def call_tool( - self, tool_name: str, input: dict[str, object] - ) -> OpenRewardToolOutput: - self.calls.append((tool_name, input)) - return OpenRewardToolOutput( - blocks=[OpenRewardTextBlock(text="Correct.")], - reward=1.0, - finished=True, - metadata={"status": "ok"}, - ) - - -class FakeOpenRewardSplit: - def __init__(self, name: str, type: str): - self.name = name - self.type = type - - -class FakeOpenRewardEnvironment: - def __init__(self): - self.sessions: list[FakeOpenRewardSession] = [] - self.task_range_calls: list[tuple[str, int | None, int | None]] = [] - self.splits = [ - FakeOpenRewardSplit(name="train", type="train"), - FakeOpenRewardSplit(name="official-test", type="test"), - ] - - def list_splits(self) -> list[FakeOpenRewardSplit]: - return self.splits - - def list_tasks(self, split: str) -> list[OpenRewardTask]: - return [ - OpenRewardTask( - server_name="owner/env", - environment_name="env", - namespace="owner", - task_spec={"id": f"{split}-0"}, - ) - ] - - def get_task_range( - self, split: str, start: int | None = None, stop: int | None = None - ) -> list[OpenRewardTask]: - self.task_range_calls.append((split, start, stop)) - return [ - OpenRewardTask( - server_name="owner/env", - environment_name="env", - namespace="owner", - task_spec={"id": f"{split}-{index}"}, - ) - for index in range(start or 0, stop or 0) - ] - - def session(self, task: OpenRewardTask) -> FakeOpenRewardSession: - session = FakeOpenRewardSession(task) - self.sessions.append(session) - return session - - -class FakeOpenRewardEnvironmentsAPI: - def __init__(self, environment: FakeOpenRewardEnvironment): - self.environment = environment - self.get_calls: list[dict[str, object]] = [] - - def get( - self, - name: str, - variant: str | None = None, - base_url: str | None = None, - ) -> FakeOpenRewardEnvironment: - self.get_calls.append({"name": name, "variant": variant, "base_url": base_url}) - return self.environment - - -class FakeOpenRewardClient: - instances: list["FakeOpenRewardClient"] = [] - - def __init__(self, environment: FakeOpenRewardEnvironment): - self.environments = FakeOpenRewardEnvironmentsAPI(environment) - self.closed = False - FakeOpenRewardClient.instances.append(self) - - def __enter__(self) -> "FakeOpenRewardClient": - return self - - def __exit__(self, *exc: object) -> None: - self.close() - - def close(self) -> None: - self.closed = True - - -@pytest.fixture -def fake_openreward_client(monkeypatch): - FakeOpenRewardClient.instances.clear() - environment = FakeOpenRewardEnvironment() - - def client_factory(): - return FakeOpenRewardClient(environment) - - monkeypatch.setattr(openreward, "OpenReward", client_factory) - return environment - - -def test_openreward_taskset_loads_serializable_tasks(fake_openreward_client): - taskset = openreward.OpenRewardTaskset( - config=openreward.OpenRewardTasksetConfig( - environment="owner/env", - split="train", - num_train_examples=2, - ) - ) - - tasks = list(taskset.get_dataset()) - task = taskset.to_task(tasks[0]) - - assert fake_openreward_client.task_range_calls == [("train", 0, 2)] - assert task["openreward"]["environment"] == "owner/env" - assert task["openreward"]["task"] == { - "server_name": "owner/env", - "environment_name": "env", - "namespace": "owner", - "task_spec": {"id": "train-0"}, - } - assert set(taskset.named_toolsets) == {"openreward"} - - -def test_openreward_taskset_load_tasks_eval_uses_test_split(fake_openreward_client): - taskset = openreward.OpenRewardTaskset( - config=openreward.OpenRewardTasksetConfig( - environment="owner/env", - split="train", - num_eval_examples=2, - ) - ) - - tasks = list(taskset.load_tasks(split="eval")) - - assert fake_openreward_client.task_range_calls == [("official-test", 0, 2)] - assert tasks[0]["openreward"]["split"] == "official-test" - assert tasks[0]["openreward"]["task"]["task_spec"] == {"id": "official-test-0"} - - -def test_openreward_taskset_eval_split_empty_without_test_split( - fake_openreward_client, -): - fake_openreward_client.splits = [FakeOpenRewardSplit(name="train", type="train")] - taskset = openreward.OpenRewardTaskset( - config=openreward.OpenRewardTasksetConfig( - environment="owner/env", - num_eval_examples=2, - ) - ) - - assert len(taskset.get_eval_dataset()) == 0 - assert fake_openreward_client.task_range_calls == [] - - -@pytest.mark.asyncio -async def test_openreward_taskset_setup_and_tool_call(fake_openreward_client): - taskset = openreward.OpenRewardTaskset( - config=openreward.OpenRewardTasksetConfig( - environment="owner/env", - split="train", - num_train_examples=1, - ) - ) - env = vf.Env(taskset=taskset, harness=vf.Harness()) - task = next(iter(taskset)) - state = vf.State.for_task(task) - - await env.harness.setup_state(task, state) - await env.harness.runtime.setup_rollout(task, state) - - assert state["prompt"] == [vf.UserMessage(content="Solve the task.")] - assert state["tools"] == ["answer"] - state["trajectory"].append({"reward": None}) - tool = cast( - Callable[..., Awaitable[object]], - env.harness.runtime.tool_calls(task, state)["answer"], - ) - result = await tool(answer="4") - - session = fake_openreward_client.sessions[0] - assert session.entered is True - assert session.calls == [("answer", {"answer": "4"})] - assert result == "Correct." - assert state["trajectory"][-1]["reward"] == 1.0 - assert state["openreward_finished"] is True - - await env.harness.runtime.cleanup_rollout(task, state) - assert session.exited is True - assert FakeOpenRewardClient.instances[-1].closed is True diff --git a/tests/test_v1_replay_harness.py b/tests/test_v1_replay_harness.py deleted file mode 100644 index b8e78b5e05..0000000000 --- a/tests/test_v1_replay_harness.py +++ /dev/null @@ -1,327 +0,0 @@ -import json -from pathlib import Path - -import pytest - -import verifiers as vf -from harnesses import ReplayHarness -from tasksets import ReplayTaskset, ReplayTasksetConfig -from tasksets.replay import replay_task_record - - -class NoModelClient: - def __init__(self) -> None: - self.requests = 0 - - async def get_response(self, **kwargs: object) -> object: - _ = kwargs - self.requests += 1 - raise AssertionError("ReplayHarness must not request model completions.") - - -class InlineReplayTaskset(ReplayTaskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "eval": - return [] - return [ - { - "messages": [ - {"role": "user", "content": "Reverse abc."}, - {"role": "assistant", "content": "cba"}, - {"role": "user", "content": "Now uppercase it."}, - { - "role": "assistant", - "content": "CBA", - "reasoning_content": "uppercased the prior answer", - }, - {"role": "user", "content": "Thanks."}, - ], - } - ] - - -class ManyTurnReplayTaskset(ReplayTaskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "eval": - return [] - messages = [] - for index in range(11): - messages.append({"role": "user", "content": f"Turn {index}?"}) - messages.append({"role": "assistant", "content": f"reply {index}"}) - return [{"messages": messages}] - - -@pytest.mark.asyncio -async def test_replay_harness_prints_assistant_messages_into_trajectory() -> None: - env = vf.Env( - taskset=InlineReplayTaskset(), - harness=ReplayHarness(config=vf.HarnessConfig()), - ) - client = NoModelClient() - - state = await env.rollout( - dict(env.get_dataset()[0]), - client=client, - model="mock-model", - ) - - assert client.requests == 0 - assert state["stop_condition"] == "replayed_messages" - assert state["num_model_requests"] == 2 - assert state["prompt"] == [{"role": "user", "content": "Reverse abc."}] - assert state["completion"] == [ - {"role": "assistant", "content": "cba"}, - {"role": "user", "content": "Now uppercase it."}, - { - "role": "assistant", - "content": "CBA", - "reasoning_content": "uppercased the prior answer", - }, - ] - assert state["completion"][-1]["role"] == "assistant" - - first, second = state["trajectory"] - assert first["prompt"] == [{"role": "user", "content": "Reverse abc."}] - assert first["completion"] == [{"role": "assistant", "content": "cba"}] - assert first["tokens"] is None - assert "tokens" not in first["response"]["message"] - - assert second["prompt"] == [ - {"role": "user", "content": "Reverse abc."}, - {"role": "assistant", "content": "cba"}, - {"role": "user", "content": "Now uppercase it."}, - ] - assert second["completion"] == [ - { - "role": "assistant", - "content": "CBA", - "reasoning_content": "uppercased the prior answer", - } - ] - assert second["tokens"] is None - assert "tokens" not in second["response"]["message"] - - -@pytest.mark.asyncio -async def test_replay_harness_marks_partial_replay_as_truncated() -> None: - env = vf.Env( - taskset=InlineReplayTaskset(), - harness=ReplayHarness(config=vf.HarnessConfig(max_turns=1)), - ) - - state = await env.rollout( - dict(env.get_dataset()[0]), - client=NoModelClient(), - model="mock-model", - ) - - assert state["stop_condition"] == "max_turns_reached" - assert state["is_truncated"] is True - assert state["num_model_requests"] == 1 - assert state["completion"] == [{"role": "assistant", "content": "cba"}] - step = state["trajectory"][0] - assert step["is_truncated"] is True - assert step["response"]["message"]["is_truncated"] is True - - -@pytest.mark.asyncio -async def test_replay_harness_defaults_to_all_assistant_messages() -> None: - assert vf.HarnessConfig().max_turns == -1 - env = vf.Env( - taskset=ManyTurnReplayTaskset(), - harness=ReplayHarness(config=vf.HarnessConfig()), - ) - - state = await env.rollout( - dict(env.get_dataset()[0]), - client=NoModelClient(), - model="mock-model", - ) - - assert state["stop_condition"] == "replayed_messages" - assert state["is_truncated"] is False - assert state["num_model_requests"] == 11 - assert len(state["trajectory"]) == 11 - assert state["completion"][-1] == {"role": "assistant", "content": "reply 10"} - - -def test_replay_taskset_loads_configured_local_jsonl_data(tmp_path: Path) -> None: - data_dir = tmp_path / "data" - data_dir.mkdir(parents=True) - (data_dir / "examples.jsonl").write_text( - "\n".join( - [ - json.dumps( - { - "messages": [ - {"role": "user", "content": "Say ok."}, - {"role": "assistant", "content": "ok"}, - ] - } - ), - json.dumps( - { - "messages": [ - {"role": "user", "content": "Say yes."}, - {"role": "assistant", "content": "yes"}, - ] - } - ), - ] - ) - + "\n", - encoding="utf-8", - ) - - taskset = ReplayTaskset(config=ReplayTasksetConfig(data_dir=str(data_dir))) - - assert taskset.load_tasks() == [ - { - "messages": [ - {"role": "user", "content": "Say ok."}, - {"role": "assistant", "content": "ok"}, - ] - }, - { - "messages": [ - {"role": "user", "content": "Say yes."}, - {"role": "assistant", "content": "yes"}, - ] - }, - ] - - -def test_replay_taskset_loads_subclass_local_jsonl_data(tmp_path: Path) -> None: - data_dir = tmp_path / "data" - data_dir.mkdir(parents=True) - (data_dir / "example.jsonl").write_text( - json.dumps( - { - "messages": [ - {"role": "user", "content": "Say ok."}, - {"role": "assistant", "content": "ok"}, - ] - } - ) - + "\n", - encoding="utf-8", - ) - local_taskset_type = type( - "LocalReplayTaskset", - (ReplayTaskset,), - {"data_dir": str(data_dir)}, - ) - taskset = local_taskset_type(config=ReplayTasksetConfig()) - - assert taskset.load_tasks() == [ - { - "messages": [ - {"role": "user", "content": "Say ok."}, - {"role": "assistant", "content": "ok"}, - ] - } - ] - - -def test_replay_taskset_rejects_missing_local_source() -> None: - taskset = ReplayTaskset(config=ReplayTasksetConfig()) - - with pytest.raises(FileNotFoundError, match="requires dataset or data_dir"): - taskset.load_tasks() - - -def test_replay_taskset_rejects_conflicting_sources(tmp_path: Path) -> None: - taskset = ReplayTaskset( - config=ReplayTasksetConfig(dataset="owner/dataset", data_dir=str(tmp_path)) - ) - - with pytest.raises(ValueError, match="cannot set both dataset and data_dir"): - taskset.load_tasks() - - -def test_replay_taskset_rejects_empty_local_data_dir(tmp_path: Path) -> None: - data_dir = tmp_path / "data" - data_dir.mkdir(parents=True) - - taskset = ReplayTaskset(config=ReplayTasksetConfig(data_dir=str(data_dir))) - - with pytest.raises(FileNotFoundError, match="must contain at least one JSONL"): - taskset.load_tasks() - - -def test_replay_taskset_rejects_json_files(tmp_path: Path) -> None: - data_dir = tmp_path / "data" - data_dir.mkdir(parents=True) - (data_dir / "example.json").write_text( - json.dumps( - { - "messages": [ - {"role": "user", "content": "Say ok."}, - {"role": "assistant", "content": "ok"}, - ] - } - ), - encoding="utf-8", - ) - - taskset = ReplayTaskset(config=ReplayTasksetConfig(data_dir=str(data_dir))) - - with pytest.raises(ValueError, match=r"accepts only \.jsonl files"): - taskset.load_tasks() - - -def test_replay_taskset_rejects_non_object_jsonl_rows(tmp_path: Path) -> None: - data_dir = tmp_path / "data" - data_dir.mkdir(parents=True) - (data_dir / "example.jsonl").write_text("[]\n", encoding="utf-8") - - taskset = ReplayTaskset(config=ReplayTasksetConfig(data_dir=str(data_dir))) - - with pytest.raises(TypeError, match="example.jsonl:1 must contain one JSON object"): - taskset.load_tasks() - - -def test_replay_taskset_canonicalizes_messages() -> None: - task = replay_task_record( - { - "messages": [ - {"role": "user", "content": "Use the tool."}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": { - "name": "search", - "arguments": {"query": "abc"}, - }, - } - ], - }, - ] - } - ) - - assert task["messages"] == [ - {"role": "user", "content": "Use the tool."}, - { - "role": "assistant", - "tool_calls": [ - { - "id": "call_1", - "name": "search", - "arguments": '{"query": "abc"}', - } - ], - }, - ] - - -def test_replay_taskset_rejects_invalid_messages() -> None: - with pytest.raises(TypeError, match="messages must be a list"): - replay_task_record({"messages": "not a transcript"}) - - with pytest.raises(ValueError, match="Unknown role"): - replay_task_record({"messages": [{"role": "assistantish", "content": "no"}]}) diff --git a/tests/test_v1_rlm_swe.py b/tests/test_v1_rlm_swe.py deleted file mode 100644 index b7b536c3c0..0000000000 --- a/tests/test_v1_rlm_swe.py +++ /dev/null @@ -1,855 +0,0 @@ -import base64 -import io -import inspect -import sys -import tarfile -import types -from pathlib import Path -from typing import cast - -import pytest -from datasets import Dataset -from pydantic import BaseModel -from verifiers.types import Tool - -import verifiers as vf -from environments.rlm_swe_v1 import rlm_swe_v1 -from harnesses import RLM, RLMConfig, RLMProgramConfig -from harnesses.rlm import ( - DEFAULT_RLM_TOOL_SKILL_MARKER, - DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH, - DEFAULT_RLM_TOOL_SKILLS_MANIFEST_NAME, -) -from harnesses.utils.rlm_utils import rlm_tool_skills_archive -from harnesses.utils.rlm_utils import rlm_skills_dir -from verifiers.v1.utils.program_utils import merge_task_program, merge_task_sandbox - - -def as_dict(value: object) -> dict[str, object]: - if isinstance(value, vf.ProgramConfig): - value = value.data() - elif isinstance(value, BaseModel): - value = value.model_dump(exclude_none=True) - assert isinstance(value, dict) - return cast(dict[str, object], value) - - -def load_order_task(split: vf.TaskSplit = "train") -> vf.Tasks: - _ = split - return [{"prompt": [{"role": "user", "content": "Find order A-1."}]}] - - -class OrderTasksetConfig(vf.TasksetConfig): - pass - - -class OrderTaskset(vf.Taskset[OrderTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return load_order_task(split) - - -def tool_skills_archive(harness: vf.Harness, state: vf.State) -> bytes: - if "task" not in state: - task = vf.Task({"prompt": []}).freeze() - state = vf.State.for_task(task) - harness.runtime.prepare_state(task, state) - return base64.b64decode(rlm_tool_skills_archive(state, harness.runtime)) - - -def test_rlm_harness_builds_sandbox_program_without_eager_checkout(): - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig(local_checkout="/tmp/does-not-need-to-exist-yet") - ) - ) - program = as_dict(harness.config.program) - program_env = as_dict(program["env"]) - artifacts = as_dict(program["artifacts"]) - setup = cast(list[str], program["setup"]) - - assert isinstance(harness, vf.Harness) - assert program["sandbox"] is not False - assert isinstance(setup, list) - assert "apt-get -o Acquire::Retries=3 update" in setup[0] - assert "apt-get -o Acquire::Retries=3 install" in setup[0] - assert "RLM_MODEL" in program_env - assert "rlm_metrics" in artifacts - - -def test_rlm_harness_accepts_typed_config_surface(): - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig( - local_checkout="/tmp/checkout", - tools=["bash", "edit"], - exec_timeout=11, - env_vars={"CUSTOM": "1"}, - ) - ) - ) - program = as_dict(harness.config.program) - program_env = as_dict(program["env"]) - - assert harness.config.program.tools == ["bash", "edit"] - assert program_env["RLM_TOOLS"] == "bash,edit" - assert program_env["RLM_EXEC_TIMEOUT"] == "11" - assert program_env["CUSTOM"] == "1" - - -def test_rlm_endpoint_hides_nested_depth_requests(): - harness = RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ) - - assert harness.endpoint.trajectory_visibility({"x-rlm-depth": "0"}) == "append" - assert harness.endpoint.trajectory_visibility({"x-rlm-depth": "1"}) == "hidden" - assert ( - harness.endpoint.trajectory_visibility( - {"x-rlm-depth": "0", "x-verifiers-trajectory": "hidden"} - ) - == "hidden" - ) - - -def test_rlm_harness_preserves_program_setup_timeout_override(): - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig( - local_checkout="/tmp/checkout", - setup_timeout=123, - ), - ) - ) - program = as_dict(harness.config.program) - - assert program["setup_timeout"] == 123 - - -def test_rlm_harness_uses_sandbox_setup_timeout_default(): - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig( - local_checkout="/tmp/checkout", - sandbox=vf.SandboxConfig(setup_timeout=777), - ), - ) - ) - program = as_dict(harness.config.program) - - assert program["setup_timeout"] == 777 - - -def test_rlm_harness_keeps_minimum_setup_timeout_for_default_sandbox_config(): - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig( - local_checkout="/tmp/checkout", - sandbox=vf.SandboxConfig(), - ), - ) - ) - program = as_dict(harness.config.program) - sandbox = as_dict(harness.sandbox) - - assert program["setup_timeout"] == 600 - assert sandbox["setup_timeout"] == 600 - - -def test_rlm_harness_can_upload_skills(tmp_path: Path): - skills = tmp_path / "skills" - (skills / "edit").mkdir(parents=True) - (skills / "edit" / "SKILL.md").write_text("---\nname: edit\n---\n") - - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig(local_checkout="/tmp/checkout", skills=str(skills)) - ) - ) - program = as_dict(harness.config.program) - dirs = as_dict(program["dirs"]) - files = as_dict(program["files"]) - setup = cast(list[str], program["setup"]) - - assert dirs["/task/rlm-skills"] == str(skills) - assert files[DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH] == { - "fn": "harnesses.utils.rlm_utils:rlm_tool_skills_archive" - } - assert isinstance(setup, list) - assert DEFAULT_RLM_TOOL_SKILLS_ARCHIVE_PATH in setup[1] - assert DEFAULT_RLM_TOOL_SKILLS_MANIFEST_NAME in setup[1] - assert DEFAULT_RLM_TOOL_SKILL_MARKER in setup[1] - assert "rm -rf" in setup[1] - assert "tar -tzf" in setup[1] - - -def test_rlm_harness_uploads_taskset_skills_by_default(tmp_path: Path): - skills = tmp_path / "taskset-skills" - skills.mkdir() - (skills / "SKILL.md").write_text("---\nname: taskset\n---\n") - - class SkillTaskset(vf.Taskset): - def get_upload_dirs(self): - return {"skills": skills} - - env = vf.Env( - taskset=SkillTaskset(config=vf.TasksetConfig()), - harness=RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ), - ) - program = as_dict(env.harness.config.program) - dirs = as_dict(program["dirs"]) - - assert dirs["/task/rlm-skills"] == { - "fn": "harnesses.utils.rlm_utils:rlm_skills_dir" - } - assert rlm_skills_dir(vf.State({}), env.harness.runtime) == skills - - -def test_rlm_harness_recomputes_taskset_skills(tmp_path: Path): - first_skills = tmp_path / "first-skills" - second_skills = tmp_path / "second-skills" - first_skills.mkdir() - second_skills.mkdir() - - class SkillTasksetConfig(vf.TasksetConfig): - skills_path: str - - class SkillTaskset(vf.Taskset[SkillTasksetConfig]): - def get_upload_dirs(self): - return {"skills": Path(self.config.skills_path)} - - class NoSkillTaskset(vf.Taskset): - def get_upload_dirs(self): - return {} - - harness = RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ) - vf.Env( - taskset=SkillTaskset(config=SkillTasksetConfig(skills_path=str(first_skills))), - harness=harness, - ) - vf.Env( - taskset=SkillTaskset(config=SkillTasksetConfig(skills_path=str(second_skills))), - harness=harness, - ) - program = as_dict(harness.config.program) - dirs = as_dict(program["dirs"]) - - assert dirs["/task/rlm-skills"] == { - "fn": "harnesses.utils.rlm_utils:rlm_skills_dir" - } - assert rlm_skills_dir(vf.State({}), harness.runtime) == second_skills - - vf.Env(taskset=NoSkillTaskset(config=vf.TasksetConfig()), harness=harness) - - assert rlm_skills_dir(vf.State({}), harness.runtime) is None - - -@pytest.mark.asyncio -async def test_rlm_harness_generates_skills_for_v1_tools(): - async def lookup_order(order_id: str) -> str: - """Look up an order by ID.""" - return f"order:{order_id}" - - taskset = OrderTaskset(config=OrderTasksetConfig()) - taskset.add_toolset(vf.Toolset(tools=[lookup_order])) - env = vf.Env( - taskset=taskset, - harness=RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ), - ) - task = next(iter(env.taskset)) - state = vf.State.for_task(task) - - env.harness.runtime.prepare_state(task, state) - archive = tool_skills_archive(env.harness, state) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - source = ( - tar.extractfile("lookup_order/src/lookup_order/lookup_order.py") - .read() - .decode() - ) - skill_markdown = tar.extractfile("lookup_order/SKILL.md").read().decode() - marker = ( - tar.extractfile(f"lookup_order/{DEFAULT_RLM_TOOL_SKILL_MARKER}") - .read() - .decode() - ) - assert "async def run(order_id: str, **kwargs) -> object" in source - assert "/vf/tools/" in source - assert "dill.loads" not in source - assert "Look up an order by ID." in skill_markdown - assert "result = await lookup_order" in skill_markdown - assert marker == "1\n" - - -@pytest.mark.asyncio -async def test_vf_tool_skill_falls_back_for_runtime_bound_tools(): - async def stateful_lookup(order_id: str, state: vf.State) -> str: - """Look up an order with rollout state.""" - return f"{state['tenant']}:{order_id}" - - taskset = OrderTaskset(config=OrderTasksetConfig()) - taskset.add_toolset(vf.Toolset(tools=[stateful_lookup])) - env = vf.Env( - taskset=taskset, - harness=RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ), - ) - task = next(iter(env.taskset)) - state = vf.State.for_task(task) - - env.harness.runtime.prepare_state(task, state) - archive = tool_skills_archive(env.harness, state) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - source = ( - tar.extractfile("stateful_lookup/src/stateful_lookup/stateful_lookup.py") - .read() - .decode() - ) - - assert "/vf/tools/" in source - assert "dill.loads" not in source - - -def test_rlm_tool_skills_archive_avoids_base_skill_name_collisions(tmp_path: Path): - skills = tmp_path / "skills" - (skills / "lookup_order").mkdir(parents=True) - harness = RLM( - config=RLMConfig( - program=RLMProgramConfig(local_checkout="/tmp/checkout", skills=str(skills)) - ) - ) - tool_def = Tool( - name="lookup_order", - description="Look up an order.", - parameters={"type": "object", "properties": {}}, - ) - setattr(harness.runtime, "tool_defs", lambda state: [tool_def]) - - archive = tool_skills_archive(harness, vf.State({})) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - names = tar.getnames() - source = ( - tar.extractfile("lookup_order_2/src/lookup_order_2/lookup_order_2.py") - .read() - .decode() - ) - - assert "lookup_order_2/SKILL.md" in names - assert "/vf/tools/" in source - assert "'lookup_order'" in source - - -def test_vf_tool_skill_uses_arguments_dict_for_tool_parameters(): - harness = RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ) - tool_def = Tool( - name="reserved_param", - description="Reserved parameter.", - parameters={ - "type": "object", - "properties": { - "_call_vf_tool": {"type": "string"}, - "limit": {"type": "integer", "default": 10}, - }, - "required": ["_call_vf_tool"], - }, - ) - setattr(harness.runtime, "tool_defs", lambda state: [tool_def]) - archive = tool_skills_archive(harness, vf.State({})) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - source = ( - tar.extractfile("reserved_param/src/reserved_param/reserved_param.py") - .read() - .decode() - ) - - assert "async def run(arguments: dict | None = None, **kwargs) -> object" in source - assert "arguments = {**(arguments or {}), **kwargs}" in source - assert 'json={"arguments": arguments}' in source - assert "def _tool_arguments" not in source - assert "limit=None" not in source - - -@pytest.mark.asyncio -async def test_vf_tool_skill_filters_extra_kwargs_for_closed_schemas( - monkeypatch: pytest.MonkeyPatch, -) -> None: - harness = RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ) - tool_def = Tool( - name="list_events", - description="List calendar events.", - parameters={ - "type": "object", - "properties": {"date": {"type": "string"}}, - "required": ["date"], - "additionalProperties": False, - }, - ) - setattr(harness.runtime, "tool_defs", lambda state: [tool_def]) - archive = tool_skills_archive(harness, vf.State({})) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - source = ( - tar.extractfile("list_events/src/list_events/list_events.py") - .read() - .decode() - ) - - module = types.ModuleType("list_events") - exec(source, module.__dict__) - calls: list[dict[str, object]] = [] - - class Response: - content = b"{}" - - def json(self): - return {"result": "ok"} - - def raise_for_status(self): - return None - - class Requests: - @staticmethod - def post(url, json, headers, timeout): - calls.append(json) - return Response() - - monkeypatch.setenv("OPENAI_BASE_URL", "https://example.test/v1") - module.requests = Requests - - assert list(inspect.signature(module.run).parameters) == ["date", "kwargs"] - result = await module.run(user="Alice", date="2025-07-14") - - assert result == "ok" - assert calls == [{"arguments": {"date": "2025-07-14"}}] - - -@pytest.mark.asyncio -async def test_vf_tool_skill_omits_unset_optional_arguments( - monkeypatch: pytest.MonkeyPatch, -) -> None: - harness = RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ) - tool_def = Tool( - name="search", - description="Search documents.", - parameters={ - "type": "object", - "properties": { - "query": {"type": "string"}, - "limit": {"type": "integer"}, - }, - "required": ["query"], - }, - ) - setattr(harness.runtime, "tool_defs", lambda state: [tool_def]) - archive = tool_skills_archive(harness, vf.State({})) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - source = tar.extractfile("search/src/search/search.py").read().decode() - - module = types.ModuleType("search") - exec(source, module.__dict__) - calls: list[dict[str, object]] = [] - - class Response: - content = b"{}" - - def json(self): - return {"result": "ok"} - - def raise_for_status(self): - return None - - class Requests: - @staticmethod - def post(url, json, headers, timeout): - calls.append(json) - return Response() - - monkeypatch.setenv("OPENAI_BASE_URL", "https://example.test/v1") - module.requests = Requests - - assert ( - "async def run(query: str, limit: int | None = None, **kwargs) -> object" - in source - ) - assert "if limit is not None:" in source - assert "limit=None" not in source - result = await module.run(query="docs") - - assert result == "ok" - assert calls == [{"arguments": {"query": "docs"}}] - - -@pytest.mark.asyncio -async def test_vf_tool_skill_surfaces_verifier_tool_errors( - monkeypatch: pytest.MonkeyPatch, -) -> None: - harness = RLM( - config=RLMConfig(program=RLMProgramConfig(local_checkout="/tmp/checkout")) - ) - tool_def = Tool( - name="list_events", - description="List calendar events.", - parameters={ - "type": "object", - "properties": {"date": {"type": "string"}}, - "required": ["date"], - "additionalProperties": False, - }, - ) - setattr(harness.runtime, "tool_defs", lambda state: [tool_def]) - archive = tool_skills_archive(harness, vf.State({})) - - with tarfile.open(fileobj=io.BytesIO(archive), mode="r:gz") as tar: - source = ( - tar.extractfile("list_events/src/list_events/list_events.py") - .read() - .decode() - ) - - module = types.ModuleType("list_events") - exec(source, module.__dict__) - - class Response: - content = b"{}" - - def json(self): - return {"error": "unexpected keyword argument 'user'"} - - def raise_for_status(self): - raise AssertionError("JSON tool errors should be raised first") - - class Requests: - @staticmethod - def post(url, json, headers, timeout): - return Response() - - monkeypatch.setenv("OPENAI_BASE_URL", "https://example.test/v1") - module.requests = Requests - - with pytest.raises(RuntimeError, match="unexpected keyword argument"): - await module.run(date="2025-07-14") - - -def test_taskset_discovers_sibling_skills_dir_by_default( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - module_name = "skill_taskset_module" - module_file = tmp_path / f"{module_name}.py" - skills = tmp_path / "skills" - module_file.write_text("") - skills.mkdir() - (skills / "SKILL.md").write_text("---\nname: sibling\n---\n") - module = types.ModuleType(module_name) - module.__file__ = str(module_file) - module.__package__ = "" - monkeypatch.setitem(sys.modules, module_name, module) - skill_taskset_type = type( - "SkillTaskset", (vf.Taskset,), {"__module__": module_name} - ) - - taskset = skill_taskset_type(config=vf.TasksetConfig()) - - assert taskset.get_upload_dirs() == {"skills": skills} - - -def test_rlm_harness_explicit_skills_override_taskset_skills(tmp_path: Path): - taskset_skills = tmp_path / "taskset-skills" - explicit_skills = tmp_path / "explicit-skills" - taskset_skills.mkdir() - explicit_skills.mkdir() - - class SkillTaskset(vf.Taskset): - def get_upload_dirs(self): - return {"skills": taskset_skills} - - env = vf.Env( - taskset=SkillTaskset(config=vf.TasksetConfig()), - harness=RLM( - config=RLMConfig( - program=RLMProgramConfig( - local_checkout="/tmp/checkout", - skills=str(explicit_skills), - ) - ) - ), - ) - program = as_dict(env.harness.config.program) - dirs = as_dict(program["dirs"]) - - assert dirs["/task/rlm-skills"] == str(explicit_skills) - - -def test_rlm_swe_environment_uses_v1_r2e_taskset(monkeypatch): - calls: dict[str, object] = {} - - def fake_load_dataset(dataset_name: str, **kwargs: object) -> Dataset: - calls["dataset_name"] = dataset_name - calls["kwargs"] = kwargs - return fake_r2e_dataset() - - monkeypatch.setattr(rlm_swe_v1, "load_dataset", fake_load_dataset) - - env = rlm_swe_v1.load_environment( - config=rlm_swe_v1.RlmSweEnvConfig( - taskset=rlm_swe_v1.RlmSweTasksetConfig( - dataset_name="fake-r2e", - repo_path="/workspace/repo", - timeout_minutes=30, - env={"CUSTOM": "1"}, - ), - harness=rlm_swe_v1.RlmSweHarnessConfig( - program=rlm_swe_v1.RlmSweProgramConfig( - local_checkout="/tmp/checkout", - env_vars={"CALLER": "1"}, - ) - ), - ), - ) - task = next(iter(env.taskset)) - program = as_dict(env.harness.config.program) - program_env = as_dict(program["env"]) - merged_program = merge_task_program(program, task, kind="command") - merged_env = as_dict(merged_program["env"]) - assert env.harness.sandbox is not None - merged_sandbox = merge_task_sandbox(env.harness.sandbox, task).data() - - assert isinstance(env, vf.Env) - assert isinstance(env.taskset, rlm_swe_v1.R2ESWETaskset) - assert isinstance(env.harness, RLM) - assert calls["dataset_name"] == "fake-r2e" - assert task["taskset_id"] == "swe/r2e" - assert task["instruction"] == "Fix repo-0." - assert task["sandbox"]["image"] == ( - f"{rlm_swe_v1.REGISTRY_PREFIX}/r2e/image:latest" - ) - assert task["sandbox"]["workdir"] == "/workspace/repo" - assert task["sandbox"]["timeout_minutes"] == 30 - task_program_env = as_dict(as_dict(task["program"])["env"]) - assert task_program_env["AGENT_WORKDIR"] == "/workspace/repo" - assert "/workspace/repo/.venv/bin" in task_program_env["AGENT_PATH"] - assert task_program_env["PAGER"] == "cat" - assert task_program_env["CUSTOM"] == "1" - assert "CUSTOM" not in program_env - assert program_env["CALLER"] == "1" - assert program_env["RLM_TOOLS"] == "bash,edit" - assert merged_sandbox["workdir"] == "/workspace/repo" - assert merged_env["AGENT_WORKDIR"] == "/workspace/repo" - assert "/workspace/repo/.venv/bin" in merged_env["AGENT_PATH"] - assert merged_env["PAGER"] == "cat" - assert merged_env["CUSTOM"] == "1" - assert merged_env["CALLER"] == "1" - - -def test_rlm_swe_taskset_hooks_are_registered_with_runtime(): - taskset = rlm_swe_v1.load_taskset(config=rlm_swe_v1.RlmSweTasksetConfig()) - env = vf.Env(taskset=taskset) - - setup_names = [handler.__name__ for handler in env.harness.runtime.rollout_setup] - cleanup_names = [ - handler.__name__ for handler in env.harness.runtime.rollout_cleanup - ] - signal_names = {signal["name"] for signal in env.harness.runtime.rollout_signals} - - assert setup_names.count("setup_r2e_sandbox") == 1 - assert cleanup_names.count("cleanup_r2e_state") == 1 - assert "solved" in signal_names - - -@pytest.mark.asyncio -async def test_rlm_swe_taskset_setup_and_reward(monkeypatch): - monkeypatch.setattr( - rlm_swe_v1, "load_dataset", lambda *args, **kwargs: fake_r2e_dataset() - ) - taskset = rlm_swe_v1.load_taskset( - config=rlm_swe_v1.RlmSweTasksetConfig(timeout_minutes=30) - ) - task = next(iter(taskset)) - state = vf.State.for_task(task) - sandbox = FakeSandbox() - calls: dict[str, object] = {} - - async def fake_setup_sandbox(sandbox_arg: object, state_arg: vf.State) -> None: - calls["setup_sandbox"] = sandbox_arg - calls["setup_state"] = state_arg - - async def fake_run_tests( - sandbox_arg: object, - state_arg: vf.State, - test_timeout: int, - ) -> str: - calls["run_tests"] = (sandbox_arg, state_arg, test_timeout) - return """ -=========================== short test summary info ============================ -PASSED tests/test_example.py::test_fix -""" - - monkeypatch.setattr(taskset, "setup_sandbox", fake_setup_sandbox) - monkeypatch.setattr(taskset, "run_tests", fake_run_tests) - - await taskset.setup_r2e_sandbox(task, state, sandbox=sandbox) - reward = await taskset.solved(task, state) - await taskset.cleanup_r2e_state(task, state) - - assert calls["setup_sandbox"] is sandbox - assert calls["setup_state"] is state - assert calls["run_tests"] == (sandbox, state, 1800) - assert state["sandbox_id"] == "sandbox-1" - assert state["test_timeout"] == 1800 - assert reward == 1.0 - assert "sandbox_client" not in state - assert "_rlm_swe_sandbox" not in state - - -@pytest.mark.asyncio -async def test_rlm_swe_run_tests_quotes_env_values(): - taskset = rlm_swe_v1.load_taskset( - config=rlm_swe_v1.RlmSweTasksetConfig( - hide_tests_from_agent=False, - env={"SAFE": "two words; $(echo nope)", "QUOTE": "it's ok"}, - ) - ) - sandbox = RecordingSandbox() - - output = await taskset.run_tests(sandbox, {}, 123) - - assert output == "test output" - assert len(sandbox.background_jobs) == 1 - command = sandbox.background_jobs[0]["command"] - assert "SAFE='two words; $(echo nope)'" in command - assert "QUOTE='it'\"'\"'s ok'" in command - assert command.endswith("/bin/bash run_tests.sh > test_output.txt 2>&1") - - -def test_rlm_swe_get_env_vars_uses_configured_repo_path(): - taskset = rlm_swe_v1.load_taskset( - config=rlm_swe_v1.RlmSweTasksetConfig(repo_path="/workspace/repo") - ) - - path = taskset.get_env_vars()["PATH"] - - assert "/workspace/repo/.venv/bin" in path - assert "/testbed/.venv/bin" not in path - - -def test_rlm_swe_reward_rejects_pytest_summary_without_nodeid(): - taskset = rlm_swe_v1.load_taskset(config=rlm_swe_v1.RlmSweTasksetConfig()) - test_output = """ -=========================== short test summary info ============================ -PASSED tests/test_example.py -""" - - reward = taskset.calculate_reward( - test_output, - {"expected_output_json": '{"test_fix": "PASSED"}'}, - ) - parsed = rlm_swe_v1.parse_log_pytest(test_output) - - assert reward == 0.0 - assert "" not in parsed - - -def test_rlm_swe_parse_log_pytest_uses_leading_status_token(): - test_output = """ -=========================== short test summary info ============================ -FAILED tests/test_PASSED_handler.py::test_fix - AssertionError -PASSED tests/test_failed_handler.py::test_other -ERROR tests/test_failed_handler.py::test_error - setup failed -""" - - parsed = rlm_swe_v1.parse_log_pytest(test_output) - - assert parsed == { - "test_fix": "FAILED", - "test_other": "PASSED", - "test_error": "ERROR", - } - - -def fake_r2e_dataset() -> Dataset: - return Dataset.from_list( - [ - { - "commit_hash": f"commit-{index}", - "repo_name": "example/repo", - "problem_statement": f"Fix repo-{index}.", - "docker_image": "r2e/image:latest", - "expected_output_json": '{"test_fix": "PASSED"}', - "parsed_commit_content": '{"file_diffs": []}', - } - for index in range(12) - ] - ) - - -class FakeLease: - client = object() - - -class FakeSandbox: - id = "sandbox-1" - lease = FakeLease() - - -class FakeCommandResult: - def __init__( - self, - stdout: str = "", - stderr: str = "", - exit_code: int = 0, - ): - self.stdout = stdout - self.stderr = stderr - self.exit_code = exit_code - - -class RecordingSandbox: - def __init__(self): - self.background_jobs: list[dict[str, object]] = [] - self.commands: list[dict[str, object]] = [] - - async def run_background_job( - self, - command: str, - timeout: int | None = None, - working_dir: str | None = None, - ) -> FakeCommandResult: - self.background_jobs.append( - { - "command": command, - "timeout": timeout, - "working_dir": working_dir, - } - ) - return FakeCommandResult() - - async def execute( - self, - command: str, - timeout: int | None = None, - working_dir: str | None = None, - ) -> FakeCommandResult: - self.commands.append( - { - "command": command, - "timeout": timeout, - "working_dir": working_dir, - } - ) - return FakeCommandResult(stdout="test output") diff --git a/tests/test_v1_runtime_lifecycle.py b/tests/test_v1_runtime_lifecycle.py deleted file mode 100644 index 9a3fc71545..0000000000 --- a/tests/test_v1_runtime_lifecycle.py +++ /dev/null @@ -1,3683 +0,0 @@ -import asyncio -import json -import os -import shlex -import sys -import tempfile -import threading -import time -import urllib.request -from contextlib import AsyncExitStack -from pathlib import Path -from types import ModuleType, SimpleNamespace -from typing import Any, cast - -import pytest -from openai import OpenAI -from pydantic import BaseModel - -import verifiers as vf -from verifiers.clients import Client -from verifiers.types import ClientConfig, Messages -from verifiers.types import Response, ResponseMessage, ToolCall -from verifiers.types import Tool -from verifiers.types import Usage -from verifiers.v1.runtime import Runtime -from verifiers.v1.utils import mcp_utils, sandbox_utils -from verifiers.v1.utils.mcp_proxy_utils import MCP_PROXY_CONFIG_PATH, MCP_PROXY_PATH -from verifiers.v1.utils.mcp_proxy_utils import proxy_command, proxy_program -from verifiers.v1.utils.program_utils import command_env -from verifiers.v1.utils.runtime_registry import load_runtime -from verifiers.v1.utils.sandbox_python_utils import ( - SANDBOX_PYTHON, - SANDBOX_UV, - python_package_install_command, -) -from verifiers.v1.utils.sandbox_program_utils import ( - PACKAGE_ROOT, - RUNNER_CONFIG_PATH, - SandboxPackage, - TOOL_DEFS_BY_PROTOCOL_PATH, - TOOL_DEFS_PATH, - apply_internal_state_patch, - runner_source, - sandbox_program_package, - sandbox_runner_program, -) -from verifiers.v1.utils.sandbox_utils import ( - VF_STATE_INPUT_PATH_KEY, - run_sandbox_command, - upload_program_dirs, -) - -PROGRAM_REF_MODULE = "v1_runtime_lifecycle_refs" - - -class FakeMCPHandle: - def __init__(self, name: str): - self.name = name - self.tool_def = Tool( - name=name, - description="fake", - parameters={"type": "object", "properties": {}}, - ) - - async def __call__(self) -> str: - return "ok" - - -class FakeClient: - def __init__(self): - self.closed = False - - async def close(self) -> None: - self.closed = True - - -class FakeModelClient: - def __init__(self, responses: list[Response]): - self.responses = responses - - async def get_response(self, **kwargs: object) -> Response: - _ = kwargs - if not self.responses: - raise AssertionError("No fake model responses left.") - return self.responses.pop(0) - - -class CapturingModelClient(FakeModelClient): - def __init__(self, responses: list[Response]): - super().__init__(responses) - self.requests: list[dict[str, object]] = [] - - async def get_response(self, **kwargs: object) -> Response: - prompt = kwargs.get("prompt") - if isinstance(prompt, list): - kwargs["prompt"] = list(prompt) - self.requests.append(dict(kwargs)) - return await super().get_response(**kwargs) - - -class BlockingModelClient(CapturingModelClient): - def __init__(self, responses: list[Response]): - super().__init__(responses) - self.entered = asyncio.Event() - self.release = asyncio.Event() - - async def get_response(self, **kwargs: object) -> Response: - self.entered.set() - await self.release.wait() - return await super().get_response(**kwargs) - - -class RaisingModelClient: - def __init__(self, error: vf.Error): - self.error = error - - async def get_response(self, **kwargs: object) -> Response: - _ = kwargs - raise self.error - - -class FakeCreateSandboxRequest: - def __init__(self, **kwargs: object): - self.kwargs = kwargs - - -class FakeAPIError(Exception): - pass - - -class FakeUploadTimeoutError(Exception): - pass - - -class FakeSandboxResult: - def __init__(self, sandbox_id: str): - self.id = sandbox_id - - -class FakeCommandResult: - exit_code = 0 - stdout = "ok\n" - stderr = "" - - -class FakeSandboxClient: - created: list[str] = [] - requests: list[dict[str, object]] = [] - deleted: list[str] = [] - commands: list[tuple[str, str]] = [] - command_timeouts: list[int | None] = [] - background_jobs: list[tuple[str, str, int | None, str | None, int]] = [] - uploads: list[tuple[str, str, bytes]] = [] - wait_attempts: list[tuple[str, int]] = [] - closed = 0 - - def __init__(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - - @classmethod - def reset(cls) -> None: - cls.created = [] - cls.requests = [] - cls.deleted = [] - cls.commands = [] - cls.command_timeouts = [] - cls.background_jobs = [] - cls.uploads = [] - cls.wait_attempts = [] - cls.closed = 0 - - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - type(self).requests.append(dict(request.kwargs)) - sandbox_id = f"sbx-{len(type(self).created) + 1}" - type(self).created.append(sandbox_id) - return FakeSandboxResult(sandbox_id) - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> None: - type(self).wait_attempts.append((sandbox_id, max_attempts)) - - async def execute_command( - self, *args: object, **kwargs: object - ) -> FakeCommandResult: - sandbox_id = str(kwargs.get("sandbox_id") or args[0]) - command = str(kwargs.get("command") or args[1]) - timeout = cast(int | None, kwargs.get("timeout")) - type(self).commands.append((sandbox_id, command)) - type(self).command_timeouts.append(timeout) - return FakeCommandResult() - - async def run_background_job( - self, *args: object, **kwargs: object - ) -> FakeCommandResult: - sandbox_id = str(kwargs.get("sandbox_id") or args[0]) - command = str(kwargs.get("command") or args[1]) - timeout = cast(int | None, kwargs.get("timeout", 900)) - working_dir = cast(str | None, kwargs.get("working_dir")) - poll_interval = cast(int, kwargs.get("poll_interval", 3)) - type(self).commands.append((sandbox_id, command)) - type(self).background_jobs.append( - (sandbox_id, command, timeout, working_dir, poll_interval) - ) - return FakeCommandResult() - - async def upload_bytes(self, *args: object, **kwargs: object) -> None: - sandbox_id = str(kwargs.get("sandbox_id") or args[0]) - path = str(kwargs.get("file_path") or kwargs.get("path") or args[1]) - data = cast(bytes, kwargs.get("file_bytes") or args[2]) - type(self).uploads.append((sandbox_id, path, data)) - - async def upload_file(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - - async def read_file(self, *args: object, **kwargs: object) -> str: - _ = args, kwargs - return "" - - async def delete(self, sandbox_id: str) -> None: - type(self).deleted.append(sandbox_id) - - async def aclose(self) -> None: - type(self).closed += 1 - - def teardown(self) -> None: - type(self).closed += 1 - - -async def echo_tool(query: str) -> str: - return f"echo:{query}" - - -async def borrowed_record_tool(value: str, state) -> str: - state.setdefault("borrowed_tool_values", []).append(value) - return f"recorded:{value}" - - -async def borrowed_stage_tool(value: str, state) -> str: - state.setdefault("borrowed_stage_values", []).append(value) - return f"stage:{value}" - - -async def named_tool(name: str) -> str: - return f"name:{name}" - - -async def failing_tool(section_id: str) -> str: - _ = section_id - raise ValueError("Invalid section_id format.") - - -async def finish_tool(answer: str, state) -> str: - state["answer"] = answer - state.stop("submitted") - return "submitted" - - -def fake_response( - content: str | None = None, - tool_calls: list[ToolCall] | None = None, - usage: Usage | None = None, -) -> Response: - return Response( - id="fake", - created=0, - model="fake", - usage=usage, - message=ResponseMessage( - role="assistant", - content=content, - tool_calls=tool_calls, - finish_reason="tool_calls" if tool_calls else "stop", - is_truncated=False, - ), - ) - - -async def program_sandbox_id(sandbox) -> str: - return sandbox.id - - -async def sandbox_lifecycle_setup(task, state, sandbox) -> None: - _ = task - state["setup_sandbox_id"] = sandbox.id - await sandbox.execute("echo lifecycle-setup") - - -async def state_input_setup(task, state) -> None: - _ = task - state["state_input_setup"] = True - - -@vf.setup(priority=150) -async def early_sandbox_lifecycle_setup(task, state, sandbox) -> None: - _ = task - state["early_setup_sandbox_id"] = sandbox.id - await sandbox.execute("echo early-lifecycle-setup") - - -def endpoint_config_binding(state): - return state.get_endpoint_config(api="chat") - - -def endpoint_config_binding_ref(state): - return state.get_endpoint_config(api="chat") - - -def configure_cli_endpoint(endpoint_config) -> str: - return f"echo model={endpoint_config.model} > /tmp/endpoint.txt" - - -def configure_cli_endpoint_ref(endpoint_config) -> str: - return f"echo ref-model={endpoint_config.model} > /tmp/ref_endpoint.txt" - - -def replay_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - _ = split - return [ - { - "prompt": [{"role": "user", "content": "Return the answer."}], - "answer": "solved", - } - ] - - -def setup_runtime_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - _ = split - return [{"prompt": [], "answer": "ready", "max_turns": 3}] - - -@vf.setup -async def initialize_from_taskset(task, state) -> None: - runtime = state.runtime_state() - sampling_args = {"top_p": 1.0} - sampling_args.update(dict(runtime.get("sampling_args") or {})) - runtime["sampling_args"] = sampling_args - state.setdefault("prompt", []).append( - {"role": "user", "content": f"task {task['answer']}"} - ) - - -ref_module = ModuleType(PROGRAM_REF_MODULE) -setattr(ref_module, "endpoint_config_binding_ref", endpoint_config_binding_ref) -setattr(ref_module, "configure_cli_endpoint_ref", configure_cli_endpoint_ref) -setattr(ref_module, "replay_tasks", replay_tasks) -setattr(ref_module, "setup_runtime_tasks", setup_runtime_tasks) -sys.modules[PROGRAM_REF_MODULE] = ref_module - - -def program_ref(name: str) -> str: - return f"{PROGRAM_REF_MODULE}:{name}" - - -def config_data(config: object | None) -> dict[str, object]: - if config is None: - return {} - if isinstance(config, BaseModel): - return config.model_dump(exclude_none=True, exclude_unset=True) - if isinstance(config, dict): - return dict(config) - raise TypeError("test config must be a mapping or config object") - - -def has_runtime_toolset(value: object) -> bool: - if isinstance(value, vf.Toolset): - return True - if isinstance(value, dict): - return any(has_runtime_toolset(item) for item in value.values()) - if isinstance(value, list | tuple): - return any(has_runtime_toolset(item) for item in value) - return False - - -def make_harness(config: object | None = None, **values: object) -> vf.Harness: - data = {**config_data(config), **values} - runtime_client = data.pop("client", None) - model_value = data.pop("model", None) - sampling_args = data.pop("sampling_args", None) - if model_value is not None or sampling_args is not None: - if model_value is None: - model_data: dict[str, object] = {} - elif isinstance(model_value, str): - model_data = {"name": model_value} - elif isinstance(model_value, vf.ModelConfig): - model_data = model_value.model_dump(exclude_none=True, exclude_unset=True) - elif isinstance(model_value, dict): - model_data = dict(model_value) - else: - raise TypeError("test harness model config must be a mapping.") - if sampling_args is not None: - model_data["sampling_args"] = sampling_args - data["model"] = model_data - runtime_toolsets = data.pop("toolsets", None) - if runtime_toolsets is not None and not has_runtime_toolset(runtime_toolsets): - data["toolsets"] = runtime_toolsets - runtime_toolsets = None - harness = vf.Harness(config=vf.HarnessConfig.model_validate(data)) - if runtime_client is not None: - harness.model_client = cast(Client, runtime_client) - if runtime_toolsets is not None: - harness.add_toolset(runtime_toolsets) - return harness - - -def make_taskset(config: object | None = None, **values: object) -> vf.Taskset: - data = {**config_data(config), **values} - runtime_toolsets = data.pop("toolsets", None) - if runtime_toolsets is not None and not has_runtime_toolset(runtime_toolsets): - data["toolsets"] = runtime_toolsets - runtime_toolsets = None - taskset = vf.Taskset(config=vf.TasksetConfig.model_validate(data)) - if runtime_toolsets is not None: - taskset.add_toolset(runtime_toolsets) - return taskset - - -async def child_reads_program_sandbox(task, state) -> dict[str, object]: - _ = task - tools = state.get_tools() - state["borrowed_sandbox_id"] = await tools["program_sandbox_id"]() - return state - - -def install_fake_sandboxes(monkeypatch: pytest.MonkeyPatch) -> None: - FakeSandboxClient.reset() - module = SimpleNamespace( - AsyncSandboxClient=FakeSandboxClient, - APIError=FakeAPIError, - CreateSandboxRequest=FakeCreateSandboxRequest, - UploadTimeoutError=FakeUploadTimeoutError, - ) - monkeypatch.setitem(sys.modules, "prime_sandboxes", module) - monkeypatch.setattr( - "verifiers.utils.threaded_sandbox_client.ThreadedAsyncSandboxClient", - FakeSandboxClient, - ) - - -def disable_sandbox_retry_sleep(monkeypatch: pytest.MonkeyPatch) -> None: - async def no_sleep(seconds: float, result: object | None = None) -> object | None: - _ = seconds - return result - - monkeypatch.setattr(sandbox_utils.asyncio, "sleep", no_sleep) - - -def install_fake_endpoint_tunnel(monkeypatch: pytest.MonkeyPatch) -> None: - async def get_tunnel_url(self) -> str: - _ = self - return "http://127.0.0.1:1" - - monkeypatch.setattr( - "verifiers.v1.utils.endpoint_utils.Endpoint.get_tunnel_url", - get_tunnel_url, - ) - - -class EndpointUserConfig(vf.UserConfig): - pass - - -class EndpointUser(vf.User[EndpointUserConfig]): - async def get_response( - self, task: dict[str, object], state: dict[str, object] - ) -> list[dict[str, str]]: - _ = task - state["user_seen"] = True - return [{"role": "user", "content": "continue"}] - - -async def endpoint_program(task, state): - _ = task - root = state["endpoint_root_url"].rstrip("/") - client = state.get_client(api="chat") - config = state.get_endpoint_config(api="responses") - endpoint_client = cast(OpenAI, state.get_client(api="responses", sync=True)) - auth_headers = {"Authorization": f"Bearer {endpoint_client.api_key}"} - endpoint_client.close() - - def get_json(url: str) -> dict[str, object]: - request = urllib.request.Request(url, headers=auth_headers) - with urllib.request.urlopen(request) as response: - return json.loads(response.read().decode()) - - def post_json(url: str, payload: dict[str, object]) -> dict[str, object]: - request = urllib.request.Request( - url, - data=json.dumps(payload).encode(), - headers={"content-type": "application/json", **auth_headers}, - ) - with urllib.request.urlopen(request) as response: - return json.loads(response.read().decode()) - - tools = await asyncio.to_thread(get_json, f"{root}/vf/tools") - openai_tools = await asyncio.to_thread( - get_json, f"{root}/vf/tools?protocol=openai_chat_completions" - ) - tool_payload: dict[str, object] = {"arguments": {"query": "hi"}} - tool_result = await asyncio.to_thread( - post_json, - f"{root}/vf/tools/echo_tool", - tool_payload, - ) - user_payload: dict[str, object] = { - "transcript": [{"role": "assistant", "content": "hello"}] - } - user_result = await asyncio.to_thread( - post_json, - f"{root}/vf/user", - user_payload, - ) - state["done"] = True - stop_result = await asyncio.to_thread(post_json, f"{root}/vf/stop", {}) - return { - "endpoint_tools": tools["tools"], - "endpoint_openai_tools": openai_tools["tools"], - "endpoint_tool_result": tool_result["result"], - "endpoint_user_messages": user_result["messages"], - "endpoint_stop": stop_result, - "endpoint_client_class": type(client).__name__, - "endpoint_config": config.model_dump(), - } - - -async def endpoint_model_error_program(task, state): - _ = task - root = state["endpoint_root_url"].rstrip("/") - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - auth_headers = {"Authorization": f"Bearer {endpoint_client.api_key}"} - endpoint_client.close() - - def post_model() -> None: - request = urllib.request.Request( - f"{root}/vf/model", - data=json.dumps( - {"messages": [{"role": "user", "content": "too long"}]} - ).encode(), - headers={"content-type": "application/json", **auth_headers}, - ) - with urllib.request.urlopen(request): - pass - - try: - await asyncio.to_thread(post_model) - except Exception as exc: - raise vf.SandboxError("Sandbox command failed") from exc - raise AssertionError("Expected /vf/model to fail") - - -async def endpoint_trajectory_program(task, state): - _ = task - root = state["endpoint_root_url"].rstrip("/") - config = state.get_endpoint_config(api="chat") - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - api_key = endpoint_client.api_key - endpoint_client.close() - - def post_chat(headers: dict[str, str]) -> dict[str, object]: - payload = { - "model": config.model, - "messages": [{"role": "user", "content": "hi"}], - } - request = urllib.request.Request( - f"{root}/v1/chat/completions", - data=json.dumps(payload).encode(), - headers={ - "content-type": "application/json", - "Authorization": f"Bearer {api_key}", - **headers, - }, - ) - with urllib.request.urlopen(request) as response: - return json.loads(response.read().decode()) - - hidden = await asyncio.to_thread(post_chat, {"x-verifiers-trajectory": "hidden"}) - shown = await asyncio.to_thread(post_chat, {}) - state["endpoint_hidden_response"] = hidden - state["endpoint_shown_response"] = shown - return state - - -async def concurrent_endpoint_program(task, state): - _ = task - root = state["endpoint_root_url"].rstrip("/") - config = state.get_endpoint_config(api="chat") - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - api_key = endpoint_client.api_key - endpoint_client.close() - - def post_chat(content: str) -> dict[str, object]: - payload = { - "model": config.model, - "messages": [{"role": "user", "content": content}], - } - request = urllib.request.Request( - f"{root}/v1/chat/completions", - data=json.dumps(payload).encode(), - headers={ - "content-type": "application/json", - "Authorization": f"Bearer {api_key}", - }, - ) - with urllib.request.urlopen(request) as response: - return json.loads(response.read().decode()) - - state["endpoint_concurrent_responses"] = await asyncio.gather( - asyncio.to_thread(post_chat, "first"), - asyncio.to_thread(post_chat, "second"), - ) - return state - - -async def mcp_proxy_program(task, state): - _ = task - from mcp import ClientSession, StdioServerParameters - from mcp.client.stdio import stdio_client - - tool_auth_var = str(state["endpoint_api_key_var"]) - proxy_config = proxy_program( - {}, - tool_base_url=f"{state['endpoint_root_url'].rstrip('/')}/vf/tools", - tool_auth_var=tool_auth_var, - ) - proxy_files = cast(dict[str, str], proxy_config["files"]) - with tempfile.NamedTemporaryFile("w", suffix=".py", delete=False) as f: - f.write(proxy_files[MCP_PROXY_PATH]) - proxy_path = Path(f.name) - config_path = proxy_path.with_suffix(".json") - config_path.write_text(proxy_files[MCP_PROXY_CONFIG_PATH]) - try: - endpoint_client = cast(OpenAI, state.get_client(api="chat", sync=True)) - tool_auth_token = endpoint_client.api_key - endpoint_client.close() - server = StdioServerParameters( - command=sys.executable, - args=[str(proxy_path), str(config_path)], - env={tool_auth_var: tool_auth_token}, - ) - async with stdio_client(server) as (read_stream, write_stream): - async with ClientSession(read_stream, write_stream) as session: - await session.initialize() - listed = await session.list_tools() - result = await session.call_tool("echo_tool", {"query": "hi"}) - return { - "mcp_tools": [tool.name for tool in listed.tools], - "mcp_result": mcp_utils.mcp_result_value(result), - } - finally: - proxy_path.unlink(missing_ok=True) - config_path.unlink(missing_ok=True) - - -async def child_program(task, state): - _ = task - return { - "child_runtime": dict(state["runtime"]), - "child_trajectory_id": state["trajectory_id"], - } - - -async def parent_program(task, state): - child = make_harness(program={"fn": program_ref("child_program")}) - child_task = vf.Task({"prompt": [{"role": "user", "content": "child"}]}).freeze() - child_state = state.for_task(child_task, borrow="model") - child_state = await child.run(child_task, child_state) - return {"child_state": child_state} - - -async def mark_submitted(task, state): - _ = task - state["submitted"] = True - return state - - -async def parent_calls_owned_child_program(task, state): - child = make_harness( - program={"fn": program_ref("child_program")}, - client=cast(Client, FakeClient()), - model="child-model", - ) - child_task = vf.Task({"prompt": [{"role": "user", "content": "child"}]}).freeze() - child_state = await child.run(child_task) - return {"child_state": child_state} - - -async def update_summary_with_resolved_handles(task, state): - _ = task - child = make_harness(system_prompt="Summarize the parent rollout in one word.") - child_task = vf.Task( - {"prompt": [{"role": "user", "content": str(state["completion"])}]} - ).freeze() - child_state = state.for_task( - child_task, - borrow="model", - transcript="append", - ) - child_state = await child.run(child_task, child_state) - state["summary"] = child_state["completion"][0]["content"] - state["child_trajectory_id"] = child_state["trajectory_id"] - - -async def update_child_uses_borrowed_tool(task, state): - _ = task - child = make_harness(max_turns=2) - child_task = vf.Task({"prompt": [{"role": "user", "content": "inspect"}]}).freeze() - child_state = state.for_task( - child_task, - borrow="model", - tools="borrowed_record_tool", - transcript="append", - ) - child_state = await child.run(child_task, child_state) - state["child_completion"] = child_state["completion"][-1]["content"] - state["child_trajectory_id"] = child_state["trajectory_id"] - - -async def submitted(task, state) -> bool: - _ = task - return bool(state.get("submitted")) - - -async def state_tools_program(task, state): - _ = task - tools = state.get_tools() - state["tool_result"] = await tools["echo_tool"](query="state") - state["tool_name"] = tools["echo_tool"].__name__ - return state - - -async def state_tool_program(task, state): - _ = task - tools = state.get_tools() - state["tool_result"] = await tools["echo_tool"](query="injected") - state["tool_name"] = tools["echo_tool"].__name__ - state["tool_docs"] = tools["echo_tool"].__doc__ - return state - - -async def replay_answer_program(task, state): - state["answer"] = task["answer"] - return state - - -@vf.reward -async def replay_reward(task, state) -> float: - return float(state.get("answer") == task.get("answer")) - - -for _name, _value in { - "sandbox_lifecycle_setup": sandbox_lifecycle_setup, - "state_input_setup": state_input_setup, - "early_sandbox_lifecycle_setup": early_sandbox_lifecycle_setup, - "endpoint_config_binding": endpoint_config_binding, - "configure_cli_endpoint": configure_cli_endpoint, - "initialize_from_taskset": initialize_from_taskset, - "child_reads_program_sandbox": child_reads_program_sandbox, - "endpoint_program": endpoint_program, - "endpoint_model_error_program": endpoint_model_error_program, - "endpoint_trajectory_program": endpoint_trajectory_program, - "concurrent_endpoint_program": concurrent_endpoint_program, - "mcp_proxy_program": mcp_proxy_program, - "child_program": child_program, - "parent_program": parent_program, - "mark_submitted": mark_submitted, - "parent_calls_owned_child_program": parent_calls_owned_child_program, - "update_summary_with_resolved_handles": update_summary_with_resolved_handles, - "update_child_uses_borrowed_tool": update_child_uses_borrowed_tool, - "submitted": submitted, - "state_tools_program": state_tools_program, - "state_tool_program": state_tool_program, - "replay_answer_program": replay_answer_program, - "replay_reward": replay_reward, -}.items(): - setattr(ref_module, _name, _value) - - -def test_model_client_default_keys_are_rollout_local() -> None: - runtime = Runtime() - client = FakeClient() - state_a = vf.State.for_task(vf.Task({}).freeze()) - state_b = vf.State.for_task(vf.Task({}).freeze()) - - runtime.bind_model_client(state_a, cast(Client, client)) - runtime.bind_model_client(state_b, cast(Client, client)) - - assert state_a["runtime"]["client_key"] != state_b["runtime"]["client_key"] - assert len(runtime.model_clients) == 2 - - -@pytest.mark.asyncio -async def test_v1_records_default_metrics_usage_and_timing() -> None: - usage = Usage( - prompt_tokens=11, - reasoning_tokens=0, - completion_tokens=7, - total_tokens=18, - ) - harness = make_harness( - client=cast( - Client, - FakeModelClient([fake_response(content="ok", usage=usage)]), - ), - model="fake-model", - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["metrics"]["num_turns"] == 1.0 - assert state["token_usage"] == { - "input_tokens": 11.0, - "output_tokens": 7.0, - "final_output_tokens": 7.0, - "final_input_tokens": 11.0, - } - assert state["usage"] == state["token_usage"] - assert state["timing"]["total"] > 0.0 - assert state["timing"]["generation"]["duration"] > 0.0 - assert state["timing"]["model"]["duration"] > 0.0 - - -def test_v1_state_does_not_copy_task_answer_to_top_level() -> None: - task = vf.Task({"answer": "gold"}).freeze() - state = vf.State.for_task(task) - - assert "answer" not in state - assert state["task"]["answer"] == "gold" - - -@pytest.mark.asyncio -async def test_endpoint_exposes_tool_user_and_stop_surfaces() -> None: - harness = make_harness( - program={"fn": program_ref("endpoint_program")}, - model="test-model", - toolsets=[vf.Toolset(tools=[echo_tool])], - user=EndpointUserConfig(), - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - await harness.teardown() - - assert [tool["name"] for tool in state["endpoint_tools"]] == ["echo_tool"] - openai_tool = state["endpoint_openai_tools"][0] - assert openai_tool["type"] == "function" - assert openai_tool["function"]["name"] == "echo_tool" - assert "query" in openai_tool["function"]["parameters"]["properties"] - assert state["endpoint_tool_result"] == "echo:hi" - assert state["endpoint_user_messages"] == [{"role": "user", "content": "continue"}] - assert state["endpoint_stop"]["done"] is True - assert state["endpoint_stop"]["stop_condition"] == "state_done" - assert state["endpoint_client_class"] == "AsyncOpenAI" - assert state["endpoint_config"]["api_client_type"] == "openai_responses" - assert state["endpoint_config"]["base_url"].endswith("/v1") - assert state["endpoint_config"]["api_key_var"].startswith("VF_ENDPOINT_API_KEY_") - assert state["endpoint_config"]["api_key_var"] not in os.environ - assert "runtime_id" not in state["runtime"] - assert "endpoint_root_url" not in state - - -@pytest.mark.asyncio -async def test_vf_model_bridge_preserves_overlong_prompt_error() -> None: - harness = make_harness( - program={"fn": program_ref("endpoint_model_error_program")}, - model="test-model", - client=RaisingModelClient(vf.OverlongPromptError("too long")), - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - await harness.teardown() - - assert state["prompt_too_long"] is True - assert state["is_truncated"] is True - assert state["stop_condition"] == "prompt_too_long" - assert state.get("error") is None - - -@pytest.mark.asyncio -async def test_vf_model_bridge_preserves_model_error() -> None: - harness = make_harness( - program={"fn": program_ref("endpoint_model_error_program")}, - model="test-model", - client=RaisingModelClient(vf.ModelError("model failed")), - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - await harness.teardown() - - assert state["stop_condition"] == "has_error" - assert state["error"]["error"] == "ModelError" - assert "SandboxError" not in state["error"]["error_chain_str"] - - -@pytest.mark.asyncio -async def test_endpoint_request_can_hide_internal_model_call_from_trajectory() -> None: - client = FakeModelClient([fake_response("hidden"), fake_response("shown")]) - harness = make_harness( - program={"fn": program_ref("endpoint_trajectory_program")}, - client=client, - model="test-model", - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - await harness.teardown() - - assert len(state["trajectory"]) == 1 - assert state["trajectory"][0]["completion"][0]["content"] == "shown" - assert state["trajectory"][0]["extras"]["endpoint"] is True - assert ( - state["endpoint_hidden_response"]["choices"][0]["message"]["content"] - == "hidden" - ) - assert ( - state["endpoint_shown_response"]["choices"][0]["message"]["content"] == "shown" - ) - - -@pytest.mark.asyncio -async def test_endpoint_max_turns_counts_inflight_visible_requests() -> None: - client = BlockingModelClient( - [fake_response("allowed"), fake_response("unexpected")] - ) - harness = make_harness( - program={"fn": program_ref("concurrent_endpoint_program")}, - client=client, - model="test-model", - max_turns=1, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - run_task = asyncio.create_task(harness.run(task)) - await asyncio.wait_for(client.entered.wait(), timeout=1.0) - await asyncio.sleep(0.05) - client.release.set() - state = await run_task - await harness.teardown() - - contents = [ - response["choices"][0]["message"]["content"] - for response in state["endpoint_concurrent_responses"] - ] - assert sorted(contents) == ["", "allowed"] - assert len(client.requests) == 1 - assert len(state["trajectory"]) == 1 - assert state["stop_condition"] == "max_turns_reached" - - -@pytest.mark.asyncio -async def test_state_helpers_load_runtime_tools_while_rollout_is_active() -> None: - harness = make_harness( - program={"fn": program_ref("state_tools_program")}, - toolsets=[vf.Toolset(tools=[echo_tool])], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["tool_result"] == "echo:state" - assert state["tool_name"] == "echo_tool" - assert "runtime_id" not in state["runtime"] - - -@pytest.mark.asyncio -async def test_entrypoint_program_uses_state_tools_helper() -> None: - harness = make_harness( - program={"fn": program_ref("state_tool_program")}, - toolsets=[vf.Toolset(tools=[echo_tool])], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["tool_result"] == "echo:injected" - assert state["tool_name"] == "echo_tool" - - -@pytest.mark.asyncio -async def test_offline_replay_program_scores_without_model_client() -> None: - class ReplayTaskset(vf.Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return replay_tasks(split) - - taskset = ReplayTaskset( - config=vf.TasksetConfig(rewards=[program_ref("replay_reward")]) - ) - harness = make_harness(program={"fn": program_ref("replay_answer_program")}) - harness = vf.Env(taskset=taskset, harness=harness).harness - task = next(iter(taskset)) - - state = await harness.run(task) - await harness.teardown() - - assert state["answer"] == "solved" - assert state["reward"] == 1.0 - assert state["stop_condition"] == "program_completed" - assert state["trajectory"] == [] - assert "runtime_id" not in state["runtime"] - - -@pytest.mark.asyncio -async def test_base_program_returns_tool_errors_to_model() -> None: - client = FakeModelClient( - [ - fake_response( - tool_calls=[ - ToolCall( - id="call_1", - name="failing_tool", - arguments='{"section_id": "bad"}', - ) - ] - ), - fake_response(content="Recovered."), - ] - ) - harness = make_harness( - client=cast(Client, client), - model="fake", - toolsets=[vf.Toolset(tools=[failing_tool])], - max_turns=2, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - tool_message = state["completion"][1] - assert tool_message["role"] == "tool" - assert tool_message["content"] == "Invalid section_id format." - assert state["completion"][-1]["content"] == "Recovered." - assert state["error"] is None - - -@pytest.mark.asyncio -async def test_base_program_stops_after_tool_calls_state_stop() -> None: - client = FakeModelClient( - [ - fake_response( - tool_calls=[ - ToolCall( - id="call_1", - name="finish_tool", - arguments='{"answer": "done"}', - ) - ] - ) - ] - ) - harness = make_harness( - client=cast(Client, client), - model="fake", - toolsets=[vf.Toolset(tools=[finish_tool])], - max_turns=3, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["answer"] == "done" - assert state["is_completed"] is True - assert state["stop_condition"] == "submitted" - assert state["completion"][-1]["role"] == "tool" - assert len(client.responses) == 0 - - -@pytest.mark.asyncio -async def test_base_program_submits_system_prompt_before_prompt() -> None: - client = CapturingModelClient([fake_response(content="ok")]) - harness = make_harness(client=cast(Client, client), model="fake", max_turns=1) - task = vf.Task( - { - "system_prompt": "Use a short answer.", - "prompt": [{"role": "user", "content": "hi"}], - } - ).freeze() - - state = await harness.run(task) - - prompt = cast(list[object], client.requests[0]["prompt"]) - assert [getattr(message, "role", None) for message in prompt] == [ - "system", - "user", - ] - assert state["system_prompt"] == [ - {"role": "system", "content": "Use a short answer."} - ] - assert state["prompt"][0]["role"] == "system" - assert state["completion"][-1]["content"] == "ok" - - -@pytest.mark.asyncio -async def test_base_program_max_turns_uses_stop_condition() -> None: - client = CapturingModelClient( - [fake_response(content="done"), fake_response(content="unexpected")] - ) - harness = make_harness(client=cast(Client, client), model="fake", max_turns=1) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert len(client.requests) == 1 - assert len(state["trajectory"]) == 1 - assert state["stop_condition"] == "max_turns_reached" - - -@pytest.mark.asyncio -async def test_model_request_reservation_released_when_client_resolution_fails() -> ( - None -): - harness = make_harness(model="fake", max_turns=1) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, vf.State.for_task(task)) - - with pytest.raises(RuntimeError, match="no model client"): - await harness.runtime.submit_model_request( - cast(Messages, state["prompt"]), task, state - ) - - assert harness.runtime.visible_model_requests(state) == 0 - assert state["trajectory"] == [] - - -@pytest.mark.asyncio -async def test_taskset_setup_initializes_base_harness_prompt_and_sampling() -> None: - class SetupRuntimeTaskset(vf.Taskset): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - return setup_runtime_tasks(split) - - taskset = SetupRuntimeTaskset( - config=vf.TasksetConfig(setups=[program_ref("initialize_from_taskset")]) - ) - env = vf.Env(taskset=taskset) - client = CapturingModelClient([fake_response(content="ok")]) - - state = await env.rollout( - taskset.to_task(taskset.get_dataset()[0]), - cast(Client, client), - "fake", - {"temperature": 0.4}, - ) - - prompt = cast(list[object], client.requests[0]["prompt"]) - assert type(env.harness) is vf.Harness - assert [getattr(message, "role", None) for message in prompt] == ["user"] - assert getattr(prompt[0], "content", None) == "task ready" - assert client.requests[0]["sampling_args"] == { - "top_p": 1.0, - "temperature": 0.4, - } - assert state["runtime"]["max_turns"] == 3 - - -@pytest.mark.asyncio -async def test_callable_tool_can_accept_name_argument() -> None: - harness = make_harness(toolsets=[vf.Toolset(tools=[named_tool])]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - harness.runtime.prepare_state(task, state) - - result = await harness.runtime.call_tool("named_tool", task, state, name="Ada") - - assert result == "name:Ada" - - -@pytest.mark.asyncio -async def test_callable_tool_rejects_reserved_hidden_args() -> None: - harness = make_harness(toolsets=[vf.Toolset(tools=[echo_tool])]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - harness.runtime.prepare_state(task, state) - - with pytest.raises(ValueError, match="runtime is reserved"): - await harness.runtime.call_tool("echo_tool", task, state, runtime="bad") - - -@pytest.mark.asyncio -async def test_callable_tools_are_available_through_mcp_proxy() -> None: - harness = make_harness( - program={"fn": program_ref("mcp_proxy_program")}, - toolsets=[vf.Toolset(tools=[echo_tool])], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - await harness.teardown() - - assert state["mcp_tools"] == ["echo_tool"] - assert state["mcp_result"] == "echo:hi" - - -@pytest.mark.asyncio -async def test_command_env_exposes_model_endpoint_without_tool_payloads() -> None: - harness = make_harness(toolsets=[vf.Toolset(tools=[echo_tool])]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - state["endpoint_root_url"] = "http://127.0.0.1:1/rollout/test" - state["endpoint_base_url"] = "http://127.0.0.1:1/rollout/test/v1" - harness.runtime.prepare_state(task, state) - - env = await command_env({}, task, state, harness.runtime, include_base=False) - - assert env["OPENAI_BASE_URL"] == "http://127.0.0.1:1/rollout/test/v1" - assert env["OPENAI_API_KEY"] == harness.endpoint.secret - assert set(env) == { - "ANTHROPIC_API_KEY", - "ANTHROPIC_BASE_URL", - "OPENAI_API_KEY", - "OPENAI_BASE_URL", - } - - -@pytest.mark.asyncio -async def test_command_env_endpoint_auth_overrides_inherited_keys(monkeypatch) -> None: - monkeypatch.setenv("OPENAI_BASE_URL", "https://api.openai.invalid/v1") - monkeypatch.setenv("OPENAI_API_KEY", "host-openai-key") - monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://api.anthropic.invalid") - monkeypatch.setenv("ANTHROPIC_API_KEY", "host-anthropic-key") - - harness = make_harness(toolsets=[vf.Toolset(tools=[echo_tool])]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - state["endpoint_root_url"] = "http://127.0.0.1:1/rollout/test" - state["endpoint_base_url"] = "http://127.0.0.1:1/rollout/test/v1" - harness.runtime.prepare_state(task, state) - - env = await command_env({}, task, state, harness.runtime, include_base=True) - explicit_env = await command_env( - {"env": {"OPENAI_API_KEY": "program-key"}}, - task, - state, - harness.runtime, - include_base=True, - ) - - assert env["OPENAI_BASE_URL"] == "http://127.0.0.1:1/rollout/test/v1" - assert env["OPENAI_API_KEY"] == harness.endpoint.secret - assert env["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:1/rollout/test" - assert env["ANTHROPIC_API_KEY"] == harness.endpoint.secret - assert explicit_env["OPENAI_API_KEY"] == "program-key" - - -def test_sandbox_base_program_uses_openai_tool_payloads() -> None: - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - - program = sandbox_runner_program( - program={}, - task=task, - state=state, - mode="base", - fn_ref=None, - max_turns=3, - tool_defs=[ - Tool( - name="echo_tool", - description="", - parameters={"type": "object", "properties": {}}, - ) - ], - ) - - files = cast(dict[str, str], program["files"]) - runner_config = json.loads(files[RUNNER_CONFIG_PATH]) - assert runner_config == {"max_turns": 3} - tool_payloads = json.loads(files[TOOL_DEFS_PATH]) - assert tool_payloads == [ - { - "type": "function", - "function": { - "name": "echo_tool", - "description": "", - "parameters": {"type": "object", "properties": {}}, - }, - } - ] - tool_payloads_by_protocol = json.loads(files[TOOL_DEFS_BY_PROTOCOL_PATH]) - assert tool_payloads_by_protocol["openai_responses"][0]["name"] == "echo_tool" - assert tool_payloads_by_protocol["anthropic_messages"][0]["name"] == "echo_tool" - - -def test_sandbox_fn_program_installs_local_package( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - module_path = tmp_path / "local_program.py" - module_path.write_text("async def run(task, state): return state\n") - (tmp_path / "pyproject.toml").write_text( - """ -[project] -name = "local-program" -version = "0.1.0" - -[build-system] -requires = ["hatchling"] -build-backend = "hatchling.build" -""".strip() - ) - monkeypatch.syspath_prepend(str(tmp_path)) - - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - program = sandbox_runner_program( - program={"env": {"PYTHONPATH": "/custom"}}, - task=task, - state=state, - mode="fn", - fn_ref="local_program:run", - max_turns=1, - tool_defs=[], - ) - - dirs = cast(dict[str, object], program["dirs"]) - env = cast(dict[str, object], program["env"]) - command = cast(list[str], program["command"]) - setup = cast(list[str], program["setup"]) - assert dirs[PACKAGE_ROOT] == str(tmp_path.resolve()) - assert env["PYTHONPATH"] == "/custom" - assert "pip install" in setup[1] - assert shlex.quote(PACKAGE_ROOT) in setup[1] - assert command == [ - SANDBOX_PYTHON, - "/tmp/vf_program_runner.py", - "fn", - "local_program:run", - ] - - -def test_sandbox_fn_program_resolves_local_module_package( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - module_path = tmp_path / "standalone_program.py" - module_path.write_text("async def run(task, state): return state\n") - (tmp_path / "pyproject.toml").write_text( - """ -[project] -name = "standalone-program" -version = "0.1.0" - -[build-system] -requires = ["hatchling"] -build-backend = "hatchling.build" -""".strip() - ) - monkeypatch.syspath_prepend(str(tmp_path)) - - package = sandbox_program_package(mode="fn", fn_ref="standalone_program:run") - - assert package == SandboxPackage(local_root=tmp_path.resolve()) - - -def test_sandbox_fn_program_resolves_package_module_root( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - package = tmp_path / "package_program" - package.mkdir() - (package / "__init__.py").write_text("") - (package / "worker.py").write_text("async def run(task, state): return state\n") - (package / "pyproject.toml").write_text( - """ -[project] -name = "package-program" -version = "0.1.0" - -[build-system] -requires = ["hatchling"] -build-backend = "hatchling.build" -""".strip() - ) - monkeypatch.syspath_prepend(str(tmp_path)) - - package_root = sandbox_program_package( - mode="fn", fn_ref="package_program.worker:run" - ) - - assert package_root == SandboxPackage(local_root=package.resolve()) - - -def test_sandbox_fn_program_does_not_walk_to_parent_pyproject( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - src = tmp_path / "src" - src.mkdir() - (src / "nested_program.py").write_text("async def run(task, state): return state\n") - (tmp_path / "pyproject.toml").write_text( - """ -[project] -name = "parent-program" -version = "0.1.0" - -[build-system] -requires = ["hatchling"] -build-backend = "hatchling.build" -""".strip() - ) - monkeypatch.syspath_prepend(str(src)) - - with pytest.raises(ValueError, match="beside the resolved environment module"): - sandbox_program_package(mode="fn", fn_ref="nested_program:run") - - -def test_sandbox_fn_program_does_not_install_stdlib_packages() -> None: - assert sandbox_program_package(mode="fn", fn_ref="json:dumps") is None - - -def test_sandbox_fn_program_requires_local_pyproject( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - module_path = tmp_path / "unpackaged_program.py" - module_path.write_text("async def run(task, state): return state\n") - monkeypatch.syspath_prepend(str(tmp_path)) - - with pytest.raises(ValueError, match="no pyproject.toml"): - sandbox_program_package(mode="fn", fn_ref="unpackaged_program:run") - - -def test_sandbox_python_program_installs_runtime_client_deps() -> None: - harness = make_harness(program={"sandbox": True}, sandbox={"packages": ["numpy"]}) - - sandbox = harness.prepare_sandbox_config( - vf.SandboxConfig(packages=["numpy"]), - {"sandbox": True}, - ) - - assert sandbox.packages == ["numpy", "openai", "anthropic", "requests"] - - -def test_sandbox_package_install_bootstraps_managed_python() -> None: - command = python_package_install_command("mcp>=1.14.1 requests") - - assert "UV_NO_CONFIG=1" not in command - assert "UV_INDEX_URL" not in command - assert "PIP_INDEX_URL" not in command - assert "https://astral.sh/uv/install.sh" in command - assert '"$VF_UV" venv --seed --python "$VF_PYTHON_VERSION"' in command - assert '"$VF_UV" pip install --python "$VF_PYTHON"' in command - assert "--index-url" not in command - assert SANDBOX_PYTHON in command - assert SANDBOX_UV in command - assert "mcp>=1.14.1 requests" in command - - -@pytest.mark.asyncio -async def test_create_sandbox_retries_create_and_bounds_wait( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - disable_sandbox_retry_sleep(monkeypatch) - - class FlakyCreateClient: - def __init__(self) -> None: - self.create_calls = 0 - self.wait_calls: list[tuple[str, int]] = [] - self.deleted: list[str] = [] - - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - _ = request - self.create_calls += 1 - if self.create_calls == 1: - raise RuntimeError("transient create") - return FakeSandboxResult("sbx-retry") - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> None: - self.wait_calls.append((sandbox_id, max_attempts)) - - async def delete(self, sandbox_id: str) -> None: - self.deleted.append(sandbox_id) - - client = FlakyCreateClient() - - sandbox_id = await sandbox_utils.create_sandbox( - cast(sandbox_utils.SandboxClient, client), - {"image": "python:3.11-slim"}, - ) - - assert sandbox_id == "sbx-retry" - assert client.create_calls == 2 - assert client.wait_calls == [ - ("sbx-retry", sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS) - ] - assert client.deleted == [] - - -@pytest.mark.asyncio -async def test_create_sandbox_cleans_up_wait_failure_with_retry( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - disable_sandbox_retry_sleep(monkeypatch) - - class WaitFailingClient: - def __init__(self) -> None: - self.delete_calls = 0 - - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - _ = request - return FakeSandboxResult("sbx-wait") - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> None: - assert sandbox_id == "sbx-wait" - assert max_attempts == sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS - raise RuntimeError("wait failed") - - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == "sbx-wait" - self.delete_calls += 1 - if self.delete_calls == 1: - raise RuntimeError("transient delete") - - client = WaitFailingClient() - - with pytest.raises(RuntimeError, match="wait failed"): - await sandbox_utils.create_sandbox( - cast(sandbox_utils.SandboxClient, client), - {"image": "python:3.11-slim"}, - ) - - assert client.delete_calls == 2 - - -@pytest.mark.asyncio -async def test_upload_program_files_retries_transient_transfer_error( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - disable_sandbox_retry_sleep(monkeypatch) - - class FlakyUploadClient: - calls = 0 - - async def upload_bytes(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - self.calls += 1 - if self.calls == 1: - raise FakeAPIError("Upload failed: ") - - client = FlakyUploadClient() - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - - await sandbox_utils.upload_program_files( - cast(sandbox_utils.SandboxClient, client), - "sbx-upload", - {"files": {"/tmp/file.txt": "content"}}, - task, - state, - Runtime(), - ) - - assert client.calls == 2 - - -@pytest.mark.asyncio -async def test_upload_program_files_does_not_retry_non_transient_api_error( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - disable_sandbox_retry_sleep(monkeypatch) - - class FailingUploadClient: - calls = 0 - - async def upload_bytes(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - self.calls += 1 - raise FakeAPIError("Upload failed: HTTP 400: bad request") - - client = FailingUploadClient() - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - - with pytest.raises(vf.SandboxError, match="HTTP 400"): - await sandbox_utils.upload_program_files( - cast(sandbox_utils.SandboxClient, client), - "sbx-upload", - {"files": {"/tmp/file.txt": "content"}}, - task, - state, - Runtime(), - ) - - assert client.calls == 1 - - -@pytest.mark.asyncio -async def test_create_sandbox_cancellation_deletes_late_provider_result( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - disable_sandbox_retry_sleep(monkeypatch) - started = asyncio.Event() - finish = asyncio.Event() - deleted: list[str] = [] - - class SlowCreateClient: - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - _ = request - started.set() - await finish.wait() - return FakeSandboxResult("sbx-created-after-cancel") - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> None: - _ = sandbox_id, max_attempts - - async def delete(self, sandbox_id: str) -> None: - deleted.append(sandbox_id) - - task = asyncio.create_task( - sandbox_utils.create_sandbox( - cast(sandbox_utils.SandboxClient, SlowCreateClient()), - {"image": "python:3.11-slim"}, - ) - ) - await started.wait() - - task.cancel() - finish.set() - with pytest.raises(asyncio.CancelledError): - await task - - assert deleted == ["sbx-created-after-cancel"] - - -@pytest.mark.asyncio -async def test_create_sandbox_wait_cancellation_deletes_known_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - disable_sandbox_retry_sleep(monkeypatch) - waiting = asyncio.Event() - deleted: list[str] = [] - - class WaitingClient: - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - _ = request - return FakeSandboxResult("sbx-wait-cancel") - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> None: - _ = sandbox_id, max_attempts - waiting.set() - await asyncio.Event().wait() - - async def delete(self, sandbox_id: str) -> None: - deleted.append(sandbox_id) - - task = asyncio.create_task( - sandbox_utils.create_sandbox( - cast(sandbox_utils.SandboxClient, WaitingClient()), - {"image": "python:3.11-slim"}, - ) - ) - await waiting.wait() - - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - - assert deleted == ["sbx-wait-cancel"] - - -@pytest.mark.asyncio -async def test_create_sandbox_threads_v1_request_fields( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - - await sandbox_utils.create_sandbox( - cast(sandbox_utils.SandboxClient, FakeSandboxClient()), - { - "image": "custom:latest", - "start_command": "sleep infinity", - "cpu_cores": 2, - "memory_gb": 8, - "disk_size_gb": 20, - "gpu_count": 1, - "gpu_type": "a10", - "vm": True, - "network_access": False, - "timeout_minutes": 30, - "environment_vars": {"NUMBER": 7}, - "secrets": {"TOKEN": "secret"}, - "team_id": "team", - "region": "us", - "registry_credentials_id": "registry", - "guaranteed": True, - "labels": ["v1"], - }, - ) - - request = FakeSandboxClient.requests[0] - assert request["docker_image"] == "custom:latest" - assert request["memory_gb"] == 8.0 - assert request["disk_size_gb"] == 20.0 - assert request["gpu_type"] == "a10" - assert request["vm"] is True - assert request["network_access"] is False - assert request["environment_vars"] == {"NUMBER": "7"} - assert request["secrets"] == {"TOKEN": "secret"} - assert request["team_id"] == "team" - assert request["region"] == "us" - assert request["registry_credentials_id"] == "registry" - assert request["guaranteed"] is True - - -@pytest.mark.asyncio -async def test_sandbox_base_program_max_turns_zero_is_unbounded( - tmp_path: Path, -) -> None: - namespace: dict[str, object] = {} - source = runner_source().rsplit("asyncio.run(main())", 1)[0] - exec(source, namespace) - config_path = tmp_path / "runner_config.json" - config_path.write_text(json.dumps({"max_turns": 0})) - namespace["RUNNER_CONFIG_PATH"] = str(config_path) - - async def create_model_message(state, messages): - _ = state, messages - return {"role": "assistant", "content": "done"} - - async def call_user(state, messages): - _ = state, messages - return [] - - async def check_stop(state): - _ = state - return False - - namespace["create_model_message"] = create_model_message - namespace["call_user"] = call_user - namespace["check_stop"] = check_stop - - state = {"prompt": [{"role": "user", "content": "hi"}], "runtime": {}} - run_base = cast(Any, namespace["run_base"]) - result = await run_base({}, state) - - assert result["completion"] == [{"role": "assistant", "content": "done"}] - assert result["stop_condition"] == "no_tools" - - -@pytest.mark.asyncio -async def test_sandbox_base_program_model_call_uses_vf_model_bridge() -> None: - namespace: dict[str, object] = {} - source = runner_source().rsplit("asyncio.run(main())", 1)[0] - exec(source, namespace) - - posted: list[tuple[str, Any, object]] = [] - - async def vf_post(state, path, payload, timeout=None): - _ = state - posted.append((path, payload, timeout)) - return {"message": {"role": "assistant", "content": "ok"}} - - namespace["vf_post"] = vf_post - create_model_message = cast(Any, namespace["create_model_message"]) - - # Canonical Messages (incl. an image content part) are sent unchanged over the - # /vf/model bridge; the host owns client resolution + tokenization and returns - # the assistant message. - messages = [ - {"role": "user", "content": "hi"}, - { - "role": "tool", - "tool_call_id": "call_1", - "content": [ - {"type": "text", "text": "shot"}, - { - "type": "image_url", - "image_url": {"url": "data:image/png;base64,AAA"}, - }, - ], - }, - ] - message = await create_model_message({"runtime": {}}, messages) - - assert message == {"role": "assistant", "content": "ok"} - assert len(posted) == 1 - path, payload, timeout = posted[0] - assert path == "model" - assert payload["messages"] == messages # image part preserved verbatim - assert timeout is None - - -def test_sandbox_program_patch_cannot_set_lifecycle_fields() -> None: - state = vf.State.for_task(vf.Task({"prompt": []}).freeze()) - - with pytest.raises(RuntimeError, match="framework-managed"): - apply_internal_state_patch( - state, - {"stop_condition": "user_stop"}, - mode="fn", - ) - with pytest.raises(RuntimeError, match="framework-managed"): - apply_internal_state_patch( - state, - {"is_completed": True}, - mode="base", - ) - - patch = { - "stop_condition": "no_tools", - "is_truncated": True, - "error": vf.ErrorData( - error="SandboxError", - message="handled", - error_chain_repr="SandboxError('handled')", - error_chain_str="SandboxError", - ), - } - apply_internal_state_patch(state, patch, mode="base") - - assert patch == {} - assert state["stop_condition"] == "no_tools" - assert state["is_truncated"] is True - assert isinstance(state["error"], vf.SandboxError) - - -def test_program_channels_mcp_injects_proxy_into_sandbox_program() -> None: - harness = make_harness( - program={"sandbox": True, "command": ["true"], "channels": "mcp"}, - sandbox={"image": "python:3.11-slim"}, - ) - state = vf.State.for_task(vf.Task({}).freeze()) - state["endpoint_root_url"] = "http://127.0.0.1:1/rollout/test" - - program = harness.prepare_sandbox_program( - {"sandbox": True, "command": ["true"], "channels": "mcp"}, state - ) - sandbox = harness.prepare_sandbox_config( - vf.SandboxConfig(image="python:3.11-slim"), - {"sandbox": True, "command": ["true"], "channels": "mcp"}, - ) - - files = cast(dict[str, str], program["files"]) - assert MCP_PROXY_PATH in files - assert MCP_PROXY_CONFIG_PATH in files - config = json.loads(files[MCP_PROXY_CONFIG_PATH]) - assert config == { - "tool_base_url": "http://127.0.0.1:1/rollout/test/vf/tools", - "tool_auth_var": "OPENAI_API_KEY", - } - assert proxy_command() == [SANDBOX_PYTHON, MCP_PROXY_PATH, MCP_PROXY_CONFIG_PATH] - assert "mcp>=1.14.1" in sandbox.packages - assert "requests" in sandbox.packages - - -def test_program_channels_mcp_requires_sandbox_command() -> None: - with pytest.raises(ValueError, match="requires program.sandbox"): - make_harness(program={"command": ["true"], "channels": "mcp"}) - - -def test_program_channels_callable_rejects_command_programs() -> None: - with pytest.raises(ValueError, match="program.channels='callable'"): - make_harness(program={"command": ["true"], "channels": "callable"}) - - -@pytest.mark.asyncio -async def test_program_channels_mcp_setup_uses_bindings_after_setup_before_command( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={ - "command": ["python", "-c", "print('ok')"], - "sandbox": True, - "setup": "echo setup", - "channels": {"mcp": {"fn": program_ref("configure_cli_endpoint")}}, - "bindings": { - "configure_cli_endpoint.endpoint_config": { - "fn": program_ref("endpoint_config_binding") - } - }, - }, - sandbox={"image": "python:3.11-slim"}, - model="bound-model", - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - await harness.run(task) - - commands = [command for _, command in FakeSandboxClient.commands] - setup_index = next( - i for i, command in enumerate(commands) if command.endswith("echo setup") - ) - mcp_setup_index = next( - i - for i, command in enumerate(commands) - if command.endswith("echo model=bound-model > /tmp/endpoint.txt") - ) - command_index = next( - i - for i, command in enumerate(commands) - if command.endswith("python -c 'print('\"'\"'ok'\"'\"')'") - ) - assert setup_index < mcp_setup_index < command_index - - -@pytest.mark.asyncio -async def test_rollout_setup_receives_program_sandbox_before_program_setup( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={ - "command": ["true"], - "sandbox": True, - "setup": "echo program-setup", - }, - sandbox={"image": "python:3.11-slim"}, - setups=[ - program_ref("early_sandbox_lifecycle_setup"), - program_ref("sandbox_lifecycle_setup"), - ], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - commands = [command for _, command in FakeSandboxClient.commands] - early_setup_index = commands.index("echo early-lifecycle-setup") - lifecycle_setup_index = commands.index("echo lifecycle-setup") - program_setup_index = commands.index("echo program-setup") - command_index = commands.index("true") - assert state["setup_sandbox_id"] == "sbx-1" - assert state["early_setup_sandbox_id"] == "sbx-1" - assert early_setup_index < program_setup_index - assert program_setup_index < lifecycle_setup_index < command_index - - -@pytest.mark.asyncio -async def test_program_setup_uses_program_setup_timeout( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={ - "command": ["true"], - "sandbox": True, - "setup": "echo program-setup", - "setup_timeout": 777, - }, - sandbox={"image": "python:3.11-slim"}, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - await harness.run(task) - - setup_commands = FakeSandboxClient.commands[ - : len(FakeSandboxClient.command_timeouts) - ] - command_timeouts = dict( - zip( - [command for _, command in setup_commands], - FakeSandboxClient.command_timeouts, - strict=True, - ) - ) - assert command_timeouts["echo program-setup"] == 777 - - -@pytest.mark.asyncio -async def test_sandbox_command_uses_configured_poll_interval( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={"command": ["true"], "sandbox": True}, - sandbox={"image": "python:3.11-slim", "poll_interval": 11}, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - await harness.run(task) - - assert FakeSandboxClient.background_jobs == [("sbx-1", "true", 900, None, 11)] - - -@pytest.mark.asyncio -async def test_sandbox_handle_forwards_background_job_poll_interval() -> None: - class BackgroundJobClient: - poll_intervals: list[int] = [] - - async def run_background_job( - self, - sandbox_id: str, - command: str, - *, - poll_interval: int = 3, - **kwargs: object, - ) -> FakeCommandResult: - _ = sandbox_id, command, kwargs - self.poll_intervals.append(poll_interval) - return FakeCommandResult() - - client = BackgroundJobClient() - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, client), - "sbx-1", - "rollout", - "program", - owns_client=False, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - handle = sandbox_utils.SandboxHandle(lease, state) - - await handle.run_background_job("true", poll_interval=11) - - assert client.poll_intervals == [11] - - -@pytest.mark.asyncio -async def test_sandbox_command_marks_oom_failures( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - class SandboxOOMError(Exception): - pass - - class OOMSandboxClient(FakeSandboxClient): - def __init__(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - - async def run_background_job( - self, *args: object, **kwargs: object - ) -> FakeCommandResult: - _ = args, kwargs - raise SandboxOOMError("Container exceeded memory limit") - - monkeypatch.setattr( - "verifiers.utils.threaded_sandbox_client.ThreadedAsyncSandboxClient", - OOMSandboxClient, - ) - - harness = make_harness( - program={"command": ["true"], "sandbox": True}, - sandbox={"image": "python:3.11-slim"}, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["sandbox_oom"] is True - assert state["sandbox_failures"][0]["kind"] == "oom" - assert state["sandbox_failures"][0]["phase"] == "command" - assert state["error"]["error"] == "SandboxError" - - -@pytest.mark.asyncio -async def test_sandbox_state_input_upload_runs_after_rollout_setup( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - - harness = make_harness(setups=[program_ref("state_input_setup")]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - program = { - "command": ["true"], - VF_STATE_INPUT_PATH_KEY: "/tmp/vf_state_in.json", - } - - await run_sandbox_command( - program, - vf.SandboxConfig(image="python:3.11-slim"), - task, - state, - harness.runtime, - ) - - uploads = { - path: json.loads(data.decode()) - for _, path, data in FakeSandboxClient.uploads - if path == "/tmp/vf_state_in.json" - } - assert uploads["/tmp/vf_state_in.json"]["state_input_setup"] is True - - -@pytest.mark.asyncio -async def test_task_command_uses_background_job( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={"command": ["sleep", "120"], "sandbox": True}, - sandbox={"image": "python:3.11-slim", "workdir": "/app"}, - ) - task = vf.Task( - { - "prompt": [{"role": "user", "content": "hi"}], - "sandbox": {"command_timeout": 120}, - } - ).freeze() - - await harness.run(task) - - assert ("sbx-1", "sleep 120", 120, "/app", 3) in FakeSandboxClient.background_jobs - - -@pytest.mark.asyncio -async def test_program_channels_mcp_setup_accepts_config_ref_mappings( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={ - "command": ["true"], - "sandbox": True, - "channels": {"mcp": [{"fn": program_ref("configure_cli_endpoint_ref")}]}, - "bindings": { - "configure_cli_endpoint_ref.endpoint_config": { - "fn": program_ref("endpoint_config_binding_ref") - } - }, - }, - sandbox={"image": "python:3.11-slim"}, - model="toml-model", - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - await harness.run(task) - - commands = [command for _, command in FakeSandboxClient.commands] - assert any( - command.endswith("echo ref-model=toml-model > /tmp/ref_endpoint.txt") - for command in commands - ) - - -def test_program_bindings_must_match_owned_callables() -> None: - with pytest.raises(ValueError, match="does not match a callable"): - make_harness( - program={ - "command": ["true"], - "sandbox": True, - "bindings": {"missing.value": "task.value"}, - }, - sandbox={"image": "python:3.11-slim"}, - ) - - -def test_program_setup_is_not_a_binding_target() -> None: - with pytest.raises(ValueError, match="setup callables cannot use"): - make_harness( - program={ - "command": ["true"], - "sandbox": True, - "setup": {"fn": program_ref("configure_cli_endpoint")}, - "bindings": { - "configure_cli_endpoint.endpoint_config": { - "fn": program_ref("endpoint_config_binding") - } - }, - }, - sandbox={"image": "python:3.11-slim"}, - ) - - -async def write_real_sandbox_file(text: str, sandbox, state) -> str: - command = ( - "python - <<'PY'\n" - "from pathlib import Path\n" - f"Path('/tmp/host_tool.txt').write_text({text!r})\n" - "print(Path('/tmp/host_tool.txt').read_text())\n" - "PY\n" - ) - result = await sandbox.execute(command, timeout=120, working_dir="/tmp") - output = str(getattr(result, "stdout", "") or "").strip() - state["real_sandbox_tool_output"] = output - return output - - -REAL_MCP_PROXY_SCRIPT = r""" -import asyncio -import json -import os - -from mcp import ClientSession, StdioServerParameters -from mcp.client.stdio import stdio_client - - -async def main(): - with open("/tmp/vf_mcp_tools.json") as f: - config = json.load(f) - tool_auth_var = str(config["tool_auth_var"]) - server = StdioServerParameters( - command="python3", - args=["/tmp/vf_mcp_tools.py", "/tmp/vf_mcp_tools.json"], - env={tool_auth_var: os.environ[tool_auth_var]}, - ) - async with stdio_client(server) as (read_stream, write_stream): - async with ClientSession(read_stream, write_stream) as session: - await session.initialize() - listed = await session.list_tools() - result = await session.call_tool("echo_tool", {"query": "real-mcp"}) - payload = { - "tools": [tool.name for tool in listed.tools], - "result": result.content[0].text, - } - print(json.dumps(payload)) - - -asyncio.run(main()) -""" - - -@pytest.mark.asyncio -@pytest.mark.prime_sandbox -async def test_real_sandbox_base_program_calls_host_callable_tool() -> None: - client = FakeModelClient( - [ - fake_response( - tool_calls=[ - ToolCall( - id="call_1", - name="write_real_sandbox_file", - arguments='{"text": "from-real-sandbox"}', - ) - ] - ), - fake_response(content="done"), - ] - ) - harness = make_harness( - client=cast(Client, client), - model="fake", - program={"sandbox": True, "channels": "callable"}, - sandbox={ - "image": "python:3.11-slim", - "scope": "group", - "network_access": True, - "timeout_minutes": 20, - "command_timeout": 120, - }, - toolsets=[ - vf.Toolset( - tools=[write_real_sandbox_file], - write=True, - sandbox="program", - ) - ], - max_turns=2, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "write file"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["group_key"] = "real-sandbox-callable-tools" - - state = await harness.run(task, state) - - try: - assert state["sandbox_id"] - assert state.get("real_sandbox_tool_output") == "from-real-sandbox", { - "endpoint_root_url": state.get("endpoint_root_url"), - "endpoint_base_url": state.get("endpoint_base_url"), - "error": state.get("error"), - } - assert state["completion"][-1]["content"] == "done" - finally: - await harness.cleanup_group([task], [state]) - await harness.teardown() - - -@pytest.mark.asyncio -@pytest.mark.prime_sandbox -async def test_real_sandbox_command_program_uses_mcp_tool_proxy() -> None: - harness = make_harness( - program={ - "sandbox": True, - "command": ["python", "/tmp/call_mcp.py"], - "channels": "mcp", - "files": {"/tmp/call_mcp.py": REAL_MCP_PROXY_SCRIPT}, - }, - sandbox={ - "image": "python:3.9-slim", - "network_access": True, - "timeout_minutes": 20, - "command_timeout": 120, - }, - toolsets=[vf.Toolset(tools=[echo_tool])], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "call mcp"}]}).freeze() - state = vf.State.for_task(task) - - state = await harness.run(task, state) - - try: - stdout = state["command"]["stdout"].strip() - payload = json.loads(stdout) - assert payload == {"tools": ["echo_tool"], "result": "echo:real-mcp"} - finally: - await harness.cleanup_group([task], [state]) - await harness.teardown() - - -@pytest.mark.asyncio -async def test_nested_harness_uses_explicit_child_model_controls() -> None: - harness = make_harness(program={"fn": program_ref("parent_program")}) - task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["model"] = "model-a" - state["runtime"]["sampling_args"] = {"temperature": 0.2} - harness.runtime.bind_model_client(state, cast(Client, FakeClient())) - - state = await harness.run(task, state) - - child_state = state["child_state"] - assert child_state["trajectory_id"] != state["trajectory_id"] - assert child_state["child_runtime"]["model"] == "model-a" - assert child_state["child_runtime"]["sampling_args"] == {"temperature": 0.2} - assert "child_rollouts" not in state - assert "client_key" not in child_state["runtime"] - assert "client_key" not in child_state["child_runtime"] - - -@pytest.mark.asyncio -async def test_state_finalize_strips_nested_runtime_handles() -> None: - harness = make_harness(program={"fn": program_ref("parent_program")}) - task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["model"] = "model-a" - state["runtime"]["group_key"] = "group-a" - harness.runtime.bind_model_client(state, cast(Client, FakeClient())) - state["runtime"]["resolved"] = { - "model": { - "runtime_id": harness.runtime.runtime_id, - "client_key": state["runtime"]["client_key"], - } - } - - state = await harness.run(task, state) - state.finalize() - - assert "runtime_id" not in state["runtime"] - assert "client_key" not in state["runtime"] - assert "resolved" not in state["runtime"] - assert "runtime_id" not in state["child_state"]["runtime"] - assert "client_key" not in state["child_state"]["runtime"] - assert "client_key" not in state["child_state"]["child_runtime"] - - -@pytest.mark.asyncio -async def test_nested_harness_can_use_own_model_controls() -> None: - harness = make_harness( - program={"fn": program_ref("parent_calls_owned_child_program")} - ) - task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["model"] = "parent-model" - state["runtime"]["sampling_args"] = {"temperature": 0.2} - harness.runtime.bind_model_client(state, cast(Client, FakeClient())) - - state = await harness.run(task, state) - - child_state = state["child_state"] - assert child_state["child_runtime"]["model"] == "child-model" - assert "client_key" not in child_state["child_runtime"] - assert "client_key" not in state["runtime"] - - -@pytest.mark.asyncio -async def test_task_model_controls_override_harness_model_controls() -> None: - client = CapturingModelClient([fake_response("done")]) - harness = make_harness( - client=cast(Client, client), - model="harness-model", - sampling_args={"temperature": 0.1, "top_p": 1.0}, - ) - task = vf.Task( - { - "prompt": [{"role": "user", "content": "Use task model."}], - "model": { - "name": "task-model", - "sampling_args": {"temperature": 0.4}, - }, - } - ).freeze() - - await harness.run(task) - - assert task["model"] == { - "name": "task-model", - "sampling_args": {"temperature": 0.4}, - } - assert client.requests[0]["model"] == "task-model" - assert client.requests[0]["sampling_args"] == { - "temperature": 0.4, - "top_p": 1.0, - } - - -@pytest.mark.asyncio -async def test_update_child_harness_run_uses_resolved_runtime_handles() -> None: - client = CapturingModelClient( - [fake_response("parent answer"), fake_response("summary")] - ) - harness = make_harness( - updates=[program_ref("update_summary_with_resolved_handles")] - ) - task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["model"] = "model-a" - harness.runtime.bind_model_client(state, cast(Client, client)) - - state = await harness.run(task, state) - - assert state["summary"] == "summary" - assert len(client.requests) == 2 - assert len(state["trajectory"]) == 2 - assert state["trajectory"][0]["trajectory_id"] == state["trajectory_id"] - assert state["trajectory"][1]["trajectory_id"] == state["child_trajectory_id"] - assert state["num_model_requests"] == 2 - assert state["completion"][-1]["content"] == "summary" - - -@pytest.mark.asyncio -async def test_update_child_harness_can_borrow_live_tools() -> None: - client = CapturingModelClient( - [ - fake_response("parent answer"), - fake_response( - tool_calls=[ - ToolCall( - id="call_1", - name="borrowed_record_tool", - arguments='{"value": "from-child"}', - ) - ] - ), - fake_response("child judged"), - ] - ) - harness = make_harness( - updates=[program_ref("update_child_uses_borrowed_tool")], - toolsets=[vf.Toolset(tools=[borrowed_record_tool], write=True)], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["model"] = "model-a" - harness.runtime.bind_model_client(state, cast(Client, client)) - - state = await harness.run(task, state) - - assert state["borrowed_tool_values"] == ["from-child"] - assert state["child_completion"] == "child judged" - assert len(client.requests) == 3 - assert len(state["trajectory"]) == 3 - assert state["trajectory"][0]["trajectory_id"] == state["trajectory_id"] - assert state["trajectory"][1]["trajectory_id"] == state["child_trajectory_id"] - assert state["trajectory"][2]["trajectory_id"] == state["child_trajectory_id"] - assert state["num_model_requests"] == 3 - assert state["completion"][-1]["content"] == "child judged" - - -async def update_parallel_children_use_borrowed_tool(task, state): - _ = task - - async def run_child(label: str) -> vf.State: - child_task = vf.Task( - {"prompt": [{"role": "user", "content": f"inspect {label}"}]} - ).freeze() - child_state = state.for_task( - child_task, - borrow="model", - tools="borrowed_stage_tool", - transcript="append", - ) - return await make_harness(max_turns=2).run(child_task, child_state) - - children = await asyncio.gather(run_child("a"), run_child("b")) - state["update_child_trajectory_ids"] = [ - child["trajectory_id"] for child in children - ] - - -async def reward_child_uses_borrowed_tool(task, state) -> float: - _ = task - child_task = vf.Task( - {"prompt": [{"role": "user", "content": "score sandbox state"}]} - ).freeze() - child_state = state.for_task( - child_task, - borrow="model", - tools="borrowed_stage_tool", - ) - child_state = await make_harness(max_turns=2).run(child_task, child_state) - state["reward_child_completion"] = child_state["completion"][-1]["content"] - state["reward_child_requests"] = child_state["num_model_requests"] - return float("reward" in state.get("borrowed_stage_values", [])) - - -setattr( - ref_module, - "update_parallel_children_use_borrowed_tool", - update_parallel_children_use_borrowed_tool, -) -setattr(ref_module, "reward_child_uses_borrowed_tool", reward_child_uses_borrowed_tool) - - -class RoutedModelClient: - """Routes responses by inspecting the request's conversation, not call order. - - Robust to event-loop interleaving (e.g. asyncio.gather): each rollout's - request carries its own conversation context, so we never depend on which - coroutine happens to wake first. - """ - - def __init__(self) -> None: - self.requests: list[dict[str, object]] = [] - - async def get_response(self, **kwargs: object) -> Response: - prompt = kwargs.get("prompt") or [] - self.requests.append(dict(kwargs)) - - # First user message identifies the rollout (parent / update-a / update-b / reward). - first_user = next( - (m.get("content") for m in prompt if m.get("role") == "user"), "" - ) - # Presence of a tool-role message tells us we're past turn 1 (tool already executed). - has_tool_msg = any(m.get("role") == "tool" for m in prompt) - - if first_user == "parent": - return fake_response("parent answer") - if first_user.startswith("inspect "): - label = first_user.split(" ", 1)[1] - if not has_tool_msg: - return fake_response( - tool_calls=[ - ToolCall( - id=f"call_update_{label}", - name="borrowed_stage_tool", - arguments=f'{{"value": "update-{label}"}}', - ) - ] - ) - return fake_response(f"update {label} done") - if first_user == "score sandbox state": - if not has_tool_msg: - return fake_response( - tool_calls=[ - ToolCall( - id="call_reward", - name="borrowed_stage_tool", - arguments='{"value": "reward"}', - ) - ] - ) - return fake_response('{"score": 1.0}') - raise AssertionError(f"Unexpected first_user: {first_user!r}") - - -@pytest.mark.asyncio -async def test_update_and_reward_children_can_share_borrowed_live_tools() -> None: - client = RoutedModelClient() - harness = make_harness( - updates=[program_ref("update_parallel_children_use_borrowed_tool")], - rewards=[program_ref("reward_child_uses_borrowed_tool")], - toolsets=[vf.Toolset(tools=[borrowed_stage_tool], write=True)], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["model"] = "model-a" - harness.runtime.bind_model_client(state, cast(Client, client)) - - state = await harness.run(task, state) - - # All three borrowed-tool invocations should have landed; order is not - # asserted because parallel `asyncio.gather` may interleave them. - assert sorted(state["borrowed_stage_values"]) == ["reward", "update-a", "update-b"] - assert state["reward"] == 1.0 - assert state["reward_child_completion"] == '{"score": 1.0}' - assert len(client.requests) == 7 - assert len(state["trajectory"]) == 5 - assert state["trajectory"][0]["trajectory_id"] == state["trajectory_id"] - update_trajectory_ids = set(state["update_child_trajectory_ids"]) - assert len(update_trajectory_ids) == 2 - assert { - str(record["trajectory_id"]) for record in state["trajectory"][1:] - } == update_trajectory_ids - assert all( - record["completion"] != [{"role": "assistant", "content": '{"score": 1.0}'}] - for record in state["trajectory"] - ) - assert "runtime" not in state or "resolved" not in state.get("runtime", {}) - assert state["num_model_requests"] == 5 - assert state["reward_child_requests"] == 2 - - -@pytest.mark.asyncio -async def test_toolset_can_contribute_stop_condition() -> None: - harness = make_harness( - program={"fn": program_ref("mark_submitted")}, - toolsets=[vf.Toolset(stops=[submitted])], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["submitted"] is True - assert state["is_completed"] is True - assert state["stop_condition"] == "submitted" - - -@pytest.mark.asyncio -async def test_runtime_owned_model_clients_close_after_rollout( - monkeypatch: pytest.MonkeyPatch, -) -> None: - runtime = Runtime() - client = FakeClient() - state = vf.State.for_task(vf.Task({}).freeze()) - - monkeypatch.setattr("verifiers.v1.runtime.resolve_client", lambda config: client) - - runtime.bind_model_client( - state, - ClientConfig( - client_type="openai_chat_completions", - api_base_url="https://example.com/v1", - api_key_var="KEY", - ), - ) - await runtime.release_model_client(state) - - assert client.closed is True - assert runtime.model_clients == {} - - -@pytest.mark.asyncio -async def test_runtime_owned_model_clients_live_until_group_cleanup( - monkeypatch: pytest.MonkeyPatch, -) -> None: - runtime = Runtime() - client = FakeClient() - state = vf.State.for_task(vf.Task({}).freeze()) - state["runtime"]["group_key"] = "group" - - monkeypatch.setattr("verifiers.v1.runtime.resolve_client", lambda config: client) - - runtime.bind_model_client( - state, - ClientConfig( - client_type="openai_chat_completions", - api_base_url="https://example.com/v1", - api_key_var="KEY", - ), - ) - await runtime.release_model_client(state) - - assert client.closed is False - assert len(runtime.model_clients) == 1 - - await runtime.release_model_client(state, group=True) - - assert client.closed is True - assert runtime.model_clients == {} - - -@pytest.mark.asyncio -async def test_mcp_lifetime_follows_toolset_scope( - monkeypatch: pytest.MonkeyPatch, -) -> None: - async def connect_mcp_tool( - spec: vf.MCPTool, exit_stack: AsyncExitStack - ) -> list[FakeMCPHandle]: - _ = exit_stack - return [FakeMCPHandle(spec.command)] - - monkeypatch.setattr(mcp_utils, "connect_mcp_tool", connect_mcp_tool) - - harness = make_harness( - toolsets=[ - vf.Toolset(tools=[vf.MCPTool("global_tool")], scope="global"), - vf.Toolset(tools=[vf.MCPTool("rollout_tool")], scope="rollout"), - vf.Toolset(tools=[vf.MCPTool("group_tool")], scope="group"), - ] - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state_a = vf.State.for_task(task) - state_a["runtime"]["group_key"] = "group" - state_b = vf.State.for_task(task) - state_b["runtime"]["group_key"] = "group" - - await harness.runtime.ensure_mcp_tools(state_a) - await harness.runtime.ensure_mcp_tools(state_b) - - keys = sorted(harness.runtime.mcp_exit_stacks) - assert len([key for key in keys if key.startswith("global:")]) == 1 - assert len([key for key in keys if key.startswith("group:")]) == 1 - assert len([key for key in keys if key.startswith("rollout:")]) == 2 - assert sorted(harness.runtime.all_exposed_tools(state_a)) == [ - "global_tool", - "group_tool", - "rollout_tool", - ] - - await harness.runtime.close_mcp_tools(state_a) - - keys = sorted(harness.runtime.mcp_exit_stacks) - assert len([key for key in keys if key.startswith("global:")]) == 1 - assert len([key for key in keys if key.startswith("group:")]) == 1 - assert len([key for key in keys if key.startswith("rollout:")]) == 1 - - await harness.runtime.close_mcp_tools(state_b) - await harness.runtime.cleanup_group([task, task], [state_a, state_b]) - - keys = sorted(harness.runtime.mcp_exit_stacks) - assert len([key for key in keys if key.startswith("global:")]) == 1 - assert not [key for key in keys if key.startswith("group:")] - assert not [key for key in keys if key.startswith("rollout:")] - - await harness.teardown() - assert harness.runtime.mcp_exit_stacks == {} - - -@pytest.mark.asyncio -async def test_shared_sandbox_delete_retries_transient_failures( - monkeypatch: pytest.MonkeyPatch, -) -> None: - disable_sandbox_retry_sleep(monkeypatch) - - class FlakyDeleteClient: - closed = 0 - - def __init__(self) -> None: - self.delete_calls = 0 - - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == "sbx-1" - self.delete_calls += 1 - if self.delete_calls == 1: - raise RuntimeError("transient delete") - - async def aclose(self) -> None: - type(self).closed += 1 - - client = FlakyDeleteClient() - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, client), - "sbx-1", - "rollout", - "program", - ) - - await lease.delete() - await lease.delete() - - assert client.delete_calls == 2 - assert client.closed == 1 - - -@pytest.mark.asyncio -async def test_sandbox_delete_failure_leaves_lease_retryable( - monkeypatch: pytest.MonkeyPatch, -) -> None: - disable_sandbox_retry_sleep(monkeypatch) - - class DeleteFailsThenSucceeds: - calls = 0 - - async def delete(self, sandbox_id: str) -> None: - _ = sandbox_id - self.calls += 1 - if self.calls <= sandbox_utils.SANDBOX_RETRY_ATTEMPTS: - raise RuntimeError("delete failed") - - client = DeleteFailsThenSucceeds() - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, client), - "sbx-1", - "rollout", - "program", - owns_client=False, - ) - - with pytest.raises(RuntimeError, match="delete failed"): - await lease.delete() - - assert lease.deleted is False - - await lease.delete() - - assert lease.deleted is True - assert client.calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS + 1 - - -@pytest.mark.asyncio -async def test_owned_sandbox_delete_failure_keeps_client_retryable( - monkeypatch: pytest.MonkeyPatch, -) -> None: - disable_sandbox_retry_sleep(monkeypatch) - - class DeleteFailsThenSucceeds: - calls = 0 - closed = 0 - - async def delete(self, sandbox_id: str) -> None: - _ = sandbox_id - self.calls += 1 - if self.calls <= sandbox_utils.SANDBOX_RETRY_ATTEMPTS: - raise RuntimeError("delete failed") - - async def aclose(self) -> None: - self.closed += 1 - - client = DeleteFailsThenSucceeds() - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, client), - "sbx-1", - "rollout", - "program", - ) - - with pytest.raises(RuntimeError, match="delete failed"): - await lease.delete() - - assert lease.deleted is False - assert client.calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS - assert client.closed == 0 - - await lease.delete() - - assert lease.deleted is True - assert client.calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS + 1 - assert client.closed == 1 - - -@pytest.mark.asyncio -async def test_program_sandbox_creations_are_concurrent_and_bounded( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - - class SlowSandboxClient(FakeSandboxClient): - active = 0 - max_active = 0 - - def __init__(self, *args: object, **kwargs: object) -> None: - _ = args, kwargs - - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - type(self).active += 1 - type(self).max_active = max(type(self).max_active, type(self).active) - await asyncio.sleep(0.01) - try: - return await super().create(request) - finally: - type(self).active -= 1 - - monkeypatch.setattr( - "verifiers.utils.threaded_sandbox_client.ThreadedAsyncSandboxClient", - SlowSandboxClient, - ) - - class RecordingRateLimiter: - wait_calls = 0 - - async def wait(self) -> None: - self.wait_calls += 1 - - harness = make_harness(sandbox={"create_concurrency": 2}) - limiter = RecordingRateLimiter() - harness.runtime.sandbox_create_rate_limiter = limiter - sandbox = vf.SandboxConfig(image="python:3.11-slim") - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - states = [vf.State.for_task(task) for _ in range(4)] - - await asyncio.gather( - *( - harness.runtime.resolve_program_sandbox(sandbox, task, state) - for state in states - ) - ) - - assert SlowSandboxClient.max_active == 2 - assert len(FakeSandboxClient.created) == 4 - assert limiter.wait_calls == 4 - - await harness.teardown() - - -@pytest.mark.asyncio -async def test_cancelled_sandbox_awaiter_teardown_deletes_completed_creation() -> None: - deleted: list[str] = [] - started = asyncio.Event() - finish = asyncio.Event() - - class DeleteClient: - async def delete(self, sandbox_id: str) -> None: - deleted.append(sandbox_id) - - async def create_late_lease() -> sandbox_utils.SandboxLease: - started.set() - await finish.wait() - return sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, DeleteClient()), - "sbx-late", - "rollout", - "program", - owns_client=False, - ) - - runtime = Runtime() - key = ("rollout:test", "program") - waiter = asyncio.create_task(runtime.resolve_sandbox_lease(key, create_late_lease)) - await started.wait() - - waiter.cancel() - with pytest.raises(asyncio.CancelledError): - await waiter - - assert key in runtime.sandbox_creation_tasks - finish.set() - await asyncio.wait_for(runtime.sandbox_creation_tasks[key], timeout=1) - - await runtime.teardown() - - assert deleted == ["sbx-late"] - assert key not in runtime.sandbox_creation_tasks - - -@pytest.mark.asyncio -async def test_teardown_deletes_late_provider_create_before_closing_client() -> None: - started = asyncio.Event() - finish = asyncio.Event() - delete_closed_states: list[bool] = [] - - class SlowCreateClient: - closed = False - - async def create(self, request: FakeCreateSandboxRequest) -> FakeSandboxResult: - _ = request - started.set() - await finish.wait() - return FakeSandboxResult("sbx-late-provider") - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = sandbox_utils.SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> None: - _ = sandbox_id, max_attempts - - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == "sbx-late-provider" - delete_closed_states.append(self.closed) - - async def aclose(self) -> None: - self.closed = True - - client = SlowCreateClient() - runtime = Runtime() - runtime._sandbox_client = cast(sandbox_utils.SandboxClient, client) - sandbox = vf.SandboxConfig(image="python:3.11-slim") - key = ("rollout:test", "program") - waiter = asyncio.create_task( - runtime.resolve_sandbox_lease( - key, - lambda: sandbox_utils.create_sandbox_lease( - sandbox, - key[1], - client=cast(sandbox_utils.SandboxClient, client), - ), - ) - ) - await started.wait() - - teardown = asyncio.create_task(runtime.teardown()) - await asyncio.sleep(0) - assert teardown.done() is False - - finish.set() - await teardown - - assert delete_closed_states == [False] - assert client.closed is True - with pytest.raises(asyncio.CancelledError): - await waiter - - -@pytest.mark.asyncio -async def test_release_sandboxes_deletes_completed_unclaimed_creation() -> None: - deleted: list[str] = [] - - class DeleteClient: - async def delete(self, sandbox_id: str) -> None: - deleted.append(sandbox_id) - - runtime = Runtime() - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - key = (runtime.scope_key("rollout", state), "program") - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, DeleteClient()), - "sbx-unclaimed", - "rollout", - "program", - owns_client=False, - ) - runtime.sandbox_creation_tasks[key] = asyncio.create_task( - asyncio.sleep(0, result=lease) - ) - await runtime.sandbox_creation_tasks[key] - - await runtime.release_sandboxes("rollout", state) - await runtime.teardown() - - assert deleted == ["sbx-unclaimed"] - assert key not in runtime.sandbox_creation_tasks - - -@pytest.mark.asyncio -async def test_release_sandboxes_keeps_failed_delete_retryable( - monkeypatch: pytest.MonkeyPatch, -) -> None: - disable_sandbox_retry_sleep(monkeypatch) - - class RetryableDeleteClient: - calls = 0 - - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == "sbx-retryable-delete" - self.calls += 1 - if self.calls <= sandbox_utils.SANDBOX_RETRY_ATTEMPTS: - raise RuntimeError("delete failed") - - runtime = Runtime() - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - key = (runtime.scope_key("rollout", state), "program") - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, RetryableDeleteClient()), - "sbx-retryable-delete", - "rollout", - "program", - owns_client=False, - ) - runtime.sandbox_leases[key] = lease - - await runtime.release_sandboxes("rollout", state) - - assert runtime.sandbox_leases[key] is lease - assert len(state["cleanup_errors"]) == 1 - - await runtime.release_sandboxes("rollout", state) - await runtime.teardown() - - assert key not in runtime.sandbox_leases - - -@pytest.mark.asyncio -async def test_resolve_sandbox_lease_rejects_lease_being_deleted() -> None: - delete_started = asyncio.Event() - finish_delete = asyncio.Event() - - class SlowDeleteClient: - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == "sbx-deleting" - delete_started.set() - await finish_delete.wait() - - async def create_replacement() -> sandbox_utils.SandboxLease: - raise AssertionError("resolve should not create a replacement") - - runtime = Runtime() - key = ("rollout:test", "program") - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, SlowDeleteClient()), - "sbx-deleting", - "rollout", - "program", - owns_client=False, - ) - runtime.sandbox_leases[key] = lease - deletion = asyncio.create_task(runtime.close_sandbox_lease(lease)) - await delete_started.wait() - - with pytest.raises(RuntimeError, match="being deleted"): - await runtime.resolve_sandbox_lease(key, create_replacement) - - finish_delete.set() - await deletion - - -@pytest.mark.asyncio -async def test_resolve_sandbox_lease_rejects_creation_claimed_by_cleanup() -> None: - deleted: list[str] = [] - - class DeleteClient: - async def delete(self, sandbox_id: str) -> None: - deleted.append(sandbox_id) - - runtime = Runtime() - key = ("rollout:test", "program") - create_started = asyncio.Event() - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, DeleteClient()), - "sbx-cleanup-claimed", - "rollout", - "program", - owns_client=False, - ) - - async def create_lease() -> sandbox_utils.SandboxLease: - create_started.set() - try: - await asyncio.sleep(10) - except asyncio.CancelledError: - return lease - raise AssertionError("creation should be cancelled by cleanup") - - resolver = asyncio.create_task(runtime.resolve_sandbox_lease(key, create_lease)) - await create_started.wait() - creation_task = runtime.sandbox_creation_tasks[key] - - await runtime.clear_sandbox_creation_tasks([(key, creation_task)]) - - with pytest.raises(RuntimeError, match="cancelled before"): - await resolver - assert deleted == ["sbx-cleanup-claimed"] - assert key not in runtime.sandbox_creation_tasks - assert key not in runtime.sandbox_leases - - -@pytest.mark.asyncio -async def test_clear_creation_tasks_keeps_failed_delete_retryable( - monkeypatch: pytest.MonkeyPatch, -) -> None: - disable_sandbox_retry_sleep(monkeypatch) - - class RetryableDeleteClient: - calls = 0 - - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == "sbx-unclaimed-retry" - self.calls += 1 - if self.calls <= sandbox_utils.SANDBOX_RETRY_ATTEMPTS: - raise RuntimeError("delete failed") - - runtime = Runtime() - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - key = (runtime.scope_key("rollout", state), "program") - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, RetryableDeleteClient()), - "sbx-unclaimed-retry", - "rollout", - "program", - owns_client=False, - ) - - async def finished_creation() -> sandbox_utils.SandboxLease: - return lease - - creation_task = asyncio.create_task(finished_creation()) - runtime.sandbox_creation_tasks[key] = creation_task - await creation_task - - await runtime.clear_sandbox_creation_tasks( - [(key, creation_task)], - state=state, - scope="rollout", - ) - - assert key not in runtime.sandbox_creation_tasks - assert runtime.sandbox_leases[key] is lease - assert len(state["cleanup_errors"]) == 1 - - await runtime.release_sandboxes("rollout", state) - await runtime.teardown() - - assert key not in runtime.sandbox_leases - - -@pytest.mark.asyncio -async def test_teardown_keeps_failed_delete_retryable( - monkeypatch: pytest.MonkeyPatch, -) -> None: - disable_sandbox_retry_sleep(monkeypatch) - - class DeleteFailsThenSucceeds: - def __init__(self, sandbox_id: str): - self.sandbox_id = sandbox_id - self.delete_calls = 0 - self.closed = False - - async def delete(self, sandbox_id: str) -> None: - assert sandbox_id == self.sandbox_id - self.delete_calls += 1 - if self.delete_calls <= sandbox_utils.SANDBOX_RETRY_ATTEMPTS: - raise RuntimeError("delete failed") - - async def aclose(self) -> None: - self.closed = True - - client = DeleteFailsThenSucceeds("sbx-terminal-fail") - owned_client = DeleteFailsThenSucceeds("sbx-owned-terminal-fail") - runtime = Runtime() - runtime._sandbox_client = cast(sandbox_utils.SandboxClient, client) - key = ("rollout:test", "program") - lease = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, client), - "sbx-terminal-fail", - "rollout", - "program", - owns_client=False, - ) - runtime.sandbox_leases[key] = lease - owned_key = ("rollout:test", "owned") - runtime.sandbox_leases[owned_key] = sandbox_utils.SandboxLease( - cast(sandbox_utils.SandboxClient, owned_client), - "sbx-owned-terminal-fail", - "rollout", - "owned", - ) - - await runtime.teardown() - - assert client.delete_calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS - assert client.closed is False - assert owned_client.delete_calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS - assert owned_client.closed is False - assert runtime.sandbox_leases[key] is lease - assert owned_key in runtime.sandbox_leases - assert load_runtime(runtime.runtime_id) is runtime - - await runtime.teardown() - - assert client.delete_calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS + 1 - assert client.closed is True - assert owned_client.delete_calls == sandbox_utils.SANDBOX_RETRY_ATTEMPTS + 1 - assert owned_client.closed is True - assert runtime.sandbox_leases == {} - with pytest.raises(RuntimeError, match="No live v1 runtime registered"): - load_runtime(runtime.runtime_id) - - -@pytest.mark.asyncio -async def test_program_sandbox_group_scope_reuses_and_cleans( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={"sandbox": True, "command": ["python", "-c", "print('ok')"]}, - sandbox={"image": "python:3.11-slim", "scope": "group"}, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state_a = vf.State.for_task(task) - state_a["runtime"]["group_key"] = "group" - state_b = vf.State.for_task(task) - state_b["runtime"]["group_key"] = "group" - - state_a, state_b = await asyncio.gather( - harness.run(task, state_a), - harness.run(task, state_b), - ) - - assert FakeSandboxClient.created == ["sbx-1"] - assert state_a["sandbox_id"] == "sbx-1" - assert state_b["sandbox_id"] == "sbx-1" - assert FakeSandboxClient.deleted == [] - - await harness.cleanup_group([task, task], [state_a, state_b]) - - assert FakeSandboxClient.deleted == ["sbx-1"] - assert FakeSandboxClient.closed == 0 - assert "resolved" not in state_a.get("runtime", {}) - assert "resolved" not in state_b.get("runtime", {}) - assert "lease_key" not in state_a.get("runtime", {}).get("sandbox", {}) - assert "lease_key" not in state_b.get("runtime", {}).get("sandbox", {}) - - await harness.teardown() - - assert FakeSandboxClient.closed == 1 - - -@pytest.mark.asyncio -async def test_program_sandbox_global_scope_lives_until_teardown( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={"sandbox": True, "command": ["python", "-c", "print('ok')"]}, - sandbox={"image": "python:3.11-slim", "scope": "global"}, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert FakeSandboxClient.created == ["sbx-1"] - assert state["sandbox_id"] == "sbx-1" - assert FakeSandboxClient.deleted == [] - - await harness.teardown() - - assert FakeSandboxClient.deleted == ["sbx-1"] - assert FakeSandboxClient.closed == 1 - assert "resolved" not in state.get("runtime", {}) - - -@pytest.mark.asyncio -async def test_upload_program_dirs_reuses_runtime_archive_cache( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path -) -> None: - source_dir = tmp_path / "source" - source_dir.mkdir() - (source_dir / "module.py").write_text("VALUE = 1\n") - archive_path = tmp_path / "cached.tar.gz" - build_calls = 0 - - def fake_build_dir_archive(local_source: Path, remote_path: str) -> Path: - nonlocal build_calls - assert local_source == source_dir - assert remote_path == "/remote/pkg" - build_calls += 1 - time.sleep(0.05) - archive_path.write_bytes(b"archive") - return archive_path - - class UploadClient: - uploads: list[tuple[str, str, str]] = [] - commands: list[tuple[str, str]] = [] - - async def upload_file(self, *args: object, **kwargs: object) -> None: - sandbox_id = str(kwargs.get("sandbox_id") or args[0]) - file_path = str(kwargs.get("file_path") or args[1]) - local_path = str(kwargs.get("local_file_path") or args[2]) - self.uploads.append((sandbox_id, file_path, local_path)) - - async def execute_command( - self, *args: object, **kwargs: object - ) -> FakeCommandResult: - sandbox_id = str(kwargs.get("sandbox_id") or args[0]) - command = str(kwargs.get("command") or args[1]) - self.commands.append((sandbox_id, command)) - return FakeCommandResult() - - monkeypatch.setattr(sandbox_utils, "build_dir_archive", fake_build_dir_archive) - runtime = Runtime() - client = UploadClient() - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - program = {"dirs": {"/remote/pkg": str(source_dir)}} - states = [vf.State.for_task(task), vf.State.for_task(task)] - - await asyncio.gather( - *( - upload_program_dirs( - cast(sandbox_utils.SandboxClient, client), - sandbox_id, - program, - task, - state, - runtime, - ) - for sandbox_id, state in zip(["sbx-1", "sbx-2"], states, strict=True) - ) - ) - - assert build_calls == 1 - assert client.uploads == [ - ("sbx-1", "/tmp/_vf_upload_remote_pkg.tar.gz", str(archive_path)), - ("sbx-2", "/tmp/_vf_upload_remote_pkg.tar.gz", str(archive_path)), - ] - assert archive_path.exists() - - (source_dir / "module.py").write_text("VALUE = 2\n") - await upload_program_dirs( - cast(sandbox_utils.SandboxClient, client), - "sbx-3", - program, - task, - vf.State.for_task(task), - runtime, - ) - - assert build_calls == 2 - assert client.uploads[-1] == ( - "sbx-3", - "/tmp/_vf_upload_remote_pkg.tar.gz", - str(archive_path), - ) - - await runtime.teardown() - - assert not archive_path.exists() - - -@pytest.mark.asyncio -async def test_cached_upload_archive_cancelled_awaiter_still_cleans_archive( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path -) -> None: - source_dir = tmp_path / "source" - source_dir.mkdir() - (source_dir / "module.py").write_text("VALUE = 1\n") - archive_path = tmp_path / "cached.tar.gz" - started = threading.Event() - finish = threading.Event() - - def fake_build_dir_archive(local_source: Path, remote_path: str) -> Path: - assert local_source == source_dir - assert remote_path == "/remote/pkg" - started.set() - finish.wait(timeout=1) - archive_path.write_bytes(b"archive") - return archive_path - - monkeypatch.setattr(sandbox_utils, "build_dir_archive", fake_build_dir_archive) - runtime = Runtime() - task = asyncio.create_task(runtime.cached_upload_archive(source_dir, "/remote/pkg")) - await asyncio.to_thread(started.wait) - - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - - finish.set() - await runtime.teardown() - - assert runtime.upload_archive_tasks == {} - assert not archive_path.exists() - - -@pytest.mark.asyncio -async def test_cleanup_upload_archives_logs_unlink_errors( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path -) -> None: - archive_path = tmp_path / "cached.tar.gz" - archive_path.write_bytes(b"archive") - runtime = Runtime() - runtime.upload_archive_tasks[("remote", "source", "digest")] = asyncio.create_task( - asyncio.sleep(0, result=archive_path) - ) - warnings: list[tuple[tuple[object, ...], dict[str, object]]] = [] - - def raise_unlink(path: Path, missing_ok: bool = False) -> None: - assert path == archive_path - assert missing_ok is True - raise PermissionError("locked") - - def record_warning(*args: object, **kwargs: object) -> None: - warnings.append((args, kwargs)) - - monkeypatch.setattr(Path, "unlink", raise_unlink) - monkeypatch.setattr("verifiers.v1.runtime.logger.warning", record_warning) - - await runtime.cleanup_upload_archives() - - assert runtime.upload_archive_tasks == {} - assert warnings[0][0][0] == "Failed to delete cached upload archive %s: %s" - assert warnings[0][1]["exc_info"] is True - - -@pytest.mark.asyncio -async def test_sandbox_program_artifact_collected_by_runtime( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={ - "command": ["true"], - "sandbox": True, - "artifacts": {"command_log": {"path": "/tmp/command.log"}}, - }, - sandbox={"image": "python:3.11-slim"}, - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - - state = await harness.run(task) - - assert state["artifacts"]["command_log"] == "ok\n" - - -@pytest.mark.asyncio -async def test_optional_toolset_artifact_does_not_create_owner_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - toolset = vf.Toolset( - tools=[program_sandbox_id], - sandbox=vf.SandboxConfig(image="python:3.11-slim"), - artifacts=vf.ArtifactsConfig.model_validate( - {"tool_log": {"path": "/tmp/tool.log", "optional": True}} - ), - ) - harness = make_harness(toolsets=[toolset]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, vf.State.for_task(task)) - - await harness.runtime.collect_artifacts(task, state) - await harness.runtime.cleanup_rollout(task, state) - - assert state["artifacts"]["tool_log"] is None - assert FakeSandboxClient.created == [] - - -@pytest.mark.asyncio -async def test_toolset_artifact_reads_owned_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - toolset = vf.Toolset( - tools=[program_sandbox_id], - sandbox=vf.SandboxConfig(image="python:3.11-slim"), - artifacts=vf.ArtifactsConfig.model_validate( - {"tool_log": {"path": "/tmp/tool.log"}} - ), - ) - harness = make_harness(toolsets=[toolset]) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, vf.State.for_task(task)) - - await harness.runtime.call_tool("program_sandbox_id", task, state) - await harness.runtime.collect_artifacts(task, state) - await harness.runtime.cleanup_rollout(task, state) - - assert state["artifacts"]["tool_log"] == "ok\n" - assert FakeSandboxClient.created == ["sbx-1"] - assert FakeSandboxClient.deleted == ["sbx-1"] - - -@pytest.mark.asyncio -async def test_toolset_can_bind_to_primary_program_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={"sandbox": True, "command": ["python", "-c", "print('ok')"]}, - sandbox={"image": "python:3.11-slim", "scope": "group"}, - toolsets=[vf.Toolset(tools=[program_sandbox_id], sandbox="program")], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["group_key"] = "group" - - state = await harness.run(task, state) - result = await harness.runtime.call_tool("program_sandbox_id", task, state) - - assert result == state["sandbox_id"] - - await harness.cleanup_group([task], [state]) - - -@pytest.mark.asyncio -async def test_toolset_sandbox_prefer_program_falls_back_to_owned_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - - harness = make_harness( - toolsets=[ - vf.Toolset( - tools=[program_sandbox_id], - sandbox=vf.SandboxConfig( - prefer="program", - image="python:3.11-slim", - scope="rollout", - ), - ) - ] - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = await harness.setup_state(task, vf.State.for_task(task)) - - result = await state.get_tools()["program_sandbox_id"]() - - assert result == "sbx-1" - assert FakeSandboxClient.created == ["sbx-1"] - - await harness.runtime.cleanup_rollout(task, state) - - assert FakeSandboxClient.deleted == ["sbx-1"] - - -@pytest.mark.asyncio -async def test_toolset_sandbox_prefer_program_uses_active_program_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - harness = make_harness( - program={"sandbox": True, "command": ["python", "-c", "print('ok')"]}, - sandbox={"image": "python:3.11-slim", "scope": "group"}, - toolsets=[ - vf.Toolset( - tools=[program_sandbox_id], - sandbox=vf.SandboxConfig( - prefer="program", - image="python:3.11-slim", - scope="rollout", - ), - ) - ], - ) - task = vf.Task({"prompt": [{"role": "user", "content": "hi"}]}).freeze() - state = vf.State.for_task(task) - state["runtime"]["group_key"] = "group" - - state = await harness.run(task, state) - result = await harness.runtime.call_tool("program_sandbox_id", task, state) - - assert result == state["sandbox_id"] - assert FakeSandboxClient.created == ["sbx-1"] - - await harness.cleanup_group([task], [state]) - - assert FakeSandboxClient.deleted == ["sbx-1"] - - -@pytest.mark.asyncio -async def test_child_state_can_borrow_primary_program_sandbox( - monkeypatch: pytest.MonkeyPatch, -) -> None: - install_fake_sandboxes(monkeypatch) - install_fake_endpoint_tunnel(monkeypatch) - - parent = make_harness( - program={"sandbox": True, "command": ["python", "-c", "print('ok')"]}, - sandbox={"image": "python:3.11-slim", "scope": "group"}, - ) - child = make_harness( - program={"fn": program_ref("child_reads_program_sandbox")}, - toolsets=[vf.Toolset(tools=[program_sandbox_id], sandbox="program")], - ) - parent_task = vf.Task({"prompt": [{"role": "user", "content": "parent"}]}).freeze() - parent_state = vf.State.for_task(parent_task) - parent_state["runtime"]["group_key"] = "group" - - parent_state = await parent.run(parent_task, parent_state) - child_task = vf.Task({"prompt": [{"role": "user", "content": "child"}]}).freeze() - child_state = parent_state.for_task(child_task, borrow="sandbox") - child_state = await child.run(child_task, child_state) - - assert parent_state["sandbox_id"] == "sbx-1" - assert child_state["borrowed_sandbox_id"] == "sbx-1" - assert FakeSandboxClient.deleted == [] - - await parent.cleanup_group([parent_task], [parent_state]) - - assert FakeSandboxClient.deleted == ["sbx-1"] diff --git a/tests/test_v1_scoring_functions.py b/tests/test_v1_scoring_functions.py deleted file mode 100644 index 246798a898..0000000000 --- a/tests/test_v1_scoring_functions.py +++ /dev/null @@ -1,163 +0,0 @@ -from typing import Any, cast - -import pytest - -import verifiers as vf -from verifiers import ( - add_advantage, - add_metric, - add_reward, - build_signals, - collect_signals, - score_group, - score_rollout, -) - - -@vf.metric -async def num_tool_calls(task: dict, state: dict) -> float: - return float(len(state.get("tool_calls", []))) - - -@vf.metric -async def config_metric(task: dict, state: dict) -> float: - return float(task["x"] + state["y"]) - - -@vf.reward(weight=2.0) -async def exact_answer(task: dict, state: dict) -> float: - return float(state.get("answer") == task["answer"]) - - -@vf.reward(stage="group") -async def best_answer_bonus(tasks: list[dict], states: list[dict]) -> list[float]: - return [ - float(state.get("answer") == task["answer"]) - for task, state in zip(tasks, states) - ] - - -@vf.advantage -async def explicit_advantage(tasks: list[dict], states: list[dict]) -> list[float]: - _ = tasks - return [float(index) for index, _ in enumerate(states)] - - -def task_and_state( - task_data: dict[str, Any], state_data: dict[str, Any] -) -> tuple[vf.Task, vf.State]: - task = vf.Task(task_data).freeze() - state = vf.State.for_task(task) - state.update(state_data) - return task, state - - -@pytest.mark.asyncio -async def test_programmatic_metric_and_reward_share_signal_path() -> None: - signals = build_signals() - add_metric(signals, num_tool_calls) - add_reward(signals, exact_answer) - task, state = task_and_state( - {"answer": "4"}, {"answer": "4", "tool_calls": ["a", "b"]} - ) - - await score_rollout(signals, task, state) - - metrics = cast(dict[str, float], state["metrics"]) - assert state["reward"] == 2.0 - assert metrics["exact_answer"] == 1.0 - assert metrics["num_tool_calls"] == 2.0 - - -@pytest.mark.asyncio -async def test_config_overrides_default_signal_metadata() -> None: - signals = build_signals( - scoring={"exact_answer": {"weight": 0.5}}, - rewards=[exact_answer], - ) - task, state = task_and_state({"answer": "4"}, {"answer": "4"}) - - await score_rollout(signals, task, state) - - assert state["reward"] == 0.5 - - -@pytest.mark.asyncio -async def test_config_tunes_imported_signal_by_name() -> None: - signals = build_signals( - metrics=[config_metric], - scoring={"config_metric": {"priority": 10}}, - ) - task, state = task_and_state({"x": 2}, {"y": 3}) - - await score_rollout(signals, task, state) - - metrics = cast(dict[str, float], state["metrics"]) - assert metrics["config_metric"] == 5.0 - - -def test_signal_name_collisions_hard_fail() -> None: - taskset_signals = build_signals(metrics=[num_tool_calls]) - harness_signals = build_signals(metrics=[num_tool_calls]) - - with pytest.raises(ValueError, match="defined twice"): - collect_signals(taskset_signals, harness_signals) - - -@pytest.mark.asyncio -async def test_group_signal_reports_unresolved_required_args() -> None: - @vf.metric(stage="group") - async def bad_group_metric(task: dict, state: dict) -> float: - return 0.0 - - signals = build_signals(metrics=[bad_group_metric]) - - with pytest.raises(TypeError, match="metric signal 'bad_group_metric'.*task"): - await score_group(signals, [{"answer": "a"}], [{"answer": "a"}]) - - -@pytest.mark.asyncio -async def test_group_reward_scores_each_state() -> None: - signals = build_signals(rewards=[best_answer_bonus]) - task_a, state_a = task_and_state( - {"answer": "a"}, - {"answer": "a", "trajectory": [{"advantage": None}, {"advantage": 9.0}]}, - ) - task_b, state_b = task_and_state( - {"answer": "b"}, {"answer": "c", "trajectory": [{"advantage": None}]} - ) - tasks = [task_a, task_b] - states = [state_a, state_b] - - await score_group(signals, tasks, states) - - assert states[0]["reward"] == 1.0 - assert states[1]["reward"] == 0.0 - assert "advantage" not in states[0] - assert "advantage" not in states[1] - trajectory = cast(list[dict[str, Any]], states[0]["trajectory"]) - assert trajectory[0]["advantage"] is None - assert trajectory[1]["advantage"] == 9.0 - - -@pytest.mark.asyncio -async def test_advantage_signal_writes_group_advantages() -> None: - signals = build_signals(rewards=[best_answer_bonus]) - add_advantage(signals, explicit_advantage) - task_a, state_a = task_and_state({"answer": "a"}, {"answer": "a"}) - task_b, state_b = task_and_state({"answer": "b"}, {"answer": "c"}) - tasks = [task_a, task_b] - states = [state_a, state_b] - - await score_group(signals, tasks, states) - - assert states[0]["advantage"] == 0.0 - assert states[1]["advantage"] == 1.0 - - -def test_advantage_requires_group_plural_args() -> None: - async def bad_advantage(task: dict, state: dict) -> float: - return 0.0 - - with pytest.raises(ValueError, match="stage='group'"): - build_signals(advantages=[bad_advantage]) diff --git a/tests/test_v1_taskset_bindings.py b/tests/test_v1_taskset_bindings.py deleted file mode 100644 index 3a8b785256..0000000000 --- a/tests/test_v1_taskset_bindings.py +++ /dev/null @@ -1,374 +0,0 @@ -import re -import sys -from types import ModuleType - -import pytest -from pydantic import BaseModel - -import verifiers as vf - - -REF_MODULE = "v1_taskset_binding_refs" -ref_module = ModuleType(REF_MODULE) -sys.modules[REF_MODULE] = ref_module - - -def ref(name: str) -> str: - return f"{REF_MODULE}:{name}" - - -def config_data(config: object | None) -> dict[str, object]: - if config is None: - return {} - if isinstance(config, BaseModel): - return config.model_dump(exclude_none=True) - if isinstance(config, dict): - return {str(key): item for key, item in config.items()} - raise TypeError("test config must be a mapping or config object") - - -def make_taskset(config: object | None = None, **values: object) -> vf.Taskset: - data = {**config_data(config), **values} - return BindingTaskset(config=BindingTasksetConfig.model_validate(data)) - - -def load_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - return [ - { - "prompt": [{"role": "user", "content": "reply ok"}], - "answer": "ok", - } - ] - - -def load_prefixed_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - return [ - { - "prompt": [{"role": "user", "content": "reply ok"}], - "answer": "ok", - "prefix": "bound:", - } - ] - - -def load_two_prefixed_tasks(split: vf.TaskSplit = "train") -> list[dict[str, object]]: - return [ - { - "prompt": [{"role": "user", "content": "reply ok"}], - "answer": "ok", - "prefix": "first:", - }, - { - "prompt": [{"role": "user", "content": "reply ok"}], - "answer": "ok", - "prefix": "second:", - }, - ] - - -class BindingTasksetConfig(vf.TasksetConfig): - dataset: str = "default" - - -class BindingTaskset(vf.Taskset[BindingTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if self.config.dataset == "prefixed": - return load_prefixed_tasks(split) - if self.config.dataset == "two_prefixed": - return load_two_prefixed_tasks(split) - return load_tasks(split) - - -class Prefixer: - def __init__(self, prefix: str): - self.prefix = prefix - - def __call__(self, value: str) -> str: - return f"{self.prefix}{value}" - - -class TagExtractor: - def __init__(self, tag: str): - self.pattern = re.compile(rf"<{tag}>(.*?)", re.DOTALL) - - def __call__(self, completion: list[dict[str, object]]) -> str: - message = vf.get_messages(completion, role="assistant")[-1] - match = self.pattern.search(str(message.content or "")) - return "" if match is None else match.group(1).strip() - - -prefixer_factory_calls = 0 - - -def load_factory_prefixer() -> Prefixer: - global prefixer_factory_calls - prefixer_factory_calls += 1 - return Prefixer("factory:") - - -def load_config_prefixer() -> Prefixer: - return Prefixer("config:") - - -def load_defaulted_prefixer(prefix: str = "defaulted:") -> Prefixer: - return Prefixer(prefix) - - -def load_bound_prefixer(prefix: str) -> Prefixer: - return Prefixer(prefix) - - -def load_token() -> str: - return "bound" - - -def load_answer_extractor() -> TagExtractor: - return TagExtractor("answer") - - -@vf.reward -async def prefix_reward(state, prefixer) -> float: - state["prefixed"] = prefixer("ok") - return 1.0 - - -@vf.reward -async def framework_state_reward(state) -> float: - state["framework_state_seen"] = True - return 1.0 - - -@vf.setup -async def setup_with_override(state, token) -> None: - state["token"] = token - - -@vf.reward -async def missing_binding_reward(state, extractor) -> float: - _ = state, extractor - return 0.0 - - -@vf.reward -async def extracted_answer_reward(task, state, extract_answer) -> float: - response = extract_answer(state.get("completion") or []) - return float(response == task["answer"]) - - -async def score_taskset(taskset: vf.Taskset) -> vf.State: - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - task = next(iter(taskset)) - state = await env.harness.setup_state(task, vf.State.for_task(task)) - await env.harness.runtime.score_rollout(task, state) - return state - - -for _name, _value in { - "load_tasks": load_tasks, - "load_prefixed_tasks": load_prefixed_tasks, - "load_two_prefixed_tasks": load_two_prefixed_tasks, - "load_factory_prefixer": load_factory_prefixer, - "load_config_prefixer": load_config_prefixer, - "load_defaulted_prefixer": load_defaulted_prefixer, - "load_bound_prefixer": load_bound_prefixer, - "load_token": load_token, - "load_answer_extractor": load_answer_extractor, - "prefix_reward": prefix_reward, - "framework_state_reward": framework_state_reward, - "setup_with_override": setup_with_override, - "missing_binding_reward": missing_binding_reward, - "extracted_answer_reward": extracted_answer_reward, -}.items(): - setattr(ref_module, _name, _value) - - -def test_taskset_object_binding_rejects_live_instance() -> None: - with pytest.raises(TypeError, match="prefixer"): - make_taskset( - rewards=[ref("prefix_reward")], - objects={"prefixer": Prefixer("inst:")}, - bindings={"prefix_reward.prefixer": "objects.prefixer"}, - ) - - -@pytest.mark.asyncio -async def test_taskset_object_factory_is_lazy_and_resolved_once() -> None: - global prefixer_factory_calls - prefixer_factory_calls = 0 - - taskset = make_taskset( - rewards=[ref("prefix_reward")], - objects={"prefixer": ref("load_factory_prefixer")}, - bindings={"prefix_reward.prefixer": "objects.prefixer"}, - ) - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - task = next(iter(taskset)) - state = await env.harness.setup_state(task, vf.State.for_task(task)) - - await env.harness.runtime.score_rollout(task, state) - await env.harness.runtime.score_rollout(task, state) - - assert prefixer_factory_calls == 1 - assert state["prefixed"] == "factory:ok" - - -@pytest.mark.asyncio -async def test_taskset_object_factory_accepts_defaulted_arguments() -> None: - taskset = make_taskset( - rewards=[ref("prefix_reward")], - objects={"prefixer": ref("load_defaulted_prefixer")}, - bindings={"prefix_reward.prefixer": "objects.prefixer"}, - ) - - state = await score_taskset(taskset) - - assert state["prefixed"] == "defaulted:ok" - - -@pytest.mark.asyncio -async def test_taskset_object_factory_accepts_bound_arguments() -> None: - taskset = make_taskset( - dataset="prefixed", - rewards=[ref("prefix_reward")], - objects={"prefixer": ref("load_bound_prefixer")}, - bindings={ - "prefixer.prefix": "task.prefix", - "prefix_reward.prefixer": "objects.prefixer", - }, - ) - - state = await score_taskset(taskset) - - assert state["prefixed"] == "bound:ok" - - -@pytest.mark.asyncio -async def test_taskset_object_factory_bindings_are_rollout_scoped() -> None: - taskset = make_taskset( - dataset="two_prefixed", - rewards=[ref("prefix_reward")], - objects={"prefixer": ref("load_bound_prefixer")}, - bindings={ - "prefixer.prefix": "task.prefix", - "prefix_reward.prefixer": "objects.prefixer", - }, - ) - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - first, second = list(taskset) - first_state = await env.harness.setup_state(first, vf.State.for_task(first)) - second_state = await env.harness.setup_state(second, vf.State.for_task(second)) - - await env.harness.runtime.score_rollout(first, first_state) - await env.harness.runtime.score_rollout(second, second_state) - - assert first_state["prefixed"] == "first:ok" - assert second_state["prefixed"] == "second:ok" - - -@pytest.mark.asyncio -async def test_taskset_object_factory_rejects_unbound_arguments() -> None: - taskset = make_taskset( - dataset="prefixed", - rewards=[ref("prefix_reward")], - objects={"prefixer": ref("load_bound_prefixer")}, - bindings={"prefix_reward.prefixer": "objects.prefixer"}, - ) - - with pytest.raises(TypeError, match="unbound factory arguments"): - await score_taskset(taskset) - - -@pytest.mark.asyncio -async def test_framework_args_win_over_taskset_bindings() -> None: - taskset = make_taskset( - rewards=[ref("framework_state_reward")], - bindings={"framework_state_reward.state": "objects.missing"}, - ) - - state = await score_taskset(taskset) - - assert state["framework_state_seen"] is True - - -@pytest.mark.asyncio -async def test_caller_kwargs_win_over_taskset_bindings_for_handlers() -> None: - taskset = make_taskset( - setups=[ref("setup_with_override")], - objects={"token": ref("load_token")}, - bindings={"setup_with_override.token": "objects.token"}, - ) - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - task = next(iter(taskset)) - state = vf.State.for_task(task) - - await env.harness.runtime.run_rollout_handlers( - [setup_with_override], task=task, state=state, token="explicit" - ) - - assert state["token"] == "explicit" - - -@pytest.mark.asyncio -async def test_missing_taskset_binding_error_names_signal_and_arg() -> None: - taskset = make_taskset(rewards=[ref("missing_binding_reward")]) - - with pytest.raises( - TypeError, - match="reward signal 'missing_binding_reward'.*extractor", - ): - await score_taskset(taskset) - - -@pytest.mark.asyncio -async def test_taskset_config_map_round_trips_objects_and_bindings() -> None: - config = vf.TasksetConfig( - objects={"prefixer": ref("load_config_prefixer")}, - bindings={"prefix_reward.prefixer": "objects.prefixer"}, - ) - taskset = make_taskset( - rewards=[ref("prefix_reward")], - config=config, - ) - - state = await score_taskset(taskset) - - assert state["prefixed"] == "config:ok" - - -@pytest.mark.asyncio -async def test_taskset_bindings_support_shared_extractor_pattern() -> None: - taskset = make_taskset( - rewards=[ref("extracted_answer_reward")], - objects={"extract_answer": ref("load_answer_extractor")}, - bindings={"extracted_answer_reward.extract_answer": "objects.extract_answer"}, - ) - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - task = next(iter(taskset)) - state = await env.harness.setup_state(task, vf.State.for_task(task)) - state["completion"] = [{"role": "assistant", "content": "ok"}] - - await env.harness.runtime.score_rollout(task, state) - - assert state["reward"] == 1.0 - - -@pytest.mark.asyncio -async def test_harness_object_factory_accepts_bound_arguments() -> None: - taskset = make_taskset(dataset="prefixed") - harness = vf.Harness( - config=vf.HarnessConfig( - rewards=[ref("prefix_reward")], - objects={"prefixer": ref("load_bound_prefixer")}, - bindings={ - "prefixer.prefix": "task.prefix", - "prefix_reward.prefixer": "objects.prefixer", - }, - ) - ) - env = vf.Env(taskset=taskset, harness=harness) - task = next(iter(taskset)) - state = await env.harness.setup_state(task, vf.State.for_task(task)) - - await env.harness.runtime.score_rollout(task, state) - - assert state["prefixed"] == "bound:ok" diff --git a/tests/test_v1_taskset_utils.py b/tests/test_v1_taskset_utils.py index 1498a2f248..b3badde1d1 100644 --- a/tests/test_v1_taskset_utils.py +++ b/tests/test_v1_taskset_utils.py @@ -1,33 +1,40 @@ -import json import sys import types from datasets import Dataset +import pytest -from verifiers.v1 import Env, Taskset -from verifiers.v1.utils.taskset_utils import dataset_from_result, discover_sibling_dir +from verifiers.v1 import Env, Task, Taskset +from verifiers.v1.eval import eval_inputs +from verifiers.v1.utils.taskset_utils import ( + dataset_from_result_typed, + discover_sibling_dir, + tasks_from_result_typed, +) -def task_payload(row: dict) -> dict: - return json.loads(row["info"]["task"]) +class ReverseTextTask(Task): + question: str + answer: str def test_dataset_from_result_assigns_example_id_to_iterable_records(): - dataset = dataset_from_result( + dataset = dataset_from_result_typed( [ {"question": "Reverse abc.", "answer": "cba"}, {"question": "Reverse xyz.", "answer": "zyx"}, ], - "ReverseTextTaskset", + ReverseTextTask, ) rows = list(dataset) - payloads = [task_payload(row) for row in rows] assert [row["example_id"] for row in rows] == [0, 1] - assert [payload["example_id"] for payload in payloads] == [0, 1] - assert all(len(payload["task_id"]) == 32 for payload in payloads) - assert {payload["task_id"] for payload in payloads}.isdisjoint({"0", "1"}) + assert [row["row_id"] for row in rows] == [0, 1] + assert [row["answer"] for row in rows] == ["cba", "zyx"] + assert all(len(row["task_id"]) == 24 for row in rows) + assert {row["task_id"] for row in rows}.isdisjoint({"0", "1"}) + assert rows[0]["task_id"] != rows[1]["task_id"] def test_dataset_from_result_overwrites_existing_example_id_column(): @@ -38,15 +45,39 @@ def test_dataset_from_result_overwrites_existing_example_id_column(): ] ) - dataset = dataset_from_result(raw_dataset, "ReverseTextTaskset") + dataset = dataset_from_result_typed(raw_dataset, ReverseTextTask) rows = list(dataset) - payloads = [task_payload(row) for row in rows] assert [row["example_id"] for row in rows] == [0, 1] - assert [payload["example_id"] for payload in payloads] == [0, 1] - assert all(len(payload["task_id"]) == 32 for payload in payloads) - assert {payload["task_id"] for payload in payloads}.isdisjoint({"0", "1", "99"}) + assert [row["row_id"] for row in rows] == [0, 1] + assert [row["answer"] for row in rows] == ["cba", "zyx"] + assert all(len(row["task_id"]) == 24 for row in rows) + assert {row["task_id"] for row in rows}.isdisjoint({"0", "1", "99"}) + assert rows[0]["task_id"] != rows[1]["task_id"] + + +def test_tasks_from_result_typed_validates_existing_task_objects(): + base_task = Task(prompt="Reverse abc.", row_id=3) + + with pytest.raises(ValueError): + tasks_from_result_typed([base_task], ReverseTextTask) + + typed_task = ReverseTextTask( + prompt="Reverse abc.", + question="Reverse abc.", + answer="cba", + ) + + assert tasks_from_result_typed([typed_task], ReverseTextTask) == [typed_task] + + +def test_task_system_prompt_accepts_config_mapping(): + prompt_path = {"messages": [{"role": "system", "content": "Use short answers."}]} + + task = Task(prompt="hello", system_prompt=prompt_path) + + assert task.system_prompt == [{"role": "system", "content": "Use short answers."}] def test_discover_sibling_dir_returns_empty_existing_dir(tmp_path, monkeypatch) -> None: @@ -122,17 +153,14 @@ def test_taskset_skips_empty_skills_dir(tmp_path, monkeypatch) -> None: def test_v1_env_eval_inputs_can_shuffle_taskset_dataset() -> None: class DemoTaskset(Taskset): + task_type = ReverseTextTask + def load_tasks(self, split: str = "train"): return [{"question": f"Reverse {i}.", "answer": str(i)} for i in range(6)] env = Env(taskset=DemoTaskset()) - inputs = env._get_eval_inputs( - num_examples=3, - rollouts_per_example=2, - shuffle=True, - shuffle_seed=7, - ) + inputs = eval_inputs(env, num_examples=3, rollouts_per_example=2, seed=7) expected = ( env.get_eval_dataset().shuffle(seed=7).select(range(3)).repeat(2).to_list() ) diff --git a/tests/test_v1_textarena_taskset.py b/tests/test_v1_textarena_taskset.py deleted file mode 100644 index 4f1468c5c1..0000000000 --- a/tests/test_v1_textarena_taskset.py +++ /dev/null @@ -1,288 +0,0 @@ -import sys - -import pytest - -import verifiers as vf -from tasksets import textarena - - -class FakeNltk: - def __init__(self): - self.downloads: list[tuple[str, bool]] = [] - - def download(self, package: str, quiet: bool = False): - self.downloads.append((package, quiet)) - - -class FakeTextArenaState: - def __init__(self): - self.game_state: dict[str, str] = {} - self.done = False - self.game_info: dict[int, dict[str, str]] = {} - - -class FakeTextArenaEnv: - dictionary = {"words": {"apple", "berry", "cider"}} - word_list = ["apple", "berry", "cider"] - - def __init__(self, env_id: str = "Wordle-v0"): - self.env_id = env_id - self.guesses: list[str] = [] - self.reset_calls = 0 - self.state = FakeTextArenaState() - - def reset(self, num_players: int): - assert num_players == 1 - self.reset_calls += 1 - self.guesses = [] - self.state = FakeTextArenaState() - - def get_observation(self): - if not self.guesses: - return 0, "Guess the word. [GAME] Use [word]." - return 0, "Board [GAME] Feedback:\nmiss\nY----\ntry again" - - def step(self, guess: str): - self.guesses.append(guess) - secret = self.state.game_state.get("secret_word") - if guess == f"[{secret}]": - self.state.done = True - self.state.game_info = {0: {"reason": "Solved."}} - - -class FakeTextArenaModule: - Env = FakeTextArenaEnv - State = FakeTextArenaState - - def __init__(self): - self.envs: list[FakeTextArenaEnv] = [] - - def make(self, env_id: str): - env = FakeTextArenaEnv(env_id=env_id) - self.envs.append(env) - return env - - -@pytest.fixture -def fake_textarena(monkeypatch): - fake_nltk = FakeNltk() - fake_ta = FakeTextArenaModule() - monkeypatch.setitem(sys.modules, textarena.__name__, textarena) - monkeypatch.setattr(textarena, "nltk", fake_nltk) - monkeypatch.setattr(textarena, "ta", fake_ta) - return fake_nltk, fake_ta - - -def test_textarena_taskset_imports_from_package(): - assert textarena.TextArenaTaskset - assert textarena.TextArenaTasksetConfig - - -def test_textarena_taskset_is_generic_over_config_type(fake_textarena): - class CustomTextArenaConfig(textarena.TextArenaTasksetConfig): - game: str = "FakeWordle-v0" - answer_state_key: str = "secret_word" - - class CustomTextArenaTaskset(textarena.TextArenaTaskset[CustomTextArenaConfig]): - pass - - taskset = CustomTextArenaTaskset(config=CustomTextArenaConfig()) - - assert isinstance(taskset.config, CustomTextArenaConfig) - - -def test_textarena_taskset_builds_train_and_eval_splits(fake_textarena): - fake_nltk, _ = fake_textarena - taskset = textarena.TextArenaTaskset( - config=textarena.TextArenaTasksetConfig( - game="FakeWordle-v0", - answer_state_key="secret_word", - num_train_examples=2, - num_eval_examples=2, - seed=1, - ) - ) - - tasks = list(taskset.get_dataset()) - eval_rows = list(taskset.get_eval_dataset()) - - assert taskset.config.system_prompt is None - assert [task["example_id"] for task in tasks] == [0, 1] - assert [task["example_id"] for task in eval_rows] == [0, 1] - assert all(task["answer"] in FakeTextArenaEnv.word_list for task in tasks) - assert all(task["answer"] in FakeTextArenaEnv.word_list for task in eval_rows) - assert tasks[0]["prompt"] == [ - { - "role": "user", - "content": "Guess the word. [GAME] Use [word].", - } - ] - assert fake_nltk.downloads[:2] == [ - ("words", True), - ("averaged_perceptron_tagger_eng", True), - ] - - -def test_textarena_taskset_flattens_dict_word_list(fake_textarena, monkeypatch): - monkeypatch.setattr( - FakeTextArenaEnv, - "word_list", - {"common": ["apple", "berry"], "rare": "cider"}, - ) - - taskset = textarena.TextArenaTaskset( - config=textarena.TextArenaTasksetConfig( - game="FakeWordle-v0", - answer_state_key="secret_word", - num_train_examples=3, - num_eval_examples=0, - seed=1, - ) - ) - - word_list = ["apple", "berry", "cider"] - assert all(row["answer"] in word_list for row in taskset.get_dataset()) - - -def test_textarena_taskset_loads_user(fake_textarena): - taskset = textarena.TextArenaTaskset( - config=textarena.TextArenaTasksetConfig( - game="FakeWordle-v0", - answer_state_key="secret_word", - num_train_examples=1, - num_eval_examples=0, - ) - ) - - assert isinstance(taskset.user, textarena.TextArenaUser) - - -@pytest.mark.asyncio -async def test_textarena_user_steps_env_and_stops_when_game_finishes(fake_textarena): - _, fake_ta = fake_textarena - taskset = textarena.TextArenaTaskset( - config=textarena.TextArenaTasksetConfig( - game="FakeWordle-v0", - answer_state_key="secret_word", - num_train_examples=1, - num_eval_examples=0, - ) - ) - task = taskset.to_task( - vf.Task( - { - "example_id": 0, - "prompt": [], - "answer": "apple", - "textarena": { - "game": "FakeWordle-v0", - "answer_state_key": "secret_word", - }, - } - ) - ) - state = vf.State.for_task(task) - state["completion"] = [ - vf.AssistantMessage(content="I will guess [apple].") - ] - - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - state = await env.harness.setup_state(task, state) - messages = await env.harness.runtime.user_messages(task, state) - ta_env = fake_ta.envs[-1] - - assert ta_env.guesses == ["[apple]"] - assert ta_env.state.game_state["secret_word"] == "apple" - assert messages == [{"role": "user", "content": "Solved."}] - assert state["done"] is True - assert state["stop_condition"] == "textarena_done" - - -@pytest.mark.asyncio -async def test_textarena_user_accepts_structured_assistant_content(fake_textarena): - _, fake_ta = fake_textarena - taskset = textarena.TextArenaTaskset( - config=textarena.TextArenaTasksetConfig( - game="FakeWordle-v0", - answer_state_key="secret_word", - num_train_examples=1, - num_eval_examples=0, - ) - ) - task = taskset.to_task( - vf.Task( - { - "example_id": 0, - "prompt": [], - "answer": "apple", - "textarena": { - "game": "FakeWordle-v0", - "answer_state_key": "secret_word", - }, - } - ) - ) - state = vf.State.for_task(task) - state["completion"] = [ - vf.AssistantMessage( - content=[ - vf.TextContentPart(text="I will guess "), - vf.TextContentPart(text="[apple]."), - ] - ) - ] - - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - state = await env.harness.setup_state(task, state) - messages = await env.harness.runtime.user_messages(task, state) - ta_env = fake_ta.envs[-1] - - assert ta_env.guesses == ["[apple]"] - assert messages == [{"role": "user", "content": "Solved."}] - assert state["stop_condition"] == "textarena_done" - - -@pytest.mark.asyncio -async def test_textarena_user_returns_wordle_feedback_for_unfinished_game( - fake_textarena, -): - _, fake_ta = fake_textarena - taskset = textarena.TextArenaTaskset( - config=textarena.TextArenaTasksetConfig( - game="FakeWordle-v0", - answer_state_key="secret_word", - num_train_examples=1, - num_eval_examples=0, - ) - ) - task = taskset.to_task( - vf.Task( - { - "example_id": 0, - "prompt": [], - "answer": "apple", - "textarena": { - "game": "FakeWordle-v0", - "answer_state_key": "secret_word", - }, - } - ) - ) - state = vf.State.for_task(task) - state["completion"] = [ - vf.AssistantMessage(content="I will guess [berry].") - ] - - env = vf.Env(taskset=taskset, harness=vf.Harness(config=vf.HarnessConfig())) - state = await env.harness.setup_state(task, state) - messages = await env.harness.runtime.user_messages(task, state) - ta_env = fake_ta.envs[-1] - - assert ta_env.guesses == ["[berry]"] - assert messages == [ - { - "role": "user", - "content": "Board [GAME] Feedback:\nmiss\nY----\ntry again", - } - ] - assert state.get("done") is None diff --git a/tests/test_wiki_search_v1.py b/tests/test_wiki_search_v1.py index 9c728e967c..81c2760b0e 100644 --- a/tests/test_wiki_search_v1.py +++ b/tests/test_wiki_search_v1.py @@ -1,9 +1,10 @@ -import importlib.util +import importlib import sys from pathlib import Path from types import ModuleType import pytest +from verifiers.v1.loaders import load_environment_from_components class StubEmbeddingFunction: @@ -39,6 +40,10 @@ def install_wiki_stubs(monkeypatch: pytest.MonkeyPatch) -> None: stubs = { "chromadb": stub_module("chromadb", PersistentClient=StubPersistentClient), "chromadb.api": stub_module("chromadb.api"), + "chromadb.api.models": stub_module("chromadb.api.models"), + "chromadb.api.models.Collection": stub_module( + "chromadb.api.models.Collection", Collection=object + ), "chromadb.api.types": stub_module( "chromadb.api.types", Embeddable=object, @@ -65,45 +70,69 @@ def install_wiki_stubs(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setitem(sys.modules, name, module) -def load_wiki_module(name: str, monkeypatch: pytest.MonkeyPatch) -> ModuleType: +def load_wiki_v1(monkeypatch: pytest.MonkeyPatch) -> tuple[ModuleType, ModuleType]: install_wiki_stubs(monkeypatch) - module_path = ( - Path(__file__).parents[1] / "environments" / "wiki_search" / f"{name}.py" + env_dir = Path(__file__).parents[1] / "environments" / "wiki_search_v1" + monkeypatch.syspath_prepend(str(env_dir)) + for name in ( + "wiki_search_v1", + "wiki_search_v1.taskset", + "wiki_search_v1.servers", + "wiki_search_v1.servers.wiki", + "wiki_search_v1.servers.wiki.config", + "wiki_search_v1.servers.wiki.toolset", + ): + sys.modules.pop(name, None) + return ( + importlib.import_module("wiki_search_v1"), + importlib.import_module("wiki_search_v1.taskset"), ) - spec = importlib.util.spec_from_file_location(name, module_path) - assert spec is not None and spec.loader is not None - module = importlib.util.module_from_spec(spec) - monkeypatch.setitem(sys.modules, name, module) - spec.loader.exec_module(module) - return module + + +def load_wiki_v0(monkeypatch: pytest.MonkeyPatch) -> ModuleType: + install_wiki_stubs(monkeypatch) + env_dir = Path(__file__).parents[1] / "environments" / "wiki_search" + monkeypatch.syspath_prepend(str(env_dir)) + sys.modules.pop("wiki_search", None) + return importlib.import_module("wiki_search") def test_wiki_search_v1_default_and_explicit_toolsets( monkeypatch: pytest.MonkeyPatch, ) -> None: - module = load_wiki_module("wiki_search_v1", monkeypatch) - wrapper = load_wiki_module("wiki_search", monkeypatch) - - env = wrapper.load_environment( - v1=True, - corpus_dataset="test/corpus", - corpus_split="validation", - chroma_db_dir="/tmp/wiki", - embed_model="test-embed", + package, module = load_wiki_v1(monkeypatch) + + env = load_environment_from_components( + package, + { + "config": { + "taskset": { + "toolsets": { + "wiki": { + "corpus_dataset": "test/corpus", + "corpus_split": "validation", + } + } + } + } + }, ) - assert env.taskset.config.corpus_dataset == "test/corpus" - assert env.taskset.config.corpus_split == "validation" - assert list(env.taskset.named_toolsets) == ["wiki"] - assert len(env.taskset.toolsets) == 1 - assert len(env.taskset.rewards) == 1 + assert env.taskset.toolsets["wiki"].corpus_dataset == "test/corpus" + assert env.taskset.toolsets["wiki"].corpus_split == "validation" + assert list(env.taskset.toolsets) == ["wiki"] + assert [signal["name"] for signal in env.taskset.signals] == ["answer_in_response"] monkeypatch.setattr( module, "load_dataset", lambda *args, **kwargs: [{"question": "question?", "answer": "answer"}], ) - rows = list(module.load_tasks(max_turns=3)) + rows = list( + module.WikiSearchTaskset( + module.WikiSearchTasksetConfig(max_turns=3) + ).load_tasks() + ) assert rows[0]["max_turns"] == 3 assert "judge_model" not in rows[0] @@ -111,26 +140,70 @@ def test_wiki_search_v1_default_and_explicit_toolsets( assert "judge_api_key_var" not in rows[0] taskset = module.WikiSearchTaskset( - config=module.WikiSearchTasksetConfig(toolsets={"custom": {"tools": []}}) + config=module.WikiSearchTasksetConfig( + toolsets={"custom": module.WikiToolsetConfig()} + ) + ) + + assert list(taskset.toolsets) == ["wiki", "custom"] + + configured_taskset = module.WikiSearchTaskset( + config={ + "toolsets": { + "custom": { + "source": "wiki_search_v1.servers.wiki.config:WikiToolsetConfig", + "corpus_dataset": "custom/corpus", + } + } + } ) - assert list(taskset.named_toolsets) == ["wiki", "custom"] - assert len(taskset.toolsets) == 2 + assert configured_taskset.toolsets["custom"].corpus_dataset == "custom/corpus" - configured_env = module.load_environment( - config=module.WikiSearchEnvConfig(harness={"max_turns": 7}) + configured_env = load_environment_from_components( + package, {"config": {"harness": {"max_turns": 7}}} ) assert configured_env.harness.config.max_turns == 7 -def test_wiki_search_v1_rejects_legacy_judge_endpoint_kwargs( +@pytest.mark.asyncio +async def test_wiki_search_v1_lexical_tools(monkeypatch: pytest.MonkeyPatch) -> None: + _, module = load_wiki_v1(monkeypatch) + monkeypatch.setattr( + module, + "load_dataset", + lambda *args, **kwargs: [ + { + "id": "earth", + "title": "Earth", + "content": "# Overview\nEarth is the third planet from the Sun.", + }, + { + "id": "mars", + "title": "Mars", + "content": "# Overview\nMars is a cold desert world.", + }, + ], + ) + + wiki = module.load_wiki(module.WikiToolsetConfig()) + results = await module.search_pages("third planet", wiki) + assert results[0] == {"page_id": "earth", "title": "Earth"} + + sections = await module.view_sections("earth", wiki) + assert sections == [{"section_id": "earth:overview", "section_name": "Overview"}] + assert await module.read_section("earth:overview", wiki) == ( + "# Overview\nEarth is the third planet from the Sun." + ) + + +def test_wiki_search_v0_is_v0_only( monkeypatch: pytest.MonkeyPatch, ) -> None: - wrapper = load_wiki_module("wiki_search", monkeypatch) + wrapper = load_wiki_v0(monkeypatch) - with pytest.raises(ValueError, match="state.get_endpoint_config"): + with pytest.raises(TypeError): wrapper.load_environment( v1=True, - judge_base_url="https://judge.example/v1", ) diff --git a/tests/test_wordle_v1_env.py b/tests/test_wordle_v1_env.py index 845153fb68..16b3fc086c 100644 --- a/tests/test_wordle_v1_env.py +++ b/tests/test_wordle_v1_env.py @@ -2,10 +2,10 @@ @pytest.mark.asyncio -async def test_wordle_user_extracts_latest_feedback(monkeypatch): - from environments.wordle_v1 import wordle_v1 +async def test_wordle_textarena_user_returns_observation(monkeypatch): from tasksets import textarena - import verifiers as vf + from tasksets.textarena import TextArenaTask + import verifiers.v1 as vf class FakeTextArenaState: def __init__(self): @@ -34,34 +34,45 @@ def make(self, env_id: str): return FakeTextArenaEnv() monkeypatch.setattr(textarena, "ta", FakeTextArenaModule()) - task = vf.Task( - { - "answer": "apple", - "textarena": {"game": "Wordle-v0", "answer_state_key": "secret_word"}, - } - ).freeze() - state = vf.State.for_task(task) - state["completion"] = [vf.AssistantMessage(content="[berry]")] + session = textarena.TextArenaSession() + task = TextArenaTask( + answer="apple", + textarena={"game": "Wordle-v0", "answer_state_key": "secret_word"}, + ) + state = vf.State(task_id=task.task_id) + state.transcript.append( + vf.Turn( + prompt=task.prompt, + completion=[vf.AssistantMessage(content="[berry]")], + ) + ) - taskset = wordle_v1.WordleTaskset(config=wordle_v1.WordleTasksetConfig()) - env = vf.Env(taskset=taskset) - state = await env.harness.setup_state(task, state) - response = await env.harness.runtime.user_messages(task, state) + response = textarena.textarena_respond( + session, + task.textarena.model_dump(mode="json"), + task.answer, + [message.model_dump(mode="json") for message in state.completion], + ) - assert response == [vf.UserMessage(content="\nmiss\nY----\ntry again")] + assert response["messages"] == [ + {"role": "user", "content": "intro [GAME] Feedback:\nmiss\nY----\ntry again"} + ] def test_wordle_load_environment_coerces_taskset_config(): - from environments.wordle_v1 import wordle_v1 - from tasksets.textarena import TextArenaTasksetConfig - import verifiers as vf - - env = wordle_v1.load_environment( - vf.EnvConfig( - taskset=TextArenaTasksetConfig( - game="Wordle-v0", answer_state_key="secret_word" - ) - ) + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 + from verifiers.v1.loaders import load_environment_from_components + + env = load_environment_from_components( + wordle_v1, + { + "config": { + "taskset": { + "game": "Wordle-v0", + "answer_state_key": "secret_word", + } + } + }, ) assert isinstance(env.taskset.config, wordle_v1.WordleTasksetConfig) @@ -70,17 +81,18 @@ def test_wordle_load_environment_coerces_taskset_config(): def test_wordle_taskset_uses_textarena_loaders(): - from environments.wordle_v1 import wordle_v1 + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 taskset = wordle_v1.WordleTaskset(config=wordle_v1.WordleTasksetConfig()) assert callable(taskset.load_tasks) - assert isinstance(taskset.user, wordle_v1.WordleUser) + assert taskset.user is not None + assert taskset.user.implementation_ref() == "tasksets.textarena:TextArenaUser" def test_wordle_v1_load_taskset_reads_system_prompt_path(tmp_path): - from environments.wordle_v1 import wordle_v1 - import verifiers as vf + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 + import verifiers.v1 as vf prompt = "Optimized Wordle prompt.\n\nPreserve exact text.\n" prompt_path = tmp_path / "system_prompt.txt" @@ -98,8 +110,8 @@ def test_wordle_v1_load_taskset_reads_system_prompt_path(tmp_path): def test_wordle_v1_load_taskset_rejects_empty_system_prompt_path(tmp_path): - from environments.wordle_v1 import wordle_v1 - import verifiers as vf + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 + import verifiers.v1 as vf prompt_path = tmp_path / "system_prompt.txt" prompt_path.write_text("", encoding="utf-8") @@ -114,17 +126,26 @@ def test_wordle_v1_load_taskset_rejects_empty_system_prompt_path(tmp_path): @pytest.mark.asyncio async def test_wordle_v1_rewards_match_wordle_protocol(): - from environments.wordle_v1 import wordle_v1 - import verifiers as vf - - taskset = wordle_v1.WordleTaskset.__new__(wordle_v1.WordleTaskset) - task = vf.Task({"answer": "apple"}).freeze() - state = vf.State.for_task(task) - state["completion"] = [ - vf.AssistantMessage(content="[berry]"), - vf.UserMessage(content="miss\nGY---\ntry again"), - vf.AssistantMessage(content="[apple]"), - ] + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 + from tasksets.textarena import TextArenaTask + import verifiers.v1 as vf + + taskset = wordle_v1.WordleTaskset(config=wordle_v1.WordleTasksetConfig()) + task = TextArenaTask( + answer="apple", + textarena={"game": "Wordle-v0", "answer_state_key": "secret_word"}, + ) + state = vf.State(task_id=task.task_id) + state.transcript.append( + vf.Turn( + prompt=task.prompt, + completion=[ + vf.AssistantMessage(content="[berry]"), + vf.UserMessage(content="miss\nGY---\ntry again"), + vf.AssistantMessage(content="[apple]"), + ], + ) + ) assert await taskset.correct_answer(task, state) == 1.0 assert await taskset.length_bonus(task, state) == 0.5 @@ -134,29 +155,42 @@ async def test_wordle_v1_rewards_match_wordle_protocol(): @pytest.mark.asyncio async def test_wordle_v1_partial_answer_scans_past_non_guess_messages(): - from environments.wordle_v1 import wordle_v1 - import verifiers as vf - - taskset = wordle_v1.WordleTaskset.__new__(wordle_v1.WordleTaskset) - task = vf.Task({"answer": "apple"}).freeze() - state = vf.State.for_task(task) - state["completion"] = [ - vf.UserMessage(content="miss\nGGGGG\ntry again"), - vf.AssistantMessage(content="[apple]"), - vf.AssistantMessage(content="I already found it."), - ] + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 + from tasksets.textarena import TextArenaTask + import verifiers.v1 as vf + + taskset = wordle_v1.WordleTaskset(config=wordle_v1.WordleTasksetConfig()) + task = TextArenaTask( + answer="apple", + textarena={"game": "Wordle-v0", "answer_state_key": "secret_word"}, + ) + state = vf.State(task_id=task.task_id) + state.transcript.append( + vf.Turn( + prompt=task.prompt, + completion=[ + vf.UserMessage(content="miss\nGGGGG\ntry again"), + vf.AssistantMessage(content="[apple]"), + vf.AssistantMessage(content="I already found it."), + ], + ) + ) assert await taskset.partial_answer(task, state) == 0.0 @pytest.mark.asyncio async def test_wordle_v1_rewards_treat_missing_completion_as_empty(): - from environments.wordle_v1 import wordle_v1 - import verifiers as vf + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 + from tasksets.textarena import TextArenaTask + import verifiers.v1 as vf - taskset = wordle_v1.WordleTaskset.__new__(wordle_v1.WordleTaskset) - task = vf.Task({"answer": "apple"}).freeze() - state = vf.State.for_task(task) + taskset = wordle_v1.WordleTaskset(config=wordle_v1.WordleTasksetConfig()) + task = TextArenaTask( + answer="apple", + textarena={"game": "Wordle-v0", "answer_state_key": "secret_word"}, + ) + state = vf.State(task_id=task.task_id) assert await taskset.correct_answer(task, state) == 0.0 assert await taskset.length_bonus(task, state) == 0.0 @@ -165,7 +199,7 @@ async def test_wordle_v1_rewards_treat_missing_completion_as_empty(): def test_wordle_taskset_declares_rewards_as_methods(): - from environments.wordle_v1 import wordle_v1 + from environments.wordle_v1.wordle_v1 import taskset as wordle_v1 for name in ("correct_answer", "partial_answer", "length_bonus", "format_reward"): assert getattr(getattr(wordle_v1.WordleTaskset, name), "reward") is True diff --git a/uv.lock b/uv.lock index 0d5e15155d..750648d713 100644 --- a/uv.lock +++ b/uv.lock @@ -46,7 +46,7 @@ conflicts = [[ ]] [options] -exclude-newer = "2026-05-31T06:11:52.015918Z" +exclude-newer = "2026-05-31T07:00:48.966396Z" exclude-newer-span = "P7D" [options.exclude-newer-package] @@ -7687,26 +7687,27 @@ wheels = [ [[package]] name = "ty" -version = "0.0.21" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/ee/20/2ba8fd9493c89c41dfe9dbb73bc70a28b28028463bc0d2897ba8be36230a/ty-0.0.21.tar.gz", hash = "sha256:a4c2ba5d67d64df8fcdefd8b280ac1149d24a73dbda82fa953a0dff9d21400ed", size = 5297967, upload-time = "2026-03-06T01:57:13.809Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/36/70/edf38bb37517531681d1c37f5df64744e5ad02673c02eb48447eae4bea08/ty-0.0.21-py3-none-linux_armv6l.whl", hash = "sha256:7bdf2f572378de78e1f388d24691c89db51b7caf07cf90f2bfcc1d6b18b70a76", size = 10299222, upload-time = "2026-03-06T01:57:16.64Z" }, - { url = "https://files.pythonhosted.org/packages/72/62/0047b0bd19afeefbc7286f20a5f78a2aa39f92b4d89853f0d7185ab89edc/ty-0.0.21-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:7e9613994610431ab8625025bd2880dbcb77c5c9fabdd21134cda12d840a529d", size = 10130513, upload-time = "2026-03-06T01:57:29.93Z" }, - { url = "https://files.pythonhosted.org/packages/a2/20/0b93a9e91aaed23155780258cdfdb4726ef68b6985378ac069bc427291a0/ty-0.0.21-py3-none-macosx_11_0_arm64.whl", hash = "sha256:56d3b198b64dd0a19b2b66e257deaed2ecea568e722ae5352f3c6fb62027f89d", size = 9605425, upload-time = "2026-03-06T01:57:27.115Z" }, - { url = "https://files.pythonhosted.org/packages/ea/fd/9945e2fa2996a1287b1e1d7ce050e97e1f420233b271e770934bfa0880a0/ty-0.0.21-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:d23d2c34f7a77d974bb08f0860ef700addc8a683d81a0319f71c08f87506cfd0", size = 10108298, upload-time = "2026-03-06T01:57:35.429Z" }, - { url = "https://files.pythonhosted.org/packages/52/e7/4ec52fcb15f3200826c9f048472c062549a05b0d1ef0b51f32d527b513c4/ty-0.0.21-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:56b01fd2519637a4ca88344f61c96225f540c98ff18bca321d4eaa7bb0f7aa2f", size = 10121556, upload-time = "2026-03-06T01:57:03.242Z" }, - { url = "https://files.pythonhosted.org/packages/ee/c0/ad457be2a8abea0f25549598bd098554540ced66229488daa0d558dad3c8/ty-0.0.21-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:e9de7e11c63c6afc40f3e9ba716374add171aee7fabc70b5146a510705c6d41b", size = 10603264, upload-time = "2026-03-06T01:56:52.134Z" }, - { url = "https://files.pythonhosted.org/packages/f8/5b/2ecc7a2175243a4bcb72f5298ae41feabbb93b764bb0dc45722f3752c2c2/ty-0.0.21-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:62f7f5b235c4f7876db305c36997aea07b7af29b1a068f373d0e2547e25f32ff", size = 11196428, upload-time = "2026-03-06T01:57:32.94Z" }, - { url = "https://files.pythonhosted.org/packages/37/f5/aff507d6a901f328ef96a298032b0c11aaaf950a146ed7dd3b5bf2cd3acf/ty-0.0.21-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:ee8399f7c453a425291e6688efe430cfae7ab0ac4ffd50eba9f872bf878b54f6", size = 10866355, upload-time = "2026-03-06T01:56:57.831Z" }, - { url = "https://files.pythonhosted.org/packages/be/30/822bbcb92d55b65989aa7ed06d9585f28ade9c9447369194ed4b0fb3b5b9/ty-0.0.21-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:210e7568c9f886c4d01308d751949ee714ad7ad9d7d928d2ba90d329dd880367", size = 10738177, upload-time = "2026-03-06T01:57:11.256Z" }, - { url = "https://files.pythonhosted.org/packages/57/cc/46e7991b6469e93ac2c7e533a028983e402485580150ac864c56352a3a82/ty-0.0.21-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:53508e345b11569f78b21ba8e2b4e61df38a9754947fb3cd9f2ef574367338fb", size = 10079158, upload-time = "2026-03-06T01:57:00.516Z" }, - { url = "https://files.pythonhosted.org/packages/15/c2/0bbdadfbd008240f8f1a87dc877433cb3884436097926107ccf06e618199/ty-0.0.21-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:553e43571f4a35604c36cfd07d8b61a5eb7a714e3c67f8c4ff2cf674fefbaef9", size = 10150535, upload-time = "2026-03-06T01:57:08.815Z" }, - { url = "https://files.pythonhosted.org/packages/c5/b5/2dbdb7b57b5362200ef0a39738ebd31331726328336def0143ac097ee59d/ty-0.0.21-py3-none-musllinux_1_2_i686.whl", hash = "sha256:666f6822e3b9200abfa7e95eb0ddd576460adb8d66b550c0ad2c70abc84a2048", size = 10319803, upload-time = "2026-03-06T01:57:19.106Z" }, - { url = "https://files.pythonhosted.org/packages/72/84/70e52c0b7abc7c2086f9876ef454a73b161d3125315536d8d7e911c94ca4/ty-0.0.21-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:a0854d008347ce4a5fb351af132f660a390ab2a1163444d075251d43e6f74b9b", size = 10826239, upload-time = "2026-03-06T01:57:21.727Z" }, - { url = "https://files.pythonhosted.org/packages/a1/8a/1f72480fd013bbc6cd1929002abbbcde9a0b08ead6a15154de9d7f7fa37e/ty-0.0.21-py3-none-win32.whl", hash = "sha256:bef3ab4c7b966bcc276a8ac6c11b63ba222d21355b48d471ea782c4104eee4e0", size = 9693196, upload-time = "2026-03-06T01:57:24.126Z" }, - { url = "https://files.pythonhosted.org/packages/8d/f8/1104808b875c26c640e536945753a78562d606bef4e241d9dbf3d92477f6/ty-0.0.21-py3-none-win_amd64.whl", hash = "sha256:a709d576e5bea84b745d43058d8b9cd4f27f74a0b24acb4b0cbb7d3d41e0d050", size = 10668660, upload-time = "2026-03-06T01:56:55.06Z" }, - { url = "https://files.pythonhosted.org/packages/1b/b8/25e0adc404bbf986977657b25318991f93097b49f8aea640d93c0b0db68e/ty-0.0.21-py3-none-win_arm64.whl", hash = "sha256:f72047996598ac20553fb7e21ba5741e3c82dee4e9eadf10d954551a5fe09391", size = 10104161, upload-time = "2026-03-06T01:57:06.072Z" }, +version = "0.0.40" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/5a/f8/a754c96967b71de8723f88be17df8738216bd382ffed229cd500b7a24d13/ty-0.0.40.tar.gz", hash = "sha256:883b53dd98f6e5b33ab1c8e1a3cd94b0f29c762ef22cdf1e86aaffb4fd711c67", size = 5726484, upload-time = "2026-05-27T17:55:43.615Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2c/42/d029a72165ad39f95228b67355927fbd35c821dc8e3e475d49f47c2eeb1e/ty-0.0.40-py3-none-linux_armv6l.whl", hash = "sha256:9defb4742450e569a6a09de286a04008d6c2e815112da4362c88b6eaa2f52a36", size = 11406372, upload-time = "2026-05-27T17:55:49.633Z" }, + { url = "https://files.pythonhosted.org/packages/23/99/7f8ea09b7e49afbf795cb3341a3217f30f228db7e62a2268ed8cbbf813d6/ty-0.0.40-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:868258a3330db88b683fcafe2c4e936d6226a6312799bf15b585d93557b2d38c", size = 11159782, upload-time = "2026-05-27T17:55:47.405Z" }, + { url = "https://files.pythonhosted.org/packages/04/d8/1ea745ee97a98b26ae9564d19a430a76a35297cd450e84dcaad22e1f7ee8/ty-0.0.40-py3-none-macosx_11_0_arm64.whl", hash = "sha256:589c81060cf1e7a9ffa2f45bfa35ffd9b9fbd214104e3f13959f113627efcd91", size = 10594139, upload-time = "2026-05-27T17:55:37.206Z" }, + { url = "https://files.pythonhosted.org/packages/39/1a/fbef21273c6617ff4715b4827ee1c0b6550aa7d1df4b8c43b325545c1cf4/ty-0.0.40-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7b06108990cb338d941c315ae6e9ba2fff8f518bc15d3f33e5619ff6a6c9beab", size = 11114156, upload-time = "2026-05-27T17:55:56.11Z" }, + { url = "https://files.pythonhosted.org/packages/3c/f9/389fc4976d7ec016a7473cf1274bf9c4f491bb54c66649bd022bff9f2b6a/ty-0.0.40-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:3913ef37336bec4f96bd2512f8c3a543ca34c259b7170f7eb5adf75b3ed7f04c", size = 11189050, upload-time = "2026-05-27T17:55:54.099Z" }, + { url = "https://files.pythonhosted.org/packages/fa/a9/4ecabbf4bdda7df0d99d8d3892c6edac0efc8c4cae756a5109178a3d0e86/ty-0.0.40-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:8fd1486bd5fe48779a8aa857137f3642a0a9161f5cf57d4380f4a0ecea01c8f3", size = 11664266, upload-time = "2026-05-27T17:55:28.17Z" }, + { url = "https://files.pythonhosted.org/packages/45/02/0aa78730116507c265afb1d6d5961c583b49d4c2e368c4a49fd81bcae6dc/ty-0.0.40-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:1668364d5254a734329917ee66c2c5fdd5665389d41043f6fce0f22ddb32b749", size = 12187743, upload-time = "2026-05-27T17:56:04.337Z" }, + { url = "https://files.pythonhosted.org/packages/e6/68/ccabf2d173523598271a385c1d3f864dbda23e5ebdc67f5969b9e830ea05/ty-0.0.40-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:43f77a73edb91e5dfa2ab9af7c4cac64614f8cc121f38a8875f22e830d3aba6a", size = 11862999, upload-time = "2026-05-27T17:55:58.087Z" }, + { url = "https://files.pythonhosted.org/packages/03/8d/6d7ec22771bb23d534797cdb446eb644bccfe7a62b729bb99e7235a02fc3/ty-0.0.40-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1274ce0212ecbfed01bda7c3659c46e8bd0068e32d00c46c790466a95274c3df", size = 11743896, upload-time = "2026-05-27T17:56:00.017Z" }, + { url = "https://files.pythonhosted.org/packages/cd/a4/f9fa076b010c91cb249b1fcc3476569b7b8462cb4b688da2d04c23a0622f/ty-0.0.40-py3-none-manylinux_2_31_riscv64.whl", hash = "sha256:5ee1261dbc363e5cc1a0c5bb0c8612c192bfe53491214df8bc85a540835685f9", size = 11883581, upload-time = "2026-05-27T17:56:02.319Z" }, + { url = "https://files.pythonhosted.org/packages/fd/0f/5b776a2328c756d574dd4d6afbd30fc24e1ab4b76935c7c3c23f27ebbcb9/ty-0.0.40-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:6220e2cd5cdc4683dd87fb150d195bbd9f1a021395e04cb08bd3c66ea6da6ef8", size = 11093946, upload-time = "2026-05-27T17:55:33.284Z" }, + { url = "https://files.pythonhosted.org/packages/64/c4/eb23154bae83ad7c2935e9e5916660fb3e31598a92ee232aebd79410480c/ty-0.0.40-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:46b9ed69d01d98ef046afac9983c68336f572605ea2a27b90fbe6f80bfc8d6b7", size = 11210737, upload-time = "2026-05-27T17:55:45.523Z" }, + { url = "https://files.pythonhosted.org/packages/ff/19/1fb2529703f708cacfd13a89f98613cae2907dfa941b26976467e6119803/ty-0.0.40-py3-none-musllinux_1_2_i686.whl", hash = "sha256:ddbca9fab4406260f141674ab5efcfe7b02bd468e6985e4cdde0a21626e69ffe", size = 11332563, upload-time = "2026-05-27T17:55:41.674Z" }, + { url = "https://files.pythonhosted.org/packages/87/69/b3f5a8ef26c31204e0391147b3adcdb0674eda3e7d99868478ef168a41c6/ty-0.0.40-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:b1fcc082a749e6dc11b68fe9aab0420238bbf2a2374c2c7aa3c22e8c1618b136", size = 11843216, upload-time = "2026-05-27T17:55:35.367Z" }, + { url = "https://files.pythonhosted.org/packages/ac/e8/20193069d32787f3e1a6ec8940aaa3759d3de8f48f9281bcc0c5cb0939da/ty-0.0.40-py3-none-win32.whl", hash = "sha256:75feb115b3587824c5bdf8f8305e9547b0d1e398e3077b0addc7a1988ea9bb50", size = 10670731, upload-time = "2026-05-27T17:55:31.316Z" }, + { url = "https://files.pythonhosted.org/packages/a3/f9/8b2aa4da61db81322d4a2f9db227afeb48110ca15ae31d380f64c64ceb63/ty-0.0.40-py3-none-win_amd64.whl", hash = "sha256:b0f905edaad788bd61f779a85801b60a267a25ed57fca05aaddd168d9d8896be", size = 11766211, upload-time = "2026-05-27T17:55:51.898Z" }, + { url = "https://files.pythonhosted.org/packages/04/87/369056ed46f1b235130ec0595393262f9cd2061ca3dab276d490980f9343/ty-0.0.40-py3-none-win_arm64.whl", hash = "sha256:07da2b09d9130e2c9a257d2a29beb53105835b0256ee5fdb288fe1aab83fee47", size = 11117369, upload-time = "2026-05-27T17:55:39.329Z" }, ] [[package]] @@ -8032,7 +8033,7 @@ dev = [ { name = "ruff" }, { name = "stagehand", specifier = ">=3.0.0" }, { name = "tasksets", extras = ["openenv", "openreward", "ta"], editable = "packages/tasksets" }, - { name = "ty", specifier = ">=0.0.1a29,<0.0.22" }, + { name = "ty", specifier = ">=0.0.40,<0.0.41" }, ] policy = [{ name = "semgrep", specifier = ">=1.150.0" }] diff --git a/verifiers/__init__.py b/verifiers/__init__.py index 35fbec78f0..e369155a9a 100644 --- a/verifiers/__init__.py +++ b/verifiers/__init__.py @@ -23,7 +23,7 @@ teardown, update, ) -from .types import DatasetBuilder, EndpointConfig, Endpoints, State # noqa # isort: skip +from .types import DatasetBuilder, EndpointConfig, Endpoints, State, ToolLike # noqa # isort: skip from .parsers.parser import Parser # noqa # isort: skip from .rubrics.rubric import Rubric # noqa # isort: skip @@ -51,22 +51,8 @@ setup_logging(os.getenv("VF_LOG_LEVEL")) __all__ = [ - "ArtifactConfig", - "Artifacts", - "ArtifactsConfig", "DatasetBuilder", "State", - "BindingsConfig", - "CallableConfig", - "Config", - "ConfigData", - "Handler", - "JsonData", - "Objects", - "ObjectsConfig", - "ProgramConfig", - "ProgramValue", - "PromptInput", "Parser", "ThinkParser", "MaybeThinkParser", @@ -83,34 +69,10 @@ "MCPEnv", "BrowserEnv", "OpenEnvEnv", - "Env", - "EnvConfig", - "Endpoint", "EndpointConfig", "Endpoints", - "Task", "TaskSplit", - "Tasks", - "Taskset", - "TasksetConfig", - "Harness", - "HarnessConfig", - "MCPTool", - "MCPToolConfig", - "ModelConfig", - "SandboxConfig", - "SystemPrompt", - "SystemPromptConfig", - "SystemPromptStrategy", - "Toolset", "ToolLike", - "ToolsetConfig", - "Toolsets", - "TrajectoryVisibility", - "User", - "UserConfig", - "VisibilityConfig", - "SignalConfig", "Environment", "MultiTurnEnv", "SingleTurnEnv", @@ -132,10 +94,7 @@ "log_level", "quiet_verifiers", "load_environment", - "load_harness", - "load_taskset", "print_prompt_completions_sample", - "get_messages", "cleanup", "metric", "reward", @@ -144,13 +103,6 @@ "stop", "teardown", "update", - "add_metric", - "add_reward", - "add_advantage", - "build_signals", - "collect_signals", - "score_group", - "score_rollout", "ensure_keys", "MissingKeyError", "get_model", @@ -186,8 +138,6 @@ "EnvGroup": "verifiers.envs.env_group:EnvGroup", "JudgeRubric": "verifiers.rubrics.judge_rubric:JudgeRubric", "load_environment": "verifiers.utils.env_utils:load_environment", - "load_harness": "verifiers.utils.env_utils:load_harness", - "load_taskset": "verifiers.utils.env_utils:load_taskset", "get_model": "verifiers_rl.rl.trainer.utils:get_model", "get_model_and_tokenizer": "verifiers_rl.rl.trainer.utils:get_model_and_tokenizer", "RLConfig": "verifiers_rl.rl.trainer:RLConfig", @@ -207,55 +157,6 @@ "TextArenaEnv": "verifiers.envs.integrations.textarena_env:TextArenaEnv", "BrowserEnv": "verifiers.envs.integrations.browser_env:BrowserEnv", "OpenEnvEnv": "verifiers.envs.integrations.openenv_env:OpenEnvEnv", - "Config": "verifiers.v1:Config", - "CallableConfig": "verifiers.v1:CallableConfig", - "BindingsConfig": "verifiers.v1:BindingsConfig", - "ArtifactConfig": "verifiers.v1:ArtifactConfig", - "Artifacts": "verifiers.v1:Artifacts", - "ArtifactsConfig": "verifiers.v1:ArtifactsConfig", - "Env": "verifiers.v1:Env", - "EnvConfig": "verifiers.v1:EnvConfig", - "Endpoint": "verifiers.v1:Endpoint", - "EndpointConfig": "verifiers.v1:EndpointConfig", - "ConfigData": "verifiers.v1:ConfigData", - "Handler": "verifiers.v1:Handler", - "JsonData": "verifiers.v1:JsonData", - "Objects": "verifiers.v1:Objects", - "ObjectsConfig": "verifiers.v1:ObjectsConfig", - "Task": "verifiers.v1:Task", - "TaskSplit": "verifiers.v1:TaskSplit", - "Tasks": "verifiers.v1:Tasks", - "Taskset": "verifiers.v1:Taskset", - "TasksetConfig": "verifiers.v1:TasksetConfig", - "Harness": "verifiers.v1:Harness", - "HarnessConfig": "verifiers.v1:HarnessConfig", - "ProgramConfig": "verifiers.v1:ProgramConfig", - "ProgramValue": "verifiers.v1:ProgramValue", - "PromptInput": "verifiers.v1:PromptInput", - "MCPTool": "verifiers.v1:MCPTool", - "MCPToolConfig": "verifiers.v1:MCPToolConfig", - "ModelConfig": "verifiers.v1:ModelConfig", - "SandboxConfig": "verifiers.v1:SandboxConfig", - "SignalConfig": "verifiers.v1:SignalConfig", - "SystemPrompt": "verifiers.v1:SystemPrompt", - "SystemPromptConfig": "verifiers.v1:SystemPromptConfig", - "SystemPromptStrategy": "verifiers.v1:SystemPromptStrategy", - "ToolLike": "verifiers.v1:ToolLike", - "Toolset": "verifiers.v1:Toolset", - "ToolsetConfig": "verifiers.v1:ToolsetConfig", - "Toolsets": "verifiers.v1:Toolsets", - "TrajectoryVisibility": "verifiers.v1:TrajectoryVisibility", - "User": "verifiers.v1:User", - "UserConfig": "verifiers.v1:UserConfig", - "VisibilityConfig": "verifiers.v1:VisibilityConfig", - "get_messages": "verifiers.v1:get_messages", - "add_metric": "verifiers.v1:add_metric", - "add_reward": "verifiers.v1:add_reward", - "add_advantage": "verifiers.v1:add_advantage", - "build_signals": "verifiers.v1:build_signals", - "collect_signals": "verifiers.v1:collect_signals", - "score_group": "verifiers.v1:score_group", - "score_rollout": "verifiers.v1:score_rollout", } @@ -320,58 +221,6 @@ def __getattr__(name: str): from .rubrics.math_rubric import MathRubric # noqa: F401 from .utils.env_utils import ( # noqa: F401 load_environment, - load_harness, - load_taskset, - ) - from .v1 import ( # noqa: F401 - ArtifactConfig, - Artifacts, - ArtifactsConfig, - BindingsConfig, - CallableConfig, - Config, - ConfigData, - Env, - EnvConfig, - Endpoint, - EndpointConfig, - Handler, - Harness, - HarnessConfig, - JsonData, - MCPTool, - MCPToolConfig, - ModelConfig, - Objects, - ObjectsConfig, - ProgramConfig, - ProgramValue, - PromptInput, - SandboxConfig, - SignalConfig, - SystemPrompt, - SystemPromptConfig, - SystemPromptStrategy, - Task, - Tasks, - Taskset, - TasksetConfig, - ToolLike, - Toolset, - ToolsetConfig, - Toolsets, - TrajectoryVisibility, - User, - UserConfig, - VisibilityConfig, - add_advantage, - add_metric, - add_reward, - build_signals, - collect_signals, - get_messages, - score_group, - score_rollout, ) # Optional verifiers-rl exports. Keep type-checking clean when extra is absent. diff --git a/verifiers/clients/nemorl_chat_completions_client.py b/verifiers/clients/nemorl_chat_completions_client.py index 99f02b79b2..c1132356de 100644 --- a/verifiers/clients/nemorl_chat_completions_client.py +++ b/verifiers/clients/nemorl_chat_completions_client.py @@ -67,7 +67,7 @@ async def get_response( tools: list[Tool] | None = None, **kwargs, ) -> Response: - """Annotate prior assistant messages with their trajectory tokens before delegating to the parent client.""" + """Annotate prior assistant messages with their token ids before delegation.""" state = kwargs.get("state") if state is not None: _attach_trajectory_tokens_to_prompt(prompt, state) @@ -96,12 +96,12 @@ async def to_native_prompt( def _attach_trajectory_tokens_to_prompt( prompt: Messages, state: dict[str, Any] ) -> None: - """Attach each past assistant message's token ids from the trajectory.""" - trajectory = state.get("trajectory") or [] - if not trajectory: + """Attach each past assistant message's token ids from recorded turns.""" + turns = state.get("transcript") or state.get("trajectory") or [] + if not turns: return indices = [i for i, m in enumerate(prompt) if isinstance(m, AssistantMessage)] - step_tokens = [step.get("tokens") for step in trajectory] + step_tokens = [step.get("tokens") for step in turns] n = min(len(indices), len(step_tokens)) for i, tokens in zip(indices[-n:], step_tokens[-n:]): if tokens is None: diff --git a/verifiers/clients/openai_chat_completions_token_client.py b/verifiers/clients/openai_chat_completions_token_client.py index 2d8cd701cc..edfcd13216 100644 --- a/verifiers/clients/openai_chat_completions_token_client.py +++ b/verifiers/clients/openai_chat_completions_token_client.py @@ -9,6 +9,7 @@ ChatCompletionMessageFunctionToolCallParam, Function, ) +from pydantic import TypeAdapter from verifiers.clients.openai_chat_completions_client import ( OpenAIChatCompletionsClient, @@ -18,11 +19,13 @@ OpenAITool, handle_openai_overlong_prompt, ) -from verifiers.types import SamplingArgs, State +from verifiers.types import Messages, SamplingArgs, State from verifiers.utils.client_utils import ( post_chat_completion_with_routed_experts_sidecar, ) +_MESSAGES_ADAPTER = TypeAdapter(Messages) + def _has_multimodal_content(messages) -> bool: """Check if any message contains multimodal content (images, audio). @@ -42,6 +45,14 @@ def _has_multimodal_content(messages) -> bool: return False +def _state_turns(state: State) -> list[Any]: + if hasattr(state, "get"): + turns = state.get("transcript") or state.get("trajectory") or [] + else: + turns = getattr(state, "transcript", None) or getattr(state, "trajectory", []) + return list(turns) + + # copy from vllm/entrypoints/openai/protocol.py class TokenizeResponse(BaseModel): count: int @@ -95,13 +106,16 @@ def normalize_sampling_args(sampling_args: SamplingArgs): # N) and token-stitching (TITO) produces broken prompts. Falling back # to message-based inference (MITO) lets vLLM handle expansion # correctly on every turn. + turns = _state_turns(state) has_multimodal = _has_multimodal_content(prompt) or any( - _has_multimodal_content(step["prompt"]) for step in state["trajectory"] + _has_multimodal_content(step["prompt"]) for step in turns ) - if len(state["trajectory"]) == 0 or has_multimodal: + if len(turns) == 0 or has_multimodal: return await super().get_native_response( prompt, model, sampling_args, tools, extra_headers=extra_headers ) + if hasattr(state, "get") and state.get("model") is None: + state = cast(State, {**state, "model": model}) # The bridge tokenize calls inside get_prompt_ids must run under the # same chat-template config as the engine's actual generation, # otherwise the bridge tokens won't line up with what vLLM streamed @@ -114,14 +128,15 @@ def normalize_sampling_args(sampling_args: SamplingArgs): "chat_template_kwargs", {} ) prompt_ids = await self.get_prompt_ids( - state, prompt, tools, chat_template_kwargs=chat_template_kwargs + state, + prompt, + tools, + chat_template_kwargs=chat_template_kwargs, ) if prompt_ids is None: - # Reaching this branch means we have a non-empty trajectory but - # could not stitch — surface it loudly so ops catches regressions. - self.logger.warning( - f"TITO fell back to MITO on turn {len(state['trajectory']) + 1}" - ) + # Reaching this branch means we have prior turns but could not stitch; + # surface it loudly so ops catches regressions. + self.logger.warning(f"TITO fell back to MITO on turn {len(turns) + 1}") return await super().get_native_response( prompt, model, sampling_args, tools, extra_headers=extra_headers ) @@ -148,6 +163,8 @@ async def get_prompt_ids( state: State, prompt_messages: OpenAIChatMessages, oai_tools: list[OpenAITool] | None, + *, + model: str | None = None, chat_template_kwargs: dict | None = None, ) -> list[int] | None: """ @@ -163,7 +180,7 @@ async def get_prompt_ids( def normalize_for_comparison(value: Any) -> Any: if hasattr(value, "model_dump"): - return normalize_for_comparison(value.model_dump()) + return normalize_for_comparison(value.model_dump(exclude_none=True)) if isinstance(value, Mapping): normalized = { str(key): normalize_for_comparison(val) @@ -182,17 +199,19 @@ def normalize_for_comparison(value: Any) -> Any: return value async def find_largest_prefix_match() -> tuple[list[int], bool, int] | None: - """Scan trajectory backwards for the step whose messages form the + """Scan previous turns backwards for the step whose messages form the longest prefix of prompt_messages. Returns (token_ids, is_truncated, prefix_len) or None.""" normalized_prompt_messages = normalize_for_comparison(prompt_messages) best_prefix_len = -1 best_step = None - for step in reversed(state["trajectory"]): + for step in reversed(_state_turns(state)): step_tokens = step["tokens"] if step_tokens is None: continue - step_messages = cast(Any, [*step["prompt"], *step["completion"]]) + step_messages = _MESSAGES_ADAPTER.validate_python( + cast(Any, [*step["prompt"], *step["completion"]]) + ) step_prompt_messages, _ = await self.to_native_prompt(step_messages) normalized_step_messages = normalize_for_comparison( step_prompt_messages @@ -312,18 +331,19 @@ async def find_largest_prefix_match() -> tuple[list[int], bool, int] | None: if chat_template_kwargs else {} ) + tokenize_model = model or state["model"] try: bridge_full_ids = await self.tokenize( messages=[dummy_assistant] + env_messages, tools=oai_tools, - model=state["model"], + model=tokenize_model, extra_kwargs=dict(forwarded_ctk), ) bridge_base_ids = await self.tokenize( messages=[dummy_assistant], tools=oai_tools, - model=state["model"], + model=tokenize_model, extra_kwargs=dict(add_generation_prompt=False, **forwarded_ctk), ) except Exception: diff --git a/verifiers/clients/renderer_client.py b/verifiers/clients/renderer_client.py index 64ca4ec89d..f9d919d2c8 100644 --- a/verifiers/clients/renderer_client.py +++ b/verifiers/clients/renderer_client.py @@ -100,8 +100,14 @@ def _get_value(obj: Any, key: str, default: Any = None) -> Any: return getattr(obj, key, default) +def _state_turns(state: Any) -> list[Any]: + return list( + _get_value(state, "transcript") or _get_value(state, "trajectory") or [] + ) + + def _normalize_for_comparison(value: Any, _key: str | None = None) -> Any: - # tool_call.arguments is serialized as a string on one side (our trajectory + # tool_call.arguments is serialized as a string on one side (our stored turns # uses json.dumps with default separators) and often comes back from # upstream scaffolds re-stringified with JS JSON.stringify (compact, no # spaces). Both encode the same dict; parse and normalize structurally so @@ -300,15 +306,15 @@ async def _get_incremental_prompt_ids( ) -> "tuple[RenderedTokens, int] | None": """Return the bridged prompt and routed-experts replay start. - Returns ``None`` when no prior trajectory step lines up with the new + Returns ``None`` when no prior turn lines up with the new prompt's prefix or the renderer's ``bridge_to_next_turn`` can't extend — both cases fall back to a full re-render in :func:`generate`. """ if not state: return None - trajectory = _get_value(state, "trajectory") - if not trajectory: + turns = _state_turns(state) + if not turns: return None # Each renderer's bridge_to_next_turn (or the generic fallback) decides @@ -318,7 +324,7 @@ async def _get_incremental_prompt_ids( # falls back to a full re-render — matching main's TITO-on-truncation # behavior. normalized_prompt = _normalize_for_comparison(prompt) - for step in reversed(list(trajectory)): + for step in reversed(turns): token_ids = _step_token_ids(step) if token_ids is None: continue @@ -505,6 +511,27 @@ def _get_renderer_or_pool( return self._shared_pools[cache_key] + def get_renderer( + self, + model: str, + *, + sampling_args: SamplingArgs | None = None, + ) -> Renderer | RendererPool: + args = dict(sampling_args or {}) + sampling_params = dict(args.pop("extra_body", None) or {}) + chat_template_kwargs = sampling_params.pop("chat_template_kwargs", None) + renderer_model = ( + self._config.renderer_model_name + if self._config is not None and self._config.renderer_model_name is not None + else model + ) + renderer_config = _resolve_renderer_config( + self._config.renderer_config if self._config is not None else None, + chat_template_kwargs, + renderer_model=renderer_model, + ) + return self._get_renderer_or_pool(model, renderer_config=renderer_config) + # ── Type conversions ──────────────────────────────────────────── async def to_native_prompt( diff --git a/verifiers/envs/environment.py b/verifiers/envs/environment.py index 3795b81908..2d7f84140d 100644 --- a/verifiers/envs/environment.py +++ b/verifiers/envs/environment.py @@ -26,6 +26,7 @@ final, ) +from pydantic import TypeAdapter from verifiers.clients import Client, resolve_client from verifiers.decorators import discover_decorated from verifiers.serve import ZMQEnvClient @@ -61,8 +62,10 @@ SamplingArgs, StartCallback, State, + TextMessage, TokenUsage, Tool, + UserMessage, flatten_task_input, ) from verifiers.utils.async_utils import ( @@ -72,7 +75,6 @@ with_sem, ) from verifiers.utils.error_utils import ErrorChain -from verifiers.utils.message_utils import normalize_messages from verifiers.utils.save_utils import ( GenerateOutputsBuilder, load_outputs, @@ -87,6 +89,7 @@ ) from verifiers.utils.usage_utils import StateUsageTracker +_MESSAGES_ADAPTER = TypeAdapter(Messages) _MESSAGE_TYPE_UNSET = object() @@ -535,6 +538,16 @@ def resolve_optional_args( return response + def _coerce_messages(self, value: str | list) -> Messages: + if isinstance(value, str): + message = ( + TextMessage(content=value) + if self.message_type == "completion" + else UserMessage(content=value) + ) + return [message] + return _MESSAGES_ADAPTER.validate_python(value) + @final async def init_state( self, @@ -561,7 +574,7 @@ async def init_state( # Convert prompt to Pydantic messages raw_prompt = state_input.get("prompt") if isinstance(raw_prompt, (str, list)): - state["prompt"] = normalize_messages(raw_prompt, field_name="input.prompt") + state["prompt"] = self._coerce_messages(raw_prompt) state["client"] = resolve_client(client) state["model"] = model diff --git a/verifiers/envs/experimental/cli_agent_env.py b/verifiers/envs/experimental/cli_agent_env.py index 7d0e017b4c..ee13d65548 100644 --- a/verifiers/envs/experimental/cli_agent_env.py +++ b/verifiers/envs/experimental/cli_agent_env.py @@ -13,6 +13,7 @@ CreateSandboxRequest, ) from prime_tunnel import Tunnel +from pydantic import TypeAdapter import verifiers as vf from verifiers.clients import Client @@ -38,9 +39,9 @@ synthesize_stream, ) from verifiers.utils.logging_utils import print_time, truncate -from verifiers.utils.message_utils import normalize_messages logger = logging.getLogger(__name__) +_MESSAGES_ADAPTER = TypeAdapter(Messages) class AgentError(vf.InfraError): @@ -497,7 +498,7 @@ async def normalize_intercepted_messages( Assumes that agent requests arrive in OpenAI-format. """ - return await asyncio.to_thread(normalize_messages, intercepted_messages) # type: ignore + return _MESSAGES_ADAPTER.validate_python(intercepted_messages) async def normalize_response(self, response: Response) -> Response: """Hook to normalize the model response before it is stored in the trajectory. diff --git a/verifiers/envs/integrations/openenv_env.py b/verifiers/envs/integrations/openenv_env.py index ccc47783ac..0da5b26b86 100644 --- a/verifiers/envs/integrations/openenv_env.py +++ b/verifiers/envs/integrations/openenv_env.py @@ -10,19 +10,20 @@ import requests import tenacity as tc from datasets import Dataset +from pydantic import TypeAdapter import verifiers as vf from verifiers.types import ( AssistantMessage, - Message, Messages, Tool, ToolMessage, UserMessage, ) -from verifiers.utils.message_utils import from_raw_message from verifiers.utils.tool_utils import is_valid_tool_content_parts +_MESSAGES_ADAPTER = TypeAdapter(Messages) + def _optional_openenv_type(module_name: str, attr: str) -> type[Any] | None: try: @@ -853,18 +854,7 @@ def _render_observation_messages( f"OpenEnv prompt_renderer returned invalid output for {context}: " "expected a non-empty chat messages list." ) - messages: Messages = [] - for raw_message in cast(list[Any], rendered): - if isinstance(raw_message, dict): - messages.append(from_raw_message(raw_message)) - continue - if hasattr(raw_message, "role") and hasattr(raw_message, "content"): - messages.append(cast(Message, raw_message)) - continue - raise RuntimeError( - f"OpenEnv prompt_renderer returned unsupported message type for {context}: " - f"{type(raw_message).__name__}." - ) + messages = _MESSAGES_ADAPTER.validate_python(cast(list[Any], rendered)) if not messages: raise RuntimeError( f"OpenEnv prompt_renderer returned an empty messages list for {context}." diff --git a/verifiers/envs/multiturn_env.py b/verifiers/envs/multiturn_env.py index 40d1aa3c9a..28357715ec 100644 --- a/verifiers/envs/multiturn_env.py +++ b/verifiers/envs/multiturn_env.py @@ -15,10 +15,6 @@ TimeSpan, TrajectoryStep, ) -from verifiers.utils.message_utils import ( - concat_messages, - maybe_normalize_messages, -) from verifiers.utils.response_utils import ( parse_response_message, parse_response_tokens, @@ -104,10 +100,10 @@ async def get_prompt_messages(self, state: State) -> Messages: return state["prompt"] prev_turn_prompt = state["trajectory"][-1]["prompt"] prev_turn_completion = state["trajectory"][-1]["completion"] - messages = concat_messages([prev_turn_prompt, prev_turn_completion]) + messages = [*prev_turn_prompt, *prev_turn_completion] env_response = await self.env_response(messages, state) - env_response = maybe_normalize_messages(env_response, field_name="env_response") - return concat_messages([messages, env_response]) + env_response = self._coerce_messages(env_response) + return [*messages, *env_response] async def render_completion(self, state: State): """Override for rollouts with non-linear message sequences.""" @@ -116,13 +112,11 @@ async def render_completion(self, state: State): return last_prompt = state["trajectory"][-1]["prompt"] last_completion = state["trajectory"][-1]["completion"] - full_conversation = concat_messages([last_prompt, last_completion]) + full_conversation = [*last_prompt, *last_completion] if state.get("final_env_response"): final_resp = state["final_env_response"] - final_resp = maybe_normalize_messages( - final_resp, field_name="final_env_response" - ) - full_conversation = concat_messages([full_conversation, final_resp]) + final_resp = self._coerce_messages(final_resp) + full_conversation = [*full_conversation, *final_resp] prompt_messages = state["prompt"] state["completion"] = full_conversation[len(prompt_messages) :] @@ -195,9 +189,7 @@ async def rollout_loop() -> None: TimeSpan(start=start_time, end=end_time) ) - prompt_messages = maybe_normalize_messages( - prompt_messages, field_name="prompt_messages" - ) + prompt_messages = self._coerce_messages(prompt_messages) if state.get("final_env_response") is not None: continue diff --git a/verifiers/scripts/build.py b/verifiers/scripts/build.py index 07bf55aa39..c440e1079f 100644 --- a/verifiers/scripts/build.py +++ b/verifiers/scripts/build.py @@ -41,11 +41,18 @@ def _resolve_project_dir(environments_root: Path, env_id_underscore: str) -> Pat f"Environment not found: {env_path}. Expected directory '{env_id_underscore}' under {environments_root}." ) - project_dir = env_path / "proj" - if not project_dir.exists() or not project_dir.is_dir(): + candidates = [ + env_path / "proj", + env_path / env_id_underscore / "proj", + ] + project_dir = next( + (candidate for candidate in candidates if candidate.is_dir()), + None, + ) + if project_dir is None: + expected = " or ".join(str(candidate) for candidate in candidates) raise FileNotFoundError( - f"Embedded project directory not found: {project_dir}. " - "Required structure: environments//proj/" + f"Embedded OpenEnv project directory not found. Expected one of: {expected}" ) required = [ diff --git a/verifiers/scripts/eval.py b/verifiers/scripts/eval.py index 0fe698ba46..0efc658af2 100644 --- a/verifiers/scripts/eval.py +++ b/verifiers/scripts/eval.py @@ -10,7 +10,6 @@ import argparse import asyncio -import inspect import importlib.util import json import logging @@ -39,8 +38,7 @@ run_evaluations, run_evaluations_tui, ) -from verifiers.utils.env_utils import ( - env_config_annotation, +from verifiers.v1.loaders import ( env_config_child_types, import_env_module, load_env_config, @@ -326,29 +324,16 @@ def apply_env_config_cli_overrides( return dict(env_args) module = import_env_module(env_id) - env_load_func = getattr(module, "load_environment", None) - config_type: type[EnvConfig] | None - if env_load_func is None: - config_type = EnvConfig - else: - sig = inspect.signature(env_load_func) - config_type = env_config_annotation(env_load_func, sig) - if config_type is None: - raise ValueError( - "Taskset/harness CLI overrides require a v1 loader shaped as " - "load_environment(config: vf.EnvConfig)." - ) - merged_env_args = dict(env_args) base_config_data = explicit_config_data(merged_env_args.get("config", {})) - child_types = env_config_child_types(module, config_type, base_config_data) + child_types = env_config_child_types(module, EnvConfig, base_config_data) base_config = load_env_config( module, - config_type, + EnvConfig, merged_env_args.get("config", {}), child_types=child_types, ) - cli_type = env_config_cli_type(config_type, base_config, child_types) + cli_type = env_config_cli_type(EnvConfig, base_config, child_types) try: config = parse_pydantic_config_cli( cli_type, diff --git a/verifiers/scripts/init.py b/verifiers/scripts/init.py index 3643fa61e5..377ff63e0e 100644 --- a/verifiers/scripts/init.py +++ b/verifiers/scripts/init.py @@ -104,7 +104,7 @@ build-backend = "hatchling.build" [tool.hatch.build] -include = ["{{env_file}}.py", "pyproject.toml"] +include = [{{build_include}}] [tool.verifiers.eval] num_examples = 5 @@ -128,7 +128,7 @@ build-backend = "hatchling.build" [tool.hatch.build] -include = ["{{env_file}}.py", "pyproject.toml", "README.md", "proj/**/*", "proj/.build.json"] +include = ["{{env_file}}/**/*", "pyproject.toml", "README.md"] [tool.verifiers.eval] num_examples = 5 @@ -141,6 +141,10 @@ __all__ = {exports} """ +V1_INIT_TEMPLATE = """\ +\"\"\"{env_id_dash} environment package.\"\"\" +""" + V0_ENVIRONMENT_TEMPLATE = """\ import verifiers as vf @@ -156,13 +160,25 @@ def load_environment(**kwargs) -> vf.Environment: """ V1_TASKSET_TEMPLATE = """\ -import verifiers as vf +import verifiers.v1 as vf class {taskset_config_name}(vf.TasksetConfig): \"\"\"User-facing task settings for {env_id_dash}.\"\"\" system_prompt: vf.SystemPrompt = "Answer exactly." + # Optional user/toolset configs live in servers/. Wire them in here only + # when this taskset needs user simulation or tools: + # + # from .servers.example import ExampleToolsetConfig + # from .servers.user import UserConfig + # + # user: vf.UserConfig | None = UserConfig() + # toolsets: vf.ToolsetConfigs = {"example": ExampleToolsetConfig()} + + +class {task_name}(vf.Task): + answer: str class {taskset_name}(vf.Taskset[{taskset_config_name}]): @@ -172,6 +188,8 @@ class {taskset_name}(vf.Taskset[{taskset_config_name}]): metrics, rewards, and advantages on this class. \"\"\" + task_type = {task_name} + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: \"\"\"Return serializable task records as a list, generator, or Dataset.\"\"\" if split == "eval": @@ -185,13 +203,15 @@ def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: ] @vf.reward(weight=1.0) - async def correct_answer(self, task: vf.Task, state: vf.State) -> float: + async def correct_answer(self, task: {task_name}, state: vf.State) -> float: \"\"\"Score the final assistant response for one rollout.\"\"\" - messages = vf.get_messages(state.get("completion") or [], role="assistant") + messages = [ + message for message in state.completion if message.role == "assistant" + ] if not messages: return 0.0 response = str(messages[-1].content or "").strip() - return float(response == task["answer"]) + return float(response == task.answer) def load_taskset(config: {taskset_config_name}) -> {taskset_name}: @@ -201,6 +221,7 @@ def load_taskset(config: {taskset_config_name}) -> {taskset_name}: V1_HARNESS_TEMPLATE = """\ +import verifiers.v1 as vf class {harness_config_name}(vf.HarnessConfig): \"\"\"Execution settings for {env_id_dash}.\"\"\" @@ -219,37 +240,77 @@ def load_harness(config: {harness_config_name}) -> {harness_name}: return {harness_name}(config=config) """ +V1_SERVERS_INIT_TEMPLATE = """\ +\"\"\"Optional user/toolset implementations for {env_id_dash}.\"\"\" +""" + +V1_USER_INIT_TEMPLATE = """\ +from .config import UserConfig + +__all__ = ["UserConfig"] +""" + +V1_USER_CONFIG_TEMPLATE = """\ +import verifiers.v1 as vf -V1_ENV_LOADER_TEMPLATE = """\ -def load_environment(config: vf.EnvConfig) -> vf.Env: - \"\"\"Loader pattern for all Taskset/Harness environments.\"\"\" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), +class UserConfig(vf.UserConfig): + pass +""" + +V1_USER_SERVER_TEMPLATE = """\ +import verifiers.v1 as vf + +from .config import UserConfig + + +class User(vf.User[UserConfig]): + @vf.user( + args={{ + "task": "task", + "state": "state", + "transcript": "transcript", + }} ) + def respond(self, task: dict, state: dict, transcript: list[dict]) -> dict: + \"\"\"Return user messages, extras updates, or stop signals.\"\"\" + _ = task, state, transcript + return {{"messages": []}} """ -V1_ENVIRONMENT_TEMPLATE = V1_TASKSET_TEMPLATE + V1_ENV_LOADER_TEMPLATE -V1_HARNESS_ENVIRONMENT_TEMPLATE = ( - V1_TASKSET_TEMPLATE + V1_HARNESS_TEMPLATE + V1_ENV_LOADER_TEMPLATE -) +V1_TOOLSET_INIT_TEMPLATE = """\ +from .config import ExampleToolsetConfig -OPENENV_ENVIRONMENT_TEMPLATE = """\ -import verifiers as vf +__all__ = ["ExampleToolsetConfig"] +""" + +V1_TOOLSET_CONFIG_TEMPLATE = """\ +import verifiers.v1 as vf + + +class ExampleToolsetConfig(vf.ToolsetConfig): + pass +""" + +V1_TOOLS_SERVER_TEMPLATE = """\ +import verifiers.v1 as vf + +from .config import ExampleToolsetConfig + + +class ExampleToolset(vf.Toolset[ExampleToolsetConfig]): + @vf.tool + def reverse_text(self, text: str) -> str: + \"\"\"Example tool stub.\"\"\" + return text[::-1] +""" + +OPENENV_TASKSET_TEMPLATE = """\ from tasksets import OpenEnvTaskset, OpenEnvTasksetConfig def load_taskset(config: OpenEnvTasksetConfig) -> OpenEnvTaskset: return OpenEnvTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - \"\"\"Loader pattern for all Taskset/Harness environments.\"\"\" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) """ OPENENV_PROJ_README_TEMPLATE = """\ @@ -413,6 +474,7 @@ def init_environment( env_id_dash = env.replace("_", "-") env_id_underscore = env_id_dash.replace("-", "_") taskset_config_name = _class_name(env_id_underscore, "TasksetConfig") + task_name = _class_name(env_id_underscore, "Task") taskset_name = _class_name(env_id_underscore, "Taskset") harness_config_name = _class_name(env_id_underscore, "HarnessConfig") harness_name = _class_name(env_id_underscore, "Harness") @@ -438,66 +500,122 @@ def init_environment( else: print(f"README.md already exists at {readme_file}, skipping...") + package_layout = multi_file or v1 or openenv + # create pyproject.toml if it doesn't exist pyproject_file = local_dir / "pyproject.toml" if not pyproject_file.exists(): pyproject_template = ( OPENENV_PYPROJECT_TEMPLATE if openenv else PYPROJECT_TEMPLATE ) + build_include = ( + f'"{env_id_underscore}/**/*", "pyproject.toml", "README.md"' + if package_layout + else f'"{env_id_underscore}.py", "pyproject.toml"' + ) pyproject_file.write_text( - pyproject_template.format(env_id=env_id_dash, env_file=env_id_underscore) + pyproject_template.format( + env_id=env_id_dash, + env_file=env_id_underscore, + build_include=build_include, + ) ) else: print(f"pyproject.toml already exists at {pyproject_file}, skipping...") # create environment directory if it doesn't exist - environment_dir = local_dir / env_id_underscore if multi_file else local_dir + environment_dir = local_dir / env_id_underscore if package_layout else local_dir environment_dir.mkdir(parents=True, exist_ok=True) # create init file if it doesn't exist - if multi_file: + if package_layout: init_file = environment_dir / "__init__.py" if not init_file.exists(): - exports = ["load_environment"] - if v1: - exports.append("load_taskset") - if v1 and with_harness and not openenv: - exports.append("load_harness") - init_file.write_text( - INIT_TEMPLATE.format( - env_id=env_id_underscore, - imports=", ".join(exports), - exports=repr(exports), + if v1 or openenv: + init_file.write_text(V1_INIT_TEMPLATE.format(env_id_dash=env_id_dash)) + else: + exports = ["load_environment"] + init_file.write_text( + INIT_TEMPLATE.format( + env_id=env_id_underscore, + imports=", ".join(exports), + exports=repr(exports), + ) ) - ) else: print(f"__init__.py already exists at {init_file}, skipping...") - # create environment file if it doesn't exist - environment_file = environment_dir / f"{env_id_underscore}.py" - if not environment_file.exists(): - if openenv: - template = OPENENV_ENVIRONMENT_TEMPLATE - elif v1 and with_harness: - template = V1_HARNESS_ENVIRONMENT_TEMPLATE - elif v1: - template = V1_ENVIRONMENT_TEMPLATE + if v1 or openenv: + taskset_file = environment_dir / "taskset.py" + if not taskset_file.exists(): + taskset_template = ( + OPENENV_TASKSET_TEMPLATE if openenv else V1_TASKSET_TEMPLATE + ) + taskset_file.write_text( + taskset_template.replace("{env_id_dash}", env_id_dash) + .replace("{taskset_config_name}", taskset_config_name) + .replace("{task_name}", task_name) + .replace("{taskset_name}", taskset_name) + .replace("{env_id_underscore}", env_id_underscore) + ) else: - template = V0_ENVIRONMENT_TEMPLATE - environment_file.write_text( - template.replace("{env_id_dash}", env_id_dash) - .replace("{taskset_config_name}", taskset_config_name) - .replace("{taskset_name}", taskset_name) - .replace("{harness_config_name}", harness_config_name) - .replace("{harness_name}", harness_name) - ) + print(f"taskset.py already exists at {taskset_file}, skipping...") + + if with_harness and not openenv: + harness_file = environment_dir / "harness.py" + if not harness_file.exists(): + harness_file.write_text( + V1_HARNESS_TEMPLATE.replace("{env_id_dash}", env_id_dash) + .replace("{harness_config_name}", harness_config_name) + .replace("{harness_name}", harness_name) + ) + else: + print(f"harness.py already exists at {harness_file}, skipping...") else: - print( - f"{env_id_underscore}.py already exists at {environment_file}, skipping..." - ) + # create environment file if it doesn't exist + environment_file = environment_dir / f"{env_id_underscore}.py" + if not environment_file.exists(): + environment_file.write_text( + V0_ENVIRONMENT_TEMPLATE.replace("{env_id_dash}", env_id_dash) + .replace("{taskset_config_name}", taskset_config_name) + .replace("{task_name}", task_name) + .replace("{taskset_name}", taskset_name) + .replace("{harness_config_name}", harness_config_name) + .replace("{harness_name}", harness_name) + .replace("{env_id_underscore}", env_id_underscore) + ) + else: + print( + f"{env_id_underscore}.py already exists at {environment_file}, skipping..." + ) + + if v1 and not openenv: + servers_dir = environment_dir / "servers" + servers_dir.mkdir(parents=True, exist_ok=True) + server_files = { + "__init__.py": V1_SERVERS_INIT_TEMPLATE, + "user/__init__.py": V1_USER_INIT_TEMPLATE, + "user/config.py": V1_USER_CONFIG_TEMPLATE, + "user/user.py": V1_USER_SERVER_TEMPLATE, + "example/__init__.py": V1_TOOLSET_INIT_TEMPLATE, + "example/config.py": V1_TOOLSET_CONFIG_TEMPLATE, + "example/toolset.py": V1_TOOLS_SERVER_TEMPLATE, + } + for filename, template in server_files.items(): + server_file = servers_dir / filename + server_file.parent.mkdir(parents=True, exist_ok=True) + if not server_file.exists(): + server_file.write_text( + template.format( + env_id_dash=env_id_dash, + env_id_underscore=env_id_underscore, + ) + ) + else: + print(f"{server_file} already exists, skipping...") if openenv: - _init_openenv_proj(local_dir, env_id_dash, env_id_underscore) + _init_openenv_proj(environment_dir, env_id_dash, env_id_underscore) return local_dir diff --git a/verifiers/types.py b/verifiers/types.py index 4242f8a86f..577601548e 100644 --- a/verifiers/types.py +++ b/verifiers/types.py @@ -26,6 +26,7 @@ Field, computed_field, field_validator, + model_validator, ) from verifiers.errors import Error @@ -165,6 +166,24 @@ class ToolCall(CustomBaseModel): name: str arguments: str + @model_validator(mode="before") + @classmethod + def normalize_openai_tool_call(cls, value: object) -> object: + if not isinstance(value, Mapping): + return value + data = dict(value) + function = data.get("function") + if not isinstance(function, Mapping): + return data + if "name" not in data: + data["name"] = function.get("name") + if "arguments" not in data and "arguments" in function: + arguments = function["arguments"] + data["arguments"] = ( + arguments if isinstance(arguments, str) else json.dumps(arguments) + ) + return data + ThinkingBlock: TypeAlias = AnthropicThinkingBlock | RedactedThinkingBlock diff --git a/verifiers/utils/async_utils.py b/verifiers/utils/async_utils.py index 1195599dab..78fcacb7e8 100644 --- a/verifiers/utils/async_utils.py +++ b/verifiers/utils/async_utils.py @@ -178,7 +178,12 @@ def reraise_error_from_state(result, error_types: tuple[type[Exception], ...]): reraise_one(result.get("error"), error_types) elif isinstance(result, list): for state in result: - reraise_one(state.get("error"), error_types) + if isinstance(state, dict): + reraise_one(state.get("error"), error_types) + else: + reraise_one(getattr(state, "error", None), error_types) + else: + reraise_one(getattr(result, "error", None), error_types) def log_retry(retry_state: tc.RetryCallState) -> None: """Log a warning with the exception and the number of attempts.""" diff --git a/verifiers/utils/env_utils.py b/verifiers/utils/env_utils.py index c71d36dbf6..1b3e2b8ac8 100644 --- a/verifiers/utils/env_utils.py +++ b/verifiers/utils/env_utils.py @@ -1,34 +1,22 @@ +from __future__ import annotations + import importlib import inspect import logging -import sys -from collections.abc import Mapping -from types import ModuleType, UnionType +from types import ModuleType from typing import ( Callable, - TypeAlias, - Union, - cast, - get_args, - get_origin, - get_type_hints, + TYPE_CHECKING, ) -from pydantic import BaseModel from verifiers.envs.environment import Environment from verifiers.utils.config_utils import MissingKeyError -from verifiers.v1.env import Env, EnvConfig -from verifiers.v1.harness import Harness, HarnessConfig -from verifiers.v1.taskset import Taskset, TasksetConfig -from verifiers.v1.types import ConfigData, ConfigValue -from verifiers.v1.utils.config_utils import coerce_config, explicit_config_data -EnvConfigLoadData: TypeAlias = dict[str, ConfigValue | BaseModel] -EnvConfigChildInput: TypeAlias = ConfigData | EnvConfigLoadData -EnvConfigInput: TypeAlias = EnvConfig | ConfigData +if TYPE_CHECKING: + from verifiers.v1.env import Env -def load_environment(env_id: str, **env_args) -> Environment: +def load_environment(env_id: str, **env_args) -> Environment | "Env": logger = logging.getLogger("verifiers.utils.env_utils") logger.info(f"Loading environment: {env_id}") @@ -40,6 +28,8 @@ def load_environment(env_id: str, **env_args) -> Environment: module, "load_environment", None ) if env_load_func is None: + from verifiers.v1.loaders import load_environment_from_components + env_instance = load_environment_from_components(module, env_args) env_instance.env_id = env_instance.env_id or env_id env_instance.env_args = env_instance.env_args or env_args @@ -88,8 +78,7 @@ def load_environment(env_id: str, **env_args) -> Environment: if default_values: logger.info(f"Using default args: {', '.join(default_values)}") - call_env_args = prepare_typed_env_config(module, env_load_func, sig, env_args) - env_instance: Environment = env_load_func(**call_env_args) + env_instance: Environment = env_load_func(**env_args) env_instance.env_id = env_instance.env_id or env_id env_instance.env_args = env_instance.env_args or env_args @@ -119,308 +108,3 @@ def env_module_name(env_id: str) -> str: def import_env_module(env_id: str) -> ModuleType: return importlib.import_module(env_module_name(env_id)) - - -def caller_module() -> ModuleType: - frame = inspect.currentframe() - try: - if frame is None or frame.f_back is None or frame.f_back.f_back is None: - raise RuntimeError("Could not resolve caller module.") - module_name = frame.f_back.f_back.f_globals.get("__name__") - if not isinstance(module_name, str): - raise RuntimeError("Caller module has no __name__.") - module = sys.modules.get(module_name) - if not isinstance(module, ModuleType): - raise RuntimeError(f"Caller module {module_name!r} is not loaded.") - return module - finally: - del frame - - -def load_taskset( - env_id: str | None = None, - *, - config: TasksetConfig | ConfigData | None = None, -) -> Taskset: - module = caller_module() if env_id is None else import_env_module(env_id) - return load_taskset_from_module(module, config=config) - - -def load_harness( - env_id: str | None = None, - *, - config: HarnessConfig | ConfigData | None = None, -) -> Harness: - module = caller_module() if env_id is None else import_env_module(env_id) - return load_harness_from_module(module, config=config) - - -def load_taskset_from_module( - module: ModuleType, - *, - config: TasksetConfig | ConfigData | None = None, -) -> Taskset: - factory = getattr(module, "load_taskset", None) - if factory is None: - taskset_id = child_loader_id(config, "taskset_id") - if taskset_id is not None and env_module_name(taskset_id) != module.__name__: - return load_taskset(taskset_id, config=config) - raise AttributeError( - f"Module '{module.__name__}' does not expose load_taskset, and " - "config.taskset_id is not set to a taskset loader package." - ) - config_type = factory_config_type(module, "load_taskset", TasksetConfig) - if config_type is None: - raise TypeError(f"{module.__name__}.load_taskset must accept config.") - taskset = factory( - config=coerce_config(cast(type[TasksetConfig], config_type), config) - ) - if not isinstance(taskset, Taskset): - raise TypeError(f"{module.__name__}.load_taskset must return a Taskset.") - return taskset - - -def load_harness_from_module( - module: ModuleType, - *, - config: HarnessConfig | ConfigData | None = None, -) -> Harness: - factory = getattr(module, "load_harness", None) - if factory is None: - harness_id = child_loader_id(config, "harness_id") - if harness_id is not None: - if env_module_name(harness_id) == module.__name__: - raise AttributeError( - f"Module '{module.__name__}' does not expose load_harness." - ) - return load_harness(harness_id, config=config) - return Harness(config=coerce_config(HarnessConfig, config)) - config_type = factory_config_type(module, "load_harness", HarnessConfig) - if config_type is None: - raise TypeError(f"{module.__name__}.load_harness must accept config.") - harness = factory( - config=coerce_config(cast(type[HarnessConfig], config_type), config) - ) - if not isinstance(harness, Harness): - raise TypeError(f"{module.__name__}.load_harness must return a Harness.") - return harness - - -def prepare_typed_env_config( - module: ModuleType, - env_load_func: Callable[..., Environment], - sig: inspect.Signature, - env_args: dict, -) -> dict: - config_type = env_config_annotation(env_load_func, sig) - if config_type is None: - return env_args - - config = env_args.get("config", {}) - if config is None: - raise TypeError("load_environment config must be a concrete EnvConfig object.") - - call_env_args = dict(env_args) - call_env_args["config"] = load_env_config(module, config_type, config) - return call_env_args - - -def load_environment_from_components( - module: ModuleType, - env_args: dict, -) -> Env: - extra_args = set(env_args) - {"config"} - if extra_args: - raise TypeError( - "Default Taskset/Harness environment loading only accepts config; " - f"got {sorted(extra_args)}." - ) - config = load_env_config(module, EnvConfig, env_args.get("config", {})) - return Env( - taskset=load_taskset_from_module(module, config=config.taskset), - harness=load_harness_from_module(module, config=config.harness), - ) - - -def env_config_annotation( - env_load_func: Callable[..., Environment], - sig: inspect.Signature, -) -> type[EnvConfig] | None: - if "config" not in sig.parameters: - return None - try: - annotation = get_type_hints(env_load_func).get( - "config", sig.parameters["config"].annotation - ) - except Exception: - annotation = sig.parameters["config"].annotation - return env_config_type(annotation) - - -def env_config_type(annotation: object) -> type[EnvConfig] | None: - if annotation is inspect.Parameter.empty: - return None - origin = get_origin(annotation) - if origin in (Union, UnionType): - args = [arg for arg in get_args(annotation) if arg is not type(None)] - if len(args) == 1: - annotation = args[0] - if isinstance(annotation, type) and issubclass(annotation, EnvConfig): - return annotation - return None - - -def load_env_config( - module: ModuleType, - config_type: type[EnvConfig], - value: EnvConfigInput, - *, - child_types: Mapping[str, type[BaseModel]] | None = None, -) -> EnvConfig: - data: EnvConfigLoadData - if isinstance(value, config_type): - data = dict(explicit_config_data(value)) - elif isinstance(value, BaseModel): - raise TypeError( - f"load_environment config must be {config_type.__name__}; " - f"got {type(value).__name__}." - ) - elif not isinstance(value, Mapping): - raise TypeError("load_environment config must be a mapping or EnvConfig.") - else: - data = dict(value) - resolved_child_types = ( - env_config_child_types(module, config_type, data) - if child_types is None - else child_types - ) - defaults: EnvConfig | None = None - for field_name, child_type in resolved_child_types.items(): - if field_name not in data: - defaults = config_type() if defaults is None else defaults - child = getattr(defaults, field_name) - data[field_name] = child if isinstance(child, child_type) else child_type() - continue - child = data[field_name] - if isinstance(child, child_type): - continue - if child is None: - raise TypeError(f"config.{field_name} cannot be None.") - if not isinstance(child, BaseModel | dict): - raise TypeError(f"config.{field_name} must be a mapping or config object.") - data[field_name] = child_type.model_validate(explicit_config_data(child)) - config = config_type.model_validate(data) - for field_name, child_type in resolved_child_types.items(): - child = getattr(config, field_name) - if not isinstance(child, child_type): - raise TypeError( - f"config.{field_name} must be {child_type.__name__}; " - f"got {type(child).__name__}." - ) - return config - - -def env_config_child_types( - module: ModuleType, - config_type: type[EnvConfig], - value: EnvConfigChildInput | None = None, -) -> dict[str, type[BaseModel]]: - child_types: dict[str, type[BaseModel]] = {} - for field_name, id_field, factory_name, base_type in ( - ("taskset", "taskset_id", "load_taskset", TasksetConfig), - ("harness", "harness_id", "load_harness", HarnessConfig), - ): - field_type = config_type_from_annotation( - config_type.model_fields[field_name].annotation, - base_type, - f"{config_type.__name__}.{field_name}", - ) - factory_type = factory_config_type(module, factory_name, base_type) - child_config = value.get(field_name) if value is not None else None - if ( - factory_type is None - and field_type is base_type - and child_config_requires_loader_type(child_config, base_type) - ): - loader_id = child_loader_id(child_config, id_field) - if loader_id is not None and env_module_name(loader_id) != module.__name__: - factory_type = factory_config_type( - import_env_module(loader_id), factory_name, base_type - ) - if factory_type is not None: - if not issubclass(factory_type, field_type): - raise TypeError( - f"{module.__name__}.{factory_name} config type " - f"{factory_type.__name__} does not match " - f"{config_type.__name__}.{field_name}: {field_type.__name__}." - ) - child_types[field_name] = factory_type - else: - child_types[field_name] = field_type - return child_types - - -def child_config_requires_loader_type( - config: object, - base_type: type[BaseModel], -) -> bool: - if not isinstance(config, Mapping): - return False - base_fields = set(base_type.model_fields) | {"id"} - return bool(set(config) - base_fields) - - -def child_loader_id(config: object, id_field: str) -> str | None: - if isinstance(config, BaseModel): - value = config.__dict__.get(id_field) - elif isinstance(config, Mapping): - config_data = dict(config) - value = config_data.get(id_field) or config_data.get("id") - else: - return None - if value is None: - return None - if not isinstance(value, str) or not value: - raise TypeError(f"config.{id_field} must be a non-empty string.") - return value - - -def factory_config_type( - module: ModuleType, - factory_name: str, - base_type: type[BaseModel], -) -> type[BaseModel] | None: - factory = getattr(module, factory_name, None) - if factory is None: - return None - signature = inspect.signature(factory) - if "config" not in signature.parameters: - raise TypeError(f"{module.__name__}.{factory_name} must accept config.") - try: - annotation = get_type_hints(factory).get( - "config", signature.parameters["config"].annotation - ) - except Exception: - annotation = signature.parameters["config"].annotation - return config_type_from_annotation( - annotation, - base_type, - f"{module.__name__}.{factory_name}.config", - ) - - -def config_type_from_annotation( - annotation: object, - base_type: type[BaseModel], - context: str, -) -> type[BaseModel]: - if annotation is inspect.Parameter.empty: - raise TypeError(f"{context} must be annotated.") - origin = get_origin(annotation) - if origin in (Union, UnionType): - args = [arg for arg in get_args(annotation) if arg is not type(None)] - if len(args) == 1: - annotation = args[0] - if isinstance(annotation, type) and issubclass(annotation, base_type): - return annotation - raise TypeError(f"{context} must be a {base_type.__name__} subclass.") diff --git a/verifiers/utils/eval_utils.py b/verifiers/utils/eval_utils.py index f33006dea4..773cd3a330 100644 --- a/verifiers/utils/eval_utils.py +++ b/verifiers/utils/eval_utils.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import itertools import json @@ -454,7 +456,7 @@ def eval_env_id(config: Mapping[str, Any], section: str) -> str: if isinstance(env_config, Mapping): taskset = env_config.get("taskset") if isinstance(taskset, Mapping): - taskset_id = taskset.get("id") or taskset.get("taskset_id") + taskset_id = taskset.get("id") if isinstance(taskset_id, str) and taskset_id: return taskset_id raise ValueError(f"{section} must contain env_id or taskset.id.") @@ -890,6 +892,14 @@ def _values(key: str) -> list[float]: v = t[key] if isinstance(v, dict): v = v.get("duration", 0.0) + elif isinstance(v, list): + total = 0.0 + for item in v: + if isinstance(item, dict): + total += float(item.get("duration", 0.0)) + elif isinstance(item, int | float): + total += float(item) + v = total out.append(float(v)) return out @@ -1014,6 +1024,19 @@ def get_log_level(verbose: bool) -> str: return "DEBUG" if verbose else os.getenv("VF_LOG_LEVEL", "INFO") +def effective_max_concurrent(config: EvalConfig) -> int: + if ( + not config.independent_scoring + and config.max_concurrent > 0 + and config.rollouts_per_example > 1 + ): + max_concurrent = math.ceil(config.max_concurrent / config.rollouts_per_example) + if config.num_examples > 0: + max_concurrent = min(max_concurrent, config.num_examples) + return max_concurrent + return config.max_concurrent + + @contextmanager def quiet_datasets(): prev_level = ds_logging.get_verbosity() @@ -1040,12 +1063,14 @@ async def run_evaluation( with maybe_suppress_logs: vf_env = vf.load_environment(env_id=config.env_id, **config.env_args) + from verifiers.v1.env import Env as V1Env + + results_path = config.resume_path or get_eval_results_path(config) # set extra environment kwargs - if config.extra_env_kwargs: + if config.extra_env_kwargs and not isinstance(vf_env, V1Env): logger.info(f"Setting extra environment kwargs: {config.extra_env_kwargs}") vf_env.set_kwargs(**config.extra_env_kwargs) - results_path = config.resume_path or get_eval_results_path(config) if config.client_config.endpoint_configs: pricing_urls = [ endpoint.api_base_url for endpoint in config.client_config.endpoint_configs @@ -1057,6 +1082,26 @@ async def run_evaluation( model_pricing = (await fetch_prime_pricing()).get(config.model) on_progress = _with_eval_metadata(on_progress, model_pricing, config.name) + if isinstance(vf_env, V1Env): + from verifiers.v1.eval import run_evaluation as run_v1_evaluation + + outputs = await run_v1_evaluation( + vf_env, + config, + results_path, + on_start, + on_progress, + on_log, + ) + metadata_changed = _attach_metadata_name(outputs["metadata"], config.name) + if _attach_metadata_cost( + outputs["metadata"], model_pricing, outputs["outputs"] + ): + metadata_changed = True + if metadata_changed and config.save_results: + await asyncio.to_thread(save_metadata, outputs["metadata"], results_path) + return outputs + try: if not config.disable_env_server: extra_env_kwargs = dict(config.extra_env_kwargs) @@ -1107,22 +1152,6 @@ async def run_evaluation( f"Configuration: num_examples={config.num_examples}, rollouts_per_example={config.rollouts_per_example}, shuffle={config.shuffle}, shuffle_seed={config.shuffle_seed}, max_concurrent={config.max_concurrent}" ) - effective_group_max_concurrent = config.max_concurrent - if ( - not config.independent_scoring - and config.max_concurrent > 0 - and config.rollouts_per_example > 1 - ): - # Grouped scoring applies the semaphore at group level. Convert - # rollout-level concurrency to group-level slots. - effective_group_max_concurrent = math.ceil( - config.max_concurrent / config.rollouts_per_example - ) - if config.num_examples > 0: - effective_group_max_concurrent = min( - effective_group_max_concurrent, config.num_examples - ) - outputs = await vf_env.evaluate( client=config.client_config, model=config.model, @@ -1131,7 +1160,7 @@ async def run_evaluation( rollouts_per_example=config.rollouts_per_example, shuffle=config.shuffle, shuffle_seed=config.shuffle_seed, - max_concurrent=effective_group_max_concurrent, + max_concurrent=effective_max_concurrent(config), results_path=results_path, state_columns=config.state_columns, save_results=config.save_results, diff --git a/verifiers/utils/message_utils.py b/verifiers/utils/message_utils.py index f6fc717c45..bef5e71a30 100644 --- a/verifiers/utils/message_utils.py +++ b/verifiers/utils/message_utils.py @@ -1,228 +1,14 @@ import json import logging import re -from collections.abc import Mapping, Sequence -from typing import Any, Literal, TypeAlias, cast, overload +from collections.abc import Mapping +from typing import Any from rich.text import Text - -from verifiers.types import ( - AssistantMessage, - ImageUrlContentPart, - InputAudioContentPart, - Message, - Messages, - SystemMessage, - TextContentPart, - TextMessage, - ToolMessage, - UserMessage, -) +from verifiers.types import Messages logger = logging.getLogger(__name__) -MessageLike: TypeAlias = Message | Mapping[str, object] -MessageInput: TypeAlias = str | Sequence[MessageLike] -MessageRole: TypeAlias = Literal["text", "system", "user", "assistant", "tool"] - - -def _normalize_raw_message_content(message: dict[str, Any]) -> dict[str, Any]: - content = message.get("content") - if isinstance(content, list): - normalized_parts = [] - for part in content: - if isinstance(part, dict): - part_type = part.get("type") - if part_type == "text": - normalized_parts.append(TextContentPart.model_validate(part)) - elif part_type == "image_url": - normalized_parts.append(ImageUrlContentPart.model_validate(part)) - elif part_type == "input_audio": - normalized_parts.append(InputAudioContentPart.model_validate(part)) - else: - normalized_parts.append(part) - else: - normalized_parts.append(part) - message = dict(message) - message["content"] = normalized_parts - return message - - -def _normalize_raw_tool_calls(message: dict[str, Any]) -> dict[str, Any]: - if message.get("role") != "assistant": - return message - - tool_calls = message.get("tool_calls") - if not isinstance(tool_calls, list): - return message - - normalized_tool_calls: list[Any] = [] - for tool_call in tool_calls: - if not isinstance(tool_call, dict): - normalized_tool_calls.append(tool_call) - continue - - if "name" in tool_call and "arguments" in tool_call: - normalized_tool_calls.append(tool_call) - continue - - function = tool_call.get("function") - if not isinstance(function, dict): - normalized_tool_calls.append(tool_call) - continue - - name = function.get("name") - arguments = function.get("arguments") - if not isinstance(name, str): - normalized_tool_calls.append(tool_call) - continue - - if isinstance(arguments, str): - arguments_str = arguments - else: - try: - arguments_str = json.dumps(arguments if arguments is not None else {}) - except (TypeError, ValueError): - arguments_str = str(arguments) - - tool_call_id = tool_call.get("id") - if not isinstance(tool_call_id, str): - tool_call_id = name - - normalized_tool_calls.append( - { - "id": tool_call_id, - "name": name, - "arguments": arguments_str, - } - ) - - message = dict(message) - message["tool_calls"] = normalized_tool_calls - return message - - -def from_raw_message(message: dict) -> Message: - """Convert a raw dict to the appropriate Pydantic message type.""" - message = _normalize_raw_message_content(message) - message = _normalize_raw_tool_calls(message) - if message["role"] == "text": - return TextMessage.model_validate(message) - elif message["role"] == "system": - return SystemMessage.model_validate(message) - elif message["role"] == "user": - return UserMessage.model_validate(message) - elif message["role"] == "assistant": - return AssistantMessage.model_validate(message) - elif message["role"] == "tool": - return ToolMessage.model_validate(message) - else: - raise ValueError(f"Unknown role: {message['role']}") - - -def normalize_messages( - value: MessageInput, *, field_name: str = "messages" -) -> Messages: - """Normalize raw/string message inputs into provider-agnostic Message objects.""" - if isinstance(value, str): - return [TextMessage(content=value)] - normalized: Messages = [] - for message in value: - if isinstance(message, dict): - normalized.append(from_raw_message(dict(message))) - continue - if hasattr(message, "role") and hasattr(message, "content"): - normalized.append(cast(Message, message)) - continue - raise TypeError( - f"Invalid {field_name} item type: {type(message).__name__}. " - "Expected vf.Message-like objects." - ) - return normalized - - -@overload -def get_messages( - messages: Sequence[MessageLike], role: Literal["assistant"] -) -> list[AssistantMessage]: ... - - -@overload -def get_messages( - messages: Sequence[MessageLike], role: Literal["system"] -) -> list[SystemMessage]: ... - - -@overload -def get_messages( - messages: Sequence[MessageLike], role: Literal["user"] -) -> list[UserMessage]: ... - - -@overload -def get_messages( - messages: Sequence[MessageLike], role: Literal["tool"] -) -> list[ToolMessage]: ... - - -@overload -def get_messages( - messages: Sequence[MessageLike], role: Literal["text"] -) -> list[TextMessage]: ... - - -@overload -def get_messages(messages: Sequence[MessageLike], role: None = None) -> Messages: ... - - -def get_messages( - messages: Sequence[MessageLike], role: MessageRole | None = None -) -> Messages: - """Return typed transcript messages, optionally filtered by role.""" - normalized = normalize_messages(messages) - if role is None: - return normalized - return [message for message in normalized if message.role == role] - - -def message_role(message: MessageLike) -> str | None: - if isinstance(message, Mapping): - value = message.get("role") - else: - value = getattr(message, "role", None) - return value if isinstance(value, str) else None - - -def maybe_normalize_messages( - value: Messages | str, - *, - field_name: str = "messages", -) -> Messages: - """Normalize messages only if needed, logging a warning on first occurrence.""" - from verifiers.utils.logging_utils import warning_once - - requires_normalize = not isinstance(value, list) or not all( - isinstance(m, Message) for m in value - ) - if not requires_normalize: - return cast(Messages, value) - warning_once( - logger, - f"{field_name} returned raw dicts/strings instead of vf.Messages. This" - " repeatedly triggers normalize_messages(), causing unnecessary" - " Pydantic validation overhead. Return vf.Message types (e.g." - " vf.UserMessage, vf.AssistantMessage) to avoid this.", - ) - return normalize_messages(value, field_name=field_name) - - -def concat_messages(messages_list: list[Messages]) -> Messages: - """Concatenate multiple Messages lists into one.""" - result = [] - for messages in messages_list: - result.extend(messages) - return result - def message_to_printable(message: Any) -> Any: """ diff --git a/verifiers/v1/ENVIRONMENT_BEST_PRACTICES.md b/verifiers/v1/ENVIRONMENT_BEST_PRACTICES.md deleted file mode 100644 index cb36cc1520..0000000000 --- a/verifiers/v1/ENVIRONMENT_BEST_PRACTICES.md +++ /dev/null @@ -1,186 +0,0 @@ -# v1 Environment Contract - -This is the strict authoring contract for v1 Taskset/Harness environments. Use -it before adding, migrating, reviewing, or agent-generating v1 environment code. -The full walkthrough is in `docs/byo-harness.md`. - -## Golden Loader Shape - -Environment modules expose one root loader: - -```python -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -When the environment has a custom taskset config, expose a typed child loader: - -```python -def load_taskset(config: MyTasksetConfig) -> MyTaskset: - return MyTaskset(config=config) -``` - -When the environment has a custom harness config, expose a typed child loader: - -```python -def load_harness(config: MyHarnessConfig) -> MyHarness: - return MyHarness(config=config) -``` - -Those child loader annotations are the config contract. `load_environment` -stays typed as `vf.EnvConfig`; the framework uses the child annotations to -coerce `config.taskset` and `config.harness`. - -## Hard Rules - -1. Import Verifiers as `import verifiers as vf`. -2. Use `vf.Taskset`, `vf.Harness`, and `vf.Env` for new reusable environments. -3. Use `XXXConfig` Pydantic config classes for structured settings. -4. Put task behavior on the taskset config/class. -5. Put execution behavior on the harness config/class. -6. Keep `load_environment(config: vf.EnvConfig)` as-is; implement the config surface through taskset and harness configs. -7. Do not accept root loader kwargs for taskset or harness fields. -8. Do not subclass `vf.Env` for ordinary environment packages. -9. Do not subclass `vf.EnvConfig` just to narrow child config types. -10. Do not override `Taskset.__init__`, `Harness.__init__`, or `User.__init__`. -11. Do not synthesize fallback configs, accept `None`, or mutate configs inside - loaders. -12. Do not put system messages inside `task["prompt"]`. -13. Do not pass non-serializable callables or `Path` objects through config; - use import-ref strings. -14. Do not hide one-off behavior in private helper methods or detached helper - functions at the bottom of an environment file. - -Break these rules only when there is a concrete, documented framework-boundary -reason. Do not add escape hatches to make a local implementation easier. - -## Start With Tasksets - -Start with a self-contained taskset and the base harness. Tool-use tasks, -LLM-judged tasks, multimodal tasks, sandboxed tools, and simple multi-turn user -simulations should all be tasksets first. - -Add a custom harness only when the environment owns a reusable execution -protocol: command agents, third-party agent frameworks, browser or desktop -loops, endpoint interception, primary sandbox placement, or program execution -that can attempt arbitrary tasks. - -## Ownership Rules - -Tasksets own: - -- task loading and split selection; -- task prompts, answers, task metadata, and task controls; -- task-owned toolsets and users; -- task-specific setup, update, stop, cleanup, metrics, rewards, and advantages; -- task-owned objects, bindings, artifacts, files, dirs, and sandbox overrides. - -Harnesses own: - -- rollout execution and model/client defaults; -- programs, command agents, framework adapters, endpoint interception, and - protocol translation; -- primary sandbox placement and reusable execution setup; -- harness-owned toolsets, objects, bindings, artifacts, metrics, and cleanup. - -If a tool defines the task's action space, observations, or success condition, -it belongs to the taskset. If a class only describes how a model attempts any -task, it belongs to the harness. - -## Config Rules - -- Config fields should be serializable and stable enough to appear in TOML. -- Static system prompts belong in the owning config: - `TasksetConfig.system_prompt` for task policy and - `HarnessConfig.system_prompt` for execution policy. -- System prompt resolution is per task: `task["system_prompt"]` overrides - `TasksetConfig.system_prompt` for the taskset side, then - `HarnessConfig.system_prompt_strategy` resolves that side against the harness - side. -- The default system prompt strategy is `HT`; available strategies are `HT`, - `TH`, `H_OR_T`, `T_OR_H`, `H`, `T`, and `REJECT`. -- File-backed GEPA prompts should use `vf.SystemPromptConfig(path="...")`. -- Override `load_system_prompt(config)` only for computed prompt loading. -- Shared dependencies use `objects` and `bindings`; object entries are loader - specs, not pre-initialized objects. -- Users are configured with `UserConfig` subclasses and implemented by - `User.get_response(...)`, not by passing callable users. -- Tools are exposed through `vf.Toolset`; tasks only show/hide toolsets and - tools. -- Runtime-only resources live on `state` or runtime-managed owners, not on task - data or config. -- Do not add generic split config fields that duplicate `load_tasks(split=...)`. - Use config only when a split choice is a real taskset setting - rather than an adapter detail. - -## Task Rules - -Task records are JSON-serializable and become immutable `vf.Task` objects during -rollout. Common top-level fields are: - -- `prompt` -- `system_prompt` -- `answer` -- `info` -- `max_turns` -- `toolsets` -- `tools` -- `sandbox` -- `program` -- `artifacts` - -Users should not need to manage task IDs. Include upstream IDs only as ordinary -task metadata when they matter. - -Prefer returning a `datasets.Dataset` directly when the source already exposes -standard task columns such as `question` and `answer`; the framework derives -`prompt` from `question`. Transform rows only to match the task contract, add -real reference fields, create multimodal content, or attach per-example state -that the rollout actually uses. - -Do not copy config defaults into every task row. Use `max_turns`, `sandbox`, -`program`, and tool visibility fields in task records only when they genuinely -vary by example. - -## Tools And Sandboxes - -Sandboxed tools are normal tools. Put them in `vf.Toolset`, pass runtime -resources through bindings or state helpers, and keep task rows serializable. - -Avoid local protocol scaffolding for ordinary sandbox handles. If an example -needs a local framework protocol class, hidden state-key convention, or custom -tool wrapper just to call a sandboxed tool, fix the public API or docs instead. - -## Failure Behavior - -Fail fast. Do not add fallback parsing, fallback imports, compatibility aliases, -broad best-effort branches, or silent degraded behavior to make a local -environment more permissive. - -Mutable process globals are a smell. Runtime clients, sessions, sandboxes, -caches, and registries should be owned by config-bound objects, lifecycle -methods, state, or runtime-managed owners. - -## Review Checklist - -Before approving a v1 environment: - -1. The root loader is `load_environment(config: vf.EnvConfig)`. -2. Child loaders use one concrete config type each. -3. Environment-specific fields are not on `EnvConfig`. -4. No subclass overrides final constructors. -5. No bottom-of-file helper clutter exists for single-use logic. -6. Taskset/harness ownership is clear. -7. Config values are serializable. -8. System prompt handling is config-first. -9. Task rows do not duplicate config defaults or framework-managed IDs. -10. Dataset records are returned directly when no transformation is needed. -11. Toolsets and users use first-class v1 objects. -12. The base harness is used unless a reusable protocol requires a custom one. -13. Failure paths are strict and explicit. -14. The environment has been validated through install/load/eval, not just - imported. diff --git a/verifiers/v1/PRIME_RL_V1.md b/verifiers/v1/PRIME_RL_V1.md new file mode 100644 index 0000000000..c2f846f252 --- /dev/null +++ b/verifiers/v1/PRIME_RL_V1.md @@ -0,0 +1,325 @@ +# prime-rl v1 Integration Sketch + +This document sketches the smallest `prime-rl` changes needed to consume +`verifiers.v1` environments cleanly. + +## Goal + +`prime-rl` should treat v0 and v1 as two environment protocols behind one +trainer-facing adapter: + +- v0 produces `trajectory` rollout outputs through the existing + `verifiers.Environment` and ZMQ server path. +- v1 produces strict `State` outputs with `transcript`, token-level advantages, + and typed task/state records through `verifiers.v1.Env`. + +The split belongs at the environment adapter boundary. Training code should not +scatter `if trajectory else transcript` branches across tokenization, +interleaving, logging, or advantage handling. + +## Current prime-rl Touchpoints + +`src/prime_rl/orchestrator/envs.py` currently assumes the v0 contract: + +- top-level `vf.load_environment(...)` returns a v0 `vf.Environment`. +- `REQUIRED_STATE_COLUMNS = ["trajectory"]`. +- `run_rollout(...)` passes `vf.RolloutInput`, `client`, `model`, + `sampling_args`, `state_columns`, and `env_client`. +- group scoring is detected through `env.rubric`. +- `run_group(...)` asks the v0 env to run and score the group in one call. + +`src/prime_rl/orchestrator/dispatcher.py` is also v0-shaped: + +- each scheduling unit is a `run_rollout` task, except group-scored envs where + one task calls `run_group`; +- a single `ClientConfig` is pinned per group for prefix-cache locality; +- train rollouts use the rollout inference pool, while SFT train rollouts use + the teacher pool as the sampled model; +- completed rollout tasks are normalized to `FinishedRollout.raw`, which is + expected to be a `vf.RolloutOutput`. + +`src/prime_rl/orchestrator/train_sink.py` assumes trainer-side scalar +advantages: + +- `process_rollout(...)` tokenizes immediately from v0 trajectory records; +- `process_group(...)` calls `assign_advantages(...)`; +- each `TrainingSample` receives `sample.advantage = rollout.advantage`. + +`src/prime_rl/orchestrator/trajectories.py` also assumes the v0 rollout shape: + +- rollout output carries v0 trajectory records. +- each step is a `vf.TrajectoryStep`. +- token backfill reconstructs missing step tokens from `prompt` and + `completion`. +- interleaving consumes per-step token masks, logprobs, routed experts, and + multimodal sidecars. + +## Minimal Adapter Shape + +Introduce one protocol-specific adapter object in `prime-rl`: + +```python +class EnvAdapter(Protocol): + name: str + sampling_args: dict + + @property + def requires_group_rollouts(self) -> bool: ... + + def get_dataset(self, seed: int | None = None) -> Any: ... + + async def run_rollout( + self, + *, + client: vf.ClientConfig, + model: str, + example: dict, + cache_salt: str | None, + teacher: vf1.ModelConfig | None = None, + ) -> RolloutView: ... + + async def score_group( + self, + *, + views: list[RolloutView], + ) -> list[RolloutView]: ... +``` + +`RolloutView` is a trainer-local view, not a Verifiers public type. It replaces +direct reads from `vf.RolloutOutput` inside prime-rl: + +```python +class RolloutView(Protocol): + example_id: str | int | None + error: object | None + stop_condition: str | None + reward: float + + def iter_turns(self) -> Iterable[TurnView]: ... + def has_env_token_advantages(self) -> bool: ... + def raw_for_storage(self) -> dict[str, object]: ... +``` + +The v0 implementation wraps trajectory records. The v1 implementation wraps +the live `State` and exposes `state.transcript`. + +## v1 Environment Adapter + +The v1 adapter imports `verifiers.v1 as vf1` and keeps all v1-specific logic in +one file: + +```python +class V1EnvAdapter: + def __init__(self, config: EnvConfig): + self.env: vf1.Env = vf1.load_environment(config.stripped_id, **config.args) + self.sampling_args = config.sampling.to_sampling_args() + + @property + def requires_group_rollouts(self) -> bool: + return self.env.requires_group_rollouts + + async def run_rollout( + self, + *, + client: vf.ClientConfig, + model: str, + example: dict, + cache_salt: str | None, + teacher: vf1.ModelConfig | None = None, + ) -> RolloutView: + task = self.env.taskset.to_task(example) + model_config = vf1.ModelConfig( + model=model, + client=vf1.ClientConfig.model_validate(client.model_dump()), + sampling_args=self._sampling_args_with_salt(cache_salt), + ) + state = await self.env.run_rollout( + task, + model=model_config, + teacher=teacher, + ) + return V1RolloutView(task=task, state=state) + + async def score_group(self, *, views: list[RolloutView]) -> list[RolloutView]: + v1_views = cast(list[V1RolloutView], views) + await self.env.score_group( + tasks=[view.task for view in v1_views], + states=[view.state for view in v1_views], + ) + return v1_views +``` + +The rule is simple: convert config data at the adapter boundary, never put live +clients into `State`. + +## Dataset And Task Rows + +v1 `Taskset` owns row-to-task construction. `prime-rl` should keep train/eval +sampling as row dictionaries and ask the adapter to realize each row: + +```python +task = env.taskset.to_task(example) +``` + +That lets dynamic task construction remain taskset-owned. If an environment +needs task-local setup from a server, that setup should happen inside the v1 +rollout lifecycle and write serializable data to `state.extras`, not into the +trainer row. + +## Transcript Tokenization + +Move the current trajectory functions behind turn views: + +```python +class TurnView(Protocol): + prompt: vf1.Messages + completion: vf1.Messages + tokens: vf1.TurnTokens | None + reward: float | None + prompt_advantages: list[float] | None + completion_advantages: list[float] | None +``` + +Then the existing renderer/tokenizer logic can be shared: + +- v0 adapter maps `TrajectoryStep.prompt` and `TrajectoryStep.completion`. +- v1 adapter maps `Turn.prompt` and `Turn.completion`. +- token backfill stays a trainer concern when the env did not return tokens. +- v1 token advantages are copied into `TrainingSample` during interleaving, not + carried through a rollout-level scalar. + +This keeps `transcript` canonical for v1. The v1 adapter should not emit a +derived trajectory field just to satisfy `prime-rl`. + +## Advantages + +v1 environments default to `advantage="rl"` and may provide token-level +advantages directly: + +- `Turn.tokens.prompt_advantages` +- `Turn.tokens.completion_advantages` + +`prime-rl` should treat those as authoritative. Trainer-side advantage +computation should run only for rollout groups that do not already contain +environment-provided token advantages. + +The minimal precedence rule is: + +1. token advantages on turns win; +2. otherwise compute advantages in `prime-rl`. + +The v1 adapter should expose this as one method on `RolloutView`, so the trainer +does not inspect Verifiers internals: + +```python +if all(view.has_env_token_advantages() for view in group): + interleave_rollout_view(view, advantage_source="env") +else: + assign_advantages(group, advantage_fn) +``` + +Mixed groups should fail fast. A group where some rollouts have env token +advantages and others do not is ambiguous for normalization. + +For v1, `assign_advantages(...)` should become a no-op when env token +advantages are present. For v0, it keeps the current scalar rollout advantage +behavior. + +## Teacher Flow + +Model and teacher should be symmetric config data: + +```python +student = vf1.ModelConfig( + model=model_name, + client=client_config, + sampling_args=student_sampling, +) +teacher = vf1.ModelConfig( + model=teacher_model_name, + client=teacher_client_config, + sampling_args=teacher_sampling, +) +``` + +`prime-rl` should pass the teacher config into the v1 adapter when a run +requests teacher behavior. The adapter passes it to `env.run_rollout(...)`. + +No live teacher client should be stored in `State`. The v1 harness resolves the +client from `ModelConfig` inside the live context. + +When an environment-level advantage needs tokenizer or logprob alignment, the +advantage function should request `model: vf1.ModelClient` or +`teacher: vf1.ModelClient | None` and call `model.get_renderer()` explicitly. +`State` stores only `state.model` / `state.teacher` configs, never renderer or +client handles. + +## Group Scoring + +v1 does not need `run_group`. `prime-rl` can keep its current scheduling model: + +1. launch `N` rollouts for the same example; +2. collect `N` `V1RolloutView`s; +3. call `adapter.score_group(views=group)`; +4. interleave/tokenize the scored views. + +This matches v1's `run_rollout` plus `score_group` contract and keeps group +advantage overrides in the environment. + +The main tradeoff is runtime lifetime. Today v1 closes rollout runtimes before +`score_group`. If a group reward needs live sandboxes, v1 will need an internal +group lifecycle owner before `prime-rl` can support that case. + +## Env Server Process + +The first v1 path should run in-process inside the `prime-rl` env worker rather +than forcing v1 through the v0 ZMQ `RolloutInput` API. The existing v0 path can +keep `ZMQEnvServer` unchanged. + +This keeps the v1 state contract intact and avoids inventing a trajectory +compatibility layer. + +If v1 needs a remote env server later, add a v1-specific RPC surface that sends: + +- serialized task rows or tasks, +- serialized `ModelConfig` for student and optional teacher, +- serialized `State` outputs with `transcript`. + +Do not route v1 through v0 `vf.RolloutInput`. + +## Minimal Diff Plan + +1. Split `orchestrator/envs.py` into a v0 adapter preserving the current ZMQ + implementation and a v1 adapter using in-process `vf1.Env`. +2. Change the dispatcher to reserve group scheduling for adapters whose + `requires_group_rollouts` property is true. For v0, keep the existing + environment-owned group call; for v1, dispatch individual `run_rollout` + tasks and let the sink call `score_group(...)` once the group is complete. +3. Replace `REQUIRED_STATE_COLUMNS = ["trajectory"]` with adapter-owned output + requirements. v0 keeps `"trajectory"`; v1 keeps `"transcript"`. +4. Introduce `RolloutView` / `TurnView` in prime-rl and move + `backfill_rollout_tokens`, `interleave_rollout`, filters, length penalties, + and eval metrics onto that view. +5. Keep existing v0 scalar advantage assignment. For v1, skip scalar assignment + when every rollout in the group has token advantages; otherwise run the + trainer-side advantage function and fan out scalar values over completion + tokens. +6. Add teacher config plumbing next to the existing student model/client + config. In SFT, train rollouts can keep using the teacher pool as the sampled + model; v1 additionally passes `teacher` when an env advantage or reward needs + the student/teacher pair. +7. Store v1 outputs with `transcript`, not `trajectory`. Save/log code can use + `RolloutView.raw_for_storage()` to keep v0/v1 output shapes honest. +8. Add tests with one v0 env and three v1 envs: simple rollout, group scoring, + and env-provided token advantages. + +## Open Design Tensions + +- Runtime-backed group rewards need a v1 grouped lifetime owner. `prime-rl` + should not solve that by keeping private runtime handles. +- v1 dynamic task setup should remain taskset/harness-owned. A future task + buffer can expose richer task production, but the first adapter should not add + a task-server concept. +- Trainer-local views reduce branching, but they are still another small type + surface in `prime-rl`. The alternative is direct `trajectory`/`transcript` + branching in tokenizer code, which will be harder to keep correct. diff --git a/verifiers/v1/README.md b/verifiers/v1/README.md index 8f320f3937..5d731cb6ef 100644 --- a/verifiers/v1/README.md +++ b/verifiers/v1/README.md @@ -1,339 +1,190 @@ # Verifiers v1 -`verifiers.v1` is the Taskset/Harness API for reusable eval and training -environments. - -- `Taskset` defines what is being attempted. -- `Harness` defines how the model or agent attempts it. -- `Env` adapts one taskset/harness pair to the existing eval/training worker - API. - -Start with [`docs/byo-harness.md`](../../docs/byo-harness.md) when authoring an -environment. Use [`docs/reference.md`](../../docs/reference.md) for API lookup -and [`RE_MIGRATION.md`](RE_MIGRATION.md) for migration notes. - -## Mental Model - -![Task to Harness to State](../../docs/assets/v1-task-harness-state.svg) - -v1 is data-first: - -- `Task` is immutable, serializable input data. -- `State` is mutable, serializable rollout output. -- Runtime handles such as clients, sandboxes, MCP sessions, and tool backends - are process-local and reached through state helpers while a rollout is active. -- Tasksets and harnesses are configured through strict Pydantic config objects. - -## Golden Loader Shape - -Environment packages expose typed child loaders and one tiny root loader: +v1 is an active-development rewrite of the Taskset/Harness stack. Breaking +changes are expected before release. Import it as: ```python -import verifiers as vf - - -class ReverseTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPrompt = "Reverse text exactly." - - -class ReverseTaskset(vf.Taskset[ReverseTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "eval": - return [] - return [ - { - "prompt": [{"role": "user", "content": "Reverse abc."}], - "answer": "cba", - "max_turns": 1, - } - ] - - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - return float(response.strip() == task["answer"]) - - -def load_taskset(config: ReverseTasksetConfig) -> ReverseTaskset: - return ReverseTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) +import verifiers.v1 as vf ``` -Add `load_harness(config: MyHarnessConfig)` only when the package owns reusable -execution behavior: +The top-level `verifiers` package is the v0 surface. v1 code should not import +v1 classes from top-level `verifiers`, and v0 code should not rely on +`verifiers.v1` internals. + +## Model + +- `Taskset` owns tasks, task prompts, task tools, user simulation, metrics, + rewards, and task-specific lifecycle. +- `Env` owns the selected group advantage function. The default is `"rl"`; + pass `advantage=None` to disable environment-provided token advantages. +- `EnvRun` owns one environment execution: env-scope toolsets/users, + per-rollout runtime creation, and grouped rollout coordination. Eval creates + one `EnvRun` for the evaluation. Direct `Env.run_rollout(...)` is a one-shot + convenience around `EnvRun`. +- `Group` owns the tasks and states for one grouped example and calls + `env.score_group(...)` after its member rollouts finish. +- `Harness` is the agent. Its `run(...)` method starts a standalone `EnvRun` + when no parent `Context` is supplied; nested calls reuse the parent + `Context`. Direct `Harness.run(...)` defaults to `score=False`; + `Env.run_rollout(...)` opts into rollout scoring. +- `Context` is the live per-harness execution record: task, state, runtime, + model/teacher clients, toolsets, user, parent context, and scoring flags. +- `Env` is the thin adapter that pairs taskset and harness, opens `EnvRun` + contexts, scores groups, and serializes output. +- `State` is the canonical rollout record. It is a strict Pydantic model with + `transcript: list[Turn]`; there is no live `trajectory` alias. +- `state.messages` is a convenience rendering of the latest conversation + prompt plus completion. `state.transcript` remains the canonical record for + per-request history. +- `state.extras` is the user-owned mutable rollout data surface. Taskset and + harness configs may provide typed `vf.Extras` defaults; v1 realizes one schema + from both and rejects duplicate keys. + +Runtime handles, model clients, MCP sessions, and server connections are never +stored in `Task` or `State`. + +## Runtime And Protocols + +Runtime providers expose one live `Runtime` contract: `start`, `stop`, `expose`, +`run`, `read`, and `write`. The built-in configs are `subprocess`, `docker`, +and `prime`; `modal` and `daytona` are reserved provider stubs. +Task rows may set `image` and `resources` for serializable per-task runtime +selection. Runtime config wins over task resources when the config field is set +away from its provider default. Live runtimes stay owned by the harness +lifecycle. + +Harnesses that run external agents start an `InterceptionServer` and expose it +through the active runtime. Built-in endpoint protocols cover OpenAI +chat completions, OpenAI completions, OpenAI responses, and Anthropic messages. +Custom protocols are harness-side adapters: override `Harness.load_protocols()` +and return `EndpointProtocol` objects with `routes`, `env(...)`, `parse(...)`, +and `serialize(...)`. Protocols may execute Python on the loaded harness side, +but they must exchange only JSON request/response data and must not write live +handles into `Task` or `State`. + +## Tools And Users + +Toolsets are declared with `ToolsetConfig` and implemented as `Toolset` +subclasses: ```python -class MyHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="my_env.agent:run") - - -class MyHarness(vf.Harness[MyHarnessConfig]): - pass - - -def load_harness(config: MyHarnessConfig) -> MyHarness: - return MyHarness(config=config) -``` - -Do not subclass `EnvConfig` to narrow child config types. The child loader -annotations define `[env.taskset]` and `[env.harness]`. - -Start with a taskset and the base harness. Add a custom harness only when the -environment owns a reusable execution protocol, such as a command agent, -third-party framework adapter, endpoint interceptor, primary sandbox placement, -or program runner. - -## Ownership - -| Object | Owns | -| --- | --- | -| `Taskset` | Task data, task loading, task prompts, task controls, task tools, users, metrics, rewards, and task-specific lifecycle. | -| `Harness` | Rollout execution, programs, model/client defaults, endpoint interception, primary sandbox placement, command/framework adapters, and execution artifacts. | -| `Env` | Worker adapter for one taskset/harness pair. | - -Tasksets own the domain. Harnesses own execution. If a tool defines the task's -action space or success condition, put it on the taskset. If code describes how -an arbitrary task is attempted, put it on the harness. - -## Core Contracts - -### Task - -`Task` is immutable and serializable. `task["prompt"]` must not contain system -messages. Use top-level fields for task controls: - -| Field | Meaning | -| --- | --- | -| `prompt` | User/developer/tool messages. | -| `system_prompt` | Per-task taskset-side system prompt override. | -| `answer` | Reference answer or target data. | -| `info` | Serializable metadata. | -| `max_turns` | Per-task base-loop limit. | -| `toolsets` / `tools` | Visibility controls for toolsets and tools. | -| `sandbox` | Per-task sandbox override. | -| `program` | Task-owned program files, dirs, setup, env, artifacts, bindings, and args. | - -Use `max_turns`, `sandbox`, `program`, and visibility fields in tasks only when -they genuinely vary by example. Do not copy config defaults or -framework-managed IDs into task rows. - -### State - -`State` is mutable during rollout and serializable before return. It stores -trajectory, completion, metrics, reward, timing, artifacts, errors, and any -environment output. - -Use state helpers for active runtime resources: - -- `state.get_model()` -- `state.get_client(...)` -- `state.get_endpoint_config(...)` -- `state.get_max_turns(default)` -- `state.get_tools()` -- `state.add_tool("toolset_name", tool)` - -### Config - -Config values must be serializable. Use import refs for callables in TOML or -package config. Put task fields on `TasksetConfig`; put execution fields on -`HarnessConfig`. - -Important owner config fields: - -- `system_prompt` -- `user` -- `toolsets` -- `objects` -- `bindings` -- `artifacts` -- lifecycle lists such as `setups`, `updates`, `metrics`, `rewards`, and - `cleanups` -- `scoring` - -`Taskset.__init__`, `Harness.__init__`, and `User.__init__` are final. -Customize through config, public load methods, lifecycle decorators, and -program config. - -## System Prompts - -System prompts resolve per task during `Harness.setup_state(...)`. - -- `T` is the resolved taskset side: `task["system_prompt"]` when present, - otherwise `TasksetConfig.system_prompt`. -- `H` is the harness side: `HarnessConfig.system_prompt`. - -`HarnessConfig.system_prompt_strategy` chooses the result: +class SearchToolsetConfig(vf.ToolsetConfig): + scope: vf.Scope = "rollout" -| Strategy | Meaning | -| --- | --- | -| `HT` | Harness side followed by resolved taskset side. Default. | -| `TH` | Resolved taskset side followed by harness side. | -| `H_OR_T` | Harness side when present, otherwise resolved taskset side. | -| `T_OR_H` | Resolved taskset side when present, otherwise harness side. | -| `H` | Harness side only. | -| `T` | Resolved taskset side only. | -| `REJECT` | Error if both sides are present. | -Use `vf.SystemPromptConfig(path="system_prompt.txt")` for file-backed prompts. -Override `load_system_prompt(config)` only when prompt construction is computed. - -## Tasksets - -Tasksets load train and eval data through `load_tasks(split=...)`: - -```python -class MyTaskset(vf.Taskset[MyTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: +class SearchToolset(vf.Toolset): + @vf.tool( + args={"query_context": "state.extras.query_context"}, + extends={"events": "state.extras.search_events"}, + ) + def search(self, query: str, query_context: str) -> dict: ... -``` - -`vf.Tasks` can be a `datasets.Dataset`, an iterable of serializable records, or -an iterable of `vf.Task` objects. `Taskset.get_dataset()` calls -`load_tasks(split="train")`; `Taskset.get_eval_dataset()` calls -`load_tasks(split="eval")`. - -Prefer returning a `datasets.Dataset` directly when source columns already -match the task contract, such as `question` and `answer`. Hardcode fixed -upstream split names inside `load_tasks(split=...)`. Only expose -split-name config when the upstream split choice is genuine user-space -configuration, not the way v1 decides whether eval exists. Return `[]` for -`split == "eval"` when the taskset has no explicit eval source; `vf.Env` treats the empty -split as an absent eval dataset so the base environment can fall back to train -data with its standard warning. - -Use tasksets for: -- dataset loading; -- task-owned tools; -- user simulators; -- task-specific setup/update/cleanup; -- metrics, rewards, advantages, and stop conditions. -## Harnesses And Programs - -Harnesses run tasks. The base harness is endpoint-backed and supports the -default tool loop. - -`HarnessConfig.program` controls executable behavior: - -| Form | Meaning | -| --- | --- | -| `vf.ProgramConfig()` | Base endpoint-backed tool loop. | -| `vf.ProgramConfig(base=True)` | Explicit base loop. | -| `vf.ProgramConfig(fn="pkg:run")` | Importable Python program. | -| `vf.ProgramConfig(command=["agent", "run"])` | Local or sandboxed command. | - -Preferred program signature: - -```python -async def program(task: vf.Task, state: vf.State) -> vf.State: - ... +class SearchTasksetConfig(vf.TasksetConfig): + toolsets: vf.ToolsetConfigs = {"wiki": SearchToolsetConfig()} ``` -Use custom harnesses for reusable command agents, third-party framework -adapters, endpoint routing, primary sandbox placement, or execution artifacts. -Use `vf.load_harness(config=config.harness)` otherwise. - -## Tools, Users, And Lifecycle - -Toolsets package model-visible schemas plus bindings, objects, artifacts, and -lifecycle hooks: - -```python -class SearchTaskset(vf.Taskset[SearchTasksetConfig]): - def load_toolsets(self, config: SearchTasksetConfig) -> vf.Toolsets: - return {"search": vf.Toolset(tools=[search])} +The `toolsets` key is the model-visible tool prefix. Config may override a +taskset-defined toolset by key without repeating its source, and may add a new +toolset by pointing `source` at a `ToolsetConfig` class. + +One `Toolset` may expose multiple tools. Supported scopes are: + +- `rollout`: started for one rollout and cleaned up afterward. +- `env`: started once for an `EnvRun` and reused by all rollouts in that run. + +Supported placements are: + +- `dedicated`: start the toolset/user in its own runtime. +- `colocated`: start the toolset/user in the owning rollout runtime. +- `remote`: connect to an existing URL. + +`@vf.tool(args=..., sets=..., extends=...)` is the only framework wiring path +for hidden args and state writes. Bound args are hidden from the model and +injected from serialized `task.*`, `state.*`, `extras.*`, and server-local +`resources.*` paths. `sets` replaces one `state.*` or `extras.*` path; +`extends` appends a returned list to one `state.*` or `extras.*` list path. +Multiple same-path extends in one tool-call batch are allowed, with no ordering +guarantee. + +Users use the sibling `UserConfig` / `User` path over the same server base. A +user exposes a hidden `respond` tool and returns `messages`. Toolsets use the +same response shape; the default harness converts single text tool responses +into protocol `tool` messages and appends explicit multi-message responses +after tool results. Hidden tools are callable only by the harness through the +hidden-call path, not by model-visible tool calls. + +## Authoring Pattern + +The default v1 package layout is component-first: + +```text +my_env/ + my_env/ + taskset.py + harness.py # optional + servers/ + search/ + config.py + toolset.py + user/ + config.py + user.py ``` -Tasks show all tools by default and can restrict visibility with `toolsets` and -`tools`. - -Users subclass `vf.User` and implement `get_response(...)`. Use users for -environment replies between model turns. Use tools for schema actions. Use -setup/update handlers for state changes that should not add messages. - -Lifecycle behavior belongs on the owner class: +`vf.load_environment("my-env")` imports the package, discovers `taskset.py` and +optional `harness.py`, and constructs `vf.Env` internally. ```python -class MyTaskset(vf.Taskset[MyTasksetConfig]): - @vf.update - async def extract_answer(self, task: vf.Task, state: vf.State) -> None: - ... +import verifiers.v1 as vf +from pydantic import BaseModel - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State) -> float: - ... -``` -## Runtime Composition +class MyTask(vf.Task): + answer: str -Advanced code can create child task states and borrow selected runtime handles: -```python -child_state = state.for_task(child_task, borrow="model", tools=["search"]) -child_state = await child_harness.run(child_task, child_state) -``` - -Borrowed resources remain owned by the source runtime and are stripped before -state serialization. +class MyDetails(BaseModel, extra="forbid"): + source: str -## TOML Shape -Eval and training config own run settings. v1 child config owns environment -behavior: +class MyTasksetConfig(vf.TasksetConfig): + system_prompt: vf.SystemPrompt = "Say exactly what is requested." -```toml -[[eval]] -env_id = "my-v1-env" -[eval.taskset] -system_prompt = "Answer exactly." - -[eval.harness] -max_turns = 4 -``` - -CLI overrides target typed child fields: +class MyTaskset(vf.Taskset[MyTasksetConfig]): + task_type = MyTask -```bash -prime eval run my-v1-env --taskset.system-prompt "Answer exactly." --harness.max-turns 4 -``` + def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: + return [{"prompt": [{"role": "user", "content": "Say ok."}], "answer": "ok"}] -## Packaged Implementations + @vf.reward + async def exact(self, task: MyTask, state: vf.State) -> float: + message = state.completion[-1] + return float(str(message.content).strip() == task.answer) -Reusable tasksets and harnesses live under top-level `packages/`. -```bash -uv add "verifiers[tasksets]" -uv add "verifiers[harnesses]" -uv add "verifiers[packages]" +def load_taskset(config: MyTasksetConfig) -> MyTaskset: + return MyTaskset(config=config) ``` -Tasksets include Harbor, OpenEnv, OpenReward, TextArena, and NeMoGym. Harnesses -include OpenCode, Pi, mini-swe-agent, Terminus, RLM, and NeMoGymHarness. - -They use the same loader shape as local implementations. - -TOML can also compose packages directly. In that case `[eval.taskset].id` -selects the taskset loader package and `[eval.harness].id` optionally selects -the harness loader package: - -```toml -[[eval]] - -[eval.taskset] -id = "tasksets.harbor" -tasks_dir = "tasks" - -[eval.harness] -id = "harnesses.opencode" -max_turns = 8 -``` +Config is serializable policy. Live Python functions are allowed as decorated +methods on loaded `Taskset`/`Harness` objects, not as config values, task fields, +state fields, tool definitions, or runtime specs. + +Use ordinary Pydantic models for strict nested task/config records. The v1 +library keeps its own types to framework contracts; example-specific nesting is +userspace schema. + +## Current Tensions + +- Group rewards and token-level advantages are first-class, and v1 envs default + to the built-in `"rl"` advantage. Group scoring + currently runs after per-rollout runtimes close. Supporting runtime-backed + group scoring would require an explicit group runtime lifetime. +- Env-scope toolsets are first-class. Group-specific resources should use + `state.group_id` plus env-scope toolset state rather than a third tool scope. +- The base harness is model-loop native. Command/program agents should be + implemented as `Harness` subclasses that use `Runtime`, not as generic + callable config. diff --git a/verifiers/v1/RE_MIGRATION.md b/verifiers/v1/RE_MIGRATION.md deleted file mode 100644 index d81ae69884..0000000000 --- a/verifiers/v1/RE_MIGRATION.md +++ /dev/null @@ -1,417 +0,0 @@ -# Research Environments v1 Migration - -This guide maps older research-environments packages onto the current v1 -Taskset/Harness shape. The authoritative implementation guide is -[`docs/byo-harness.md`](../../docs/byo-harness.md); this file is a migration map -for choosing the right v1 pattern. - -## Migration Contract - -Every migrated v1 package should expose the same loader boundary: - -```python -import verifiers as vf - - -class MyTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPrompt = "Answer exactly." - - -class MyTaskset(vf.Taskset[MyTasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - if split == "eval": - return [] - return [ - { - "prompt": [{"role": "user", "content": "Question?"}], - "answer": "Answer", - "max_turns": 1, - } - ] - - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - return float(response.strip() == task["answer"]) - - -def load_taskset(config: MyTasksetConfig) -> MyTaskset: - return MyTaskset(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -Add `load_harness(config: MyHarnessConfig)` only when the package owns a -reusable execution mechanism such as a command agent, framework adapter, -sandbox placement policy, or endpoint interception behavior: - -```python -class MyHarnessConfig(vf.HarnessConfig): - program: vf.ProgramConfig = vf.ProgramConfig(fn="my_env.agent:run") - - -class MyHarness(vf.Harness[MyHarnessConfig]): - pass - - -def load_harness(config: MyHarnessConfig) -> MyHarness: - return MyHarness(config=config) - - -def load_environment(config: vf.EnvConfig) -> vf.Env: - """Loader pattern for all Taskset/Harness environments.""" - return vf.Env( - taskset=vf.load_taskset(config=config.taskset), - harness=vf.load_harness(config=config.harness), - ) -``` - -Do not subclass `vf.EnvConfig` to narrow child config types. The -`load_taskset` and `load_harness` annotations are the child config contract for -Python, TOML, CLI, eval, GEPA, RL, and Hosted Training. - -## Pattern Map - -Pick the row that matches the old package, copy the referenced v1 shape, and -then port the dataset and scoring logic. - -| Old package shape | Reference | v1 pattern | -| --- | --- | --- | -| AIME, GPQA, MATH500, MMLU-Pro, SimpleQA | `environments/reverse_text/reverse_text_v1.py` | serializable task rows, base `vf.Harness`, taskset reward | -| instruction-following or string-matching tasks | `environments/reverse_text/reverse_text_v1.py` | prompt taskset plus class-local metrics/rewards | -| Python execution math/code tasks | `environments/math_python/math_python_v1.py` | sandbox-backed callable tool | -| web/search/browser tasks | `environments/wiki_search/wiki_search_v1.py` | taskset-owned `vf.Toolset` with bound private objects | -| BFCL-style task-provided schemas | `environments/bfcl_v3/bfcl_v3.py` | task-local tool visibility/schemas and state-recorded calls | -| games and simulators | `environments/tau2_bench_v1/tau2_bench_v1.py` | taskset-owned `vf.User` subclass | -| stdio MCP tasks | `environments/mcp_search_env/mcp_search_env.py` | `vf.MCPTool` entries in a `vf.Toolset` | -| helper agents or judges | `environments/hello_subagent_v1/hello_subagent_v1.py` | nested `vf.Harness.run(...)` from a tool/update/reward | -| shared sandbox helper agents | `environments/hello_parallel_sandbox_v1/hello_parallel_sandbox_v1.py` | borrowed runtime state through `state.for_task(...)` | -| Harbor task directories | `environments/opencode_harbor/opencode_harbor.py` | `HarborTaskset` plus a packaged command harness | -| OpenCode, Pi, mini-swe-agent, Terminus, RLM | `packages/harnesses/` | reusable `vf.Harness` subclasses with typed config | -| OpenEnv, OpenReward, TextArena, NeMoGym | `packages/tasksets/` | reusable `vf.Taskset` subclasses matching upstream formats | - -## Task Data - -Tasksets return serializable task records from `load_tasks(split=...)`. During -a rollout, the framework materializes them as immutable `vf.Task` objects. - -```python -yield { - "prompt": [{"role": "user", "content": question}], - "answer": answer, - "info": {"source_id": source_id}, - "max_turns": 8, -} -``` - -Use top-level fields for task controls: - -- `prompt`: user/developer/tool messages, never system messages. -- `system_prompt`: per-task system instructions. -- `answer`: reference answer or target data. -- `info`: serializable metadata. -- `max_turns`: per-task base-loop turn limit. -- `toolsets`: toolset visibility with `{"show": [...]}` or `{"hide": [...]}`. -- `tools`: per-toolset tool visibility with `{"search": {"show": [...]}}`. -- `sandbox`: task-owned sandbox override. -- `program`: task-owned program files, dirs, setup, env, artifacts, bindings, - and args. -- `artifacts`: task-owned artifacts collected after program execution. - -Do not ask users to manage task IDs. Preserve upstream IDs only when they are -meaningful task metadata. - -## System Prompts - -Static system prompts belong in the config that owns the policy. Use taskset -config for task policy and harness config for execution or agent policy: - -```python -class PromptTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPrompt = "Answer concisely." -``` - -File-backed GEPA prompts should also be config: - -```python -class PromptTasksetConfig(vf.TasksetConfig): - system_prompt: vf.SystemPromptConfig = vf.SystemPromptConfig( - path="system_prompt.txt" - ) -``` - -Override `load_system_prompt(config)` only when prompt construction is computed -from config fields or package resources. Do not put system messages in -`task["prompt"]`. System prompt resolution is per task: task prompt overrides -taskset prompt for the taskset side, then `HarnessConfig.system_prompt_strategy` -resolves the taskset side against the harness side. The available strategies are -`HT`, `TH`, `H_OR_T`, `T_OR_H`, `H`, `T`, and `REJECT`; the default is `HT`. - -## Single-Turn QA And Instruction Following - -Use the base harness unless the old environment owns a reusable execution -mechanism. Move dataset construction into `load_tasks(split=...)`, and move -each reward/metric onto the taskset class. - -```python -class QATasksetConfig(vf.TasksetConfig): - dataset_name: str = "gsm8k" - - -class QATaskset(vf.Taskset[QATasksetConfig]): - def load_tasks(self, split: vf.TaskSplit = "train") -> vf.Tasks: - dataset_split = "test" if split == "eval" else "train" - return load_dataset(self.config.dataset_name, "main", split=dataset_split) - - @vf.reward(weight=1.0) - async def exact(self, task: vf.Task, state: vf.State) -> float: - messages = vf.get_messages(state.get("completion") or [], role="assistant") - response = str(messages[-1].content or "") if messages else "" - return float(str(task["answer"]).strip() in response) -``` - -Return the source dataset directly when it already has standard fields such as -`question` and `answer`; v1 derives `prompt` from `question`. Transform records -only when the source does not match the task contract. - -Judge/extractor dependencies should be private objects with explicit bindings: - -```python -class ExtractTasksetConfig(vf.TasksetConfig): - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"extract_answer": "my_env.extractors:load_extractor"} - ) - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - {"exact.extract_answer": "objects.extract_answer"} - ) -``` - -The reward reads serializable `task`/`state` and receives the bound dependency -as an argument: - -```python -@vf.reward(weight=1.0) -async def exact(self, task: vf.Task, state: vf.State, extract_answer) -> float: - return float(extract_answer(state.get("completion") or []) == task["answer"]) -``` - -## Callable Tools - -Tools that define the task action space belong to the taskset. Expose them -through `vf.Toolset`; hide private dependencies behind `objects` and `bindings`. - -```python -async def search(query: str, exa) -> str: - return await exa.search(query) - - -async def open_page(url: str, exa) -> str: - return await exa.open(url) - - -class SearchTasksetConfig(vf.TasksetConfig): - objects: vf.ObjectsConfig = vf.ObjectsConfig.model_validate( - {"exa": "my_env.search:load_exa"} - ) - bindings: vf.BindingsConfig = vf.BindingsConfig.model_validate( - { - "search.search.exa": "objects.exa", - "search.open_page.exa": "objects.exa", - } - ) - - -class SearchTaskset(vf.Taskset[SearchTasksetConfig]): - def load_toolsets(self, config: SearchTasksetConfig) -> vf.Toolsets: - return {"search": vf.Toolset(tools=[search, open_page])} -``` - -Task rows show all toolsets/tools by default and can restrict visibility: - -```python -yield { - "prompt": [{"role": "user", "content": "Search only the docs."}], - "toolsets": {"show": ["search"]}, - "tools": {"search": {"show": ["search"]}}, -} -``` - -Use rollout-scoped toolsets for session-backed tools. Keep live backend handles -on `state`; keep task records serializable. - -## Users - -Use a `vf.User` subclass when the environment replies with user messages after -model turns. Users are not callables. - -```python -class GameUserConfig(vf.UserConfig): - pass - - -class GameUser(vf.User[GameUserConfig]): - async def get_response( - self, - task: vf.Task, - state: vf.State, - messages: list[vf.Message], - ) -> list[vf.UserMessage]: - observation = state["game"].observe(messages) - return [{"role": "user", "content": observation}] - - -class GameTasksetConfig(vf.TasksetConfig): - user: GameUserConfig = GameUserConfig() -``` - -Use a user for simulator observations. Use tools when the model must select an -explicit schema action. Use setup/update handlers when state should change -without adding conversation messages. - -## MCP Toolsets - -MCP servers are tool entries: - -```python -class FetchTasksetConfig(vf.TasksetConfig): - toolsets: dict[str, vf.ToolsetConfig] = { - "fetch": vf.ToolsetConfig( - tools=[ - vf.MCPToolConfig(command="uvx", args=["mcp-server-fetch"]), - ], - scope="rollout", - ) - } -``` - -The runtime materializes MCP tools as callable handles for Python programs and -can expose resolved toolsets through the generic `mcp` program channel for -command harnesses. `program.channels` names the program-facing channel, not a -specific tool. - -## Programs, Harnesses, And Sandboxes - -Use `HarnessConfig.program` for the executable behavior: - -```python -class AgentHarnessConfig(vf.HarnessConfig): - sandbox: vf.SandboxConfig = vf.SandboxConfig(image="python:3.11-slim") - program: vf.ProgramConfig = vf.ProgramConfig( - command=["agent", "run"], - sandbox=True, - channels={"mcp": {"setup": ["agent mcp add vf ${VF_MCP_URL}"]}}, - ) -``` - -Tasksets can contribute task-owned files, dirs, setup, env, artifacts, -bindings, and command args through `task["program"]`. The harness still owns -the program kind (`base`, `fn`, or `command`), channel wiring, and primary -sandbox placement. - -```python -yield { - "prompt": [{"role": "user", "content": instruction}], - "sandbox": {"image": "python:3.12-slim"}, - "program": { - "files": { - "/task/instruction.md": {"task": "instruction"}, - }, - "env": {"TASK_ID": {"task": "info.source_id"}}, - "artifacts": { - "agent_log": { - "path": "/workspace/agent.log", - "format": "text", - "optional": True, - } - }, - }, -} -``` - -Use task sandbox overrides only when the taskset owns per-task images, files, -resource sizing, or setup. Put reusable execution policy on the harness. - -## Nested Harnesses - -Nested harnesses are regular harness runs. A tool, update, reward, or program -can create a child task and run a child harness. Bind the child harness through -the owning object config or construct it inside the owning Python object; do not -store live harness objects in task data. - -```python -async def ask_child(name: str, child_harness: vf.Harness, state: vf.State) -> str: - child_task = vf.Task( - {"prompt": [{"role": "user", "content": f"Say hello to {name}."}]} - ).freeze() - child_state = await child_harness.run(child_task) - messages = vf.get_messages(child_state.get("completion") or [], role="assistant") - return str(messages[-1].content or "") if messages else "" -``` - -Borrow runtime handles only when the child intentionally reuses live parent -resources: - -```python -child_state = state.for_task(child_task, borrow="model", tools=["search"]) -``` - -Borrowed resources are process-local runtime handles and are stripped before -serialization. - -## Packaged Tasksets And Harnesses - -Prefer the sibling packages when an upstream format already matches: - -```bash -uv add "verifiers[tasksets]" -uv add "verifiers[harnesses]" -uv add "verifiers[openenv]" -uv add "verifiers[openreward]" -uv add "verifiers[ta]" -uv add "verifiers[nemogym]" -``` - -```python -from harnesses import OpenCode, OpenCodeConfig -from tasksets import HarborTaskset, HarborTasksetConfig - - -def load_taskset(config: HarborTasksetConfig) -> HarborTaskset: - return HarborTaskset(config=config) - - -def load_harness(config: OpenCodeConfig) -> OpenCode: - return OpenCode(config=config) -``` - -`HarborTaskset` owns Harbor task loading, task sandbox overrides, task uploads, -and test scoring. Command harnesses own installation, endpoint wiring, config -generation, channel setup, and log artifacts. The two sides communicate only -through task controls, program config, sandbox config, state, and lifecycle -handlers. - -## Migration Checklist - -1. The root loader is `load_environment(config: vf.EnvConfig)`. -2. Custom tasksets have `load_taskset(config: MyTasksetConfig)`. -3. Custom harnesses have `load_harness(config: MyHarnessConfig)`. -4. `Taskset`, `Harness`, and `User` subclasses do not override `__init__`. -5. Static prompts are config fields; computed prompts use `load_system_prompt`. -6. Task data is serializable and does not contain live handles. -7. Runtime handles live on `state` or framework runtime owners. -8. Tasksets own task data, tools, users, metrics, rewards, and task behavior. -9. Harnesses own programs, endpoint routing, sandboxes, command agents, and - execution artifacts. -10. Tools are exposed through `vf.Toolset`; tasks show/hide tools and toolsets. -11. Private dependencies use `ObjectsConfig` plus `BindingsConfig`. -12. Lifecycle logic is on the owning class with `@vf.*` decorators. -13. No one-off bottom-of-file helpers are needed for ordinary implementations. -14. The install/load/eval path has been validated with `prime eval run` or the - relevant package-install test. diff --git a/verifiers/v1/__init__.py b/verifiers/v1/__init__.py index 2ff5788bd0..0cbede5bc7 100644 --- a/verifiers/v1/__init__.py +++ b/verifiers/v1/__init__.py @@ -2,7 +2,7 @@ import importlib -from verifiers.decorators import ( +from .decorators import ( advantage, cleanup, metric, @@ -14,137 +14,205 @@ ) from verifiers.types import ( AssistantMessage, - EndpointConfig, + ClientConfig, Message, + MessageContent, Messages, SystemMessage, TextMessage, - ToolLike, + ToolCall, ToolMessage, UserMessage, ) -from verifiers.utils.message_utils import get_messages +from . import advantages +from .advantages import AdvantageConfig from .config import ( - CallableConfig, Config, - SignalConfig, ) from .env import Env, EnvConfig -from .artifact import ArtifactConfig, Artifacts, ArtifactsConfig from .harness import Harness, HarnessConfig -from .model import ModelConfig -from .program import ProgramConfig, ProgramValue -from .runtime import TrajectoryVisibility -from .sandbox import SandboxConfig -from .utils.scoring_utils import ( - add_metric, - add_reward, - add_advantage, - build_signals, - collect_signals, - score_group, - score_rollout, +from .interception import ( + EndpointProtocol, + InterceptedRequest, + InterceptionServer, + ProtocolRoute, ) -from .state import State -from .task import Task +from .lifecycle import EnvRun, Group +from .protocols import ( + AnthropicMessagesProtocol, + OpenAIChatCompletionsProtocol, + OpenAICompletionsProtocol, + OpenAIResponsesProtocol, + default_protocols, +) +from .mcp import MCPToolRegistry, ServerResponse +from .runtime import ( + CommandResult, + DaytonaRuntimeConfig, + DaytonaRuntimeProvider, + DaytonaRuntime, + DockerRuntimeConfig, + DockerRuntimeProvider, + DockerRuntime, + ModalRuntimeConfig, + ModalRuntimeProvider, + ModalRuntime, + PrimeRuntimeConfig, + PrimeRuntimeProvider, + PrimeRuntime, + RuntimeConfig, + RuntimeConfigValue, + RuntimeProvider, + Runtime, + SubprocessRuntimeConfig, + SubprocessRuntimeProvider, + SubprocessRuntime, + make_runtime_provider, + resolve_runtime_config, +) +from .state import ( + Extras, + State, + Timing, + TimeSpan, + Turn, + TurnTokens, + TurnUsage, +) +from .task import Resources, Task, TaskVisibility from .taskset import Taskset, TasksetConfig, discover_sibling_dir from .toolset import ( - MCPTool, - MCPToolConfig, + Scope, + ServerPlacement, + ServerConfig, Toolset, ToolsetConfig, - Toolsets, + ToolsetConfigs, VisibilityConfig, + resource, + tool, ) -from .utils.endpoint_utils import Endpoint -from .utils.binding_utils import BindingsConfig, ObjectsConfig from .utils.prompt_utils import SystemPrompt, SystemPromptConfig, SystemPromptStrategy from .types import ( - ConfigData, Handler, JsonData, - Objects, + JsonValue, + ModelClient, + ModelConfig, PromptInput, + Context, TaskSplit, Tasks, ) -from .user import User, UserConfig +from .user import User, UserConfig, user __all__ = [ - "BindingsConfig", - "ArtifactConfig", - "Artifacts", - "ArtifactsConfig", - "ConfigData", - "CallableConfig", "Config", "Env", "EnvConfig", - "Endpoint", - "EndpointConfig", + "EnvRun", + "Extras", + "EndpointProtocol", "AssistantMessage", + "ClientConfig", "Harness", "HarnessConfig", "Handler", + "Group", + "InterceptedRequest", + "InterceptionServer", "JsonData", - "MCPTool", - "MCPToolConfig", + "JsonValue", + "MCPToolRegistry", + "ServerResponse", "Message", + "MessageContent", "Messages", + "OpenAIChatCompletionsProtocol", + "OpenAICompletionsProtocol", + "OpenAIResponsesProtocol", + "AnthropicMessagesProtocol", + "AdvantageConfig", + "ProtocolRoute", + "Context", + "ModelClient", "ModelConfig", - "Objects", - "ObjectsConfig", - "ProgramConfig", - "ProgramValue", + "RuntimeConfig", + "RuntimeConfigValue", + "RuntimeProvider", + "Runtime", + "Resources", + "Scope", + "ServerPlacement", + "ServerConfig", + "CommandResult", + "DaytonaRuntimeConfig", + "DaytonaRuntimeProvider", + "DaytonaRuntime", + "DockerRuntimeConfig", + "DockerRuntimeProvider", + "DockerRuntime", + "ModalRuntimeConfig", + "ModalRuntimeProvider", + "ModalRuntime", + "PrimeRuntimeConfig", + "PrimeRuntimeProvider", + "PrimeRuntime", + "SubprocessRuntimeConfig", + "SubprocessRuntimeProvider", + "SubprocessRuntime", "PromptInput", - "SandboxConfig", - "SignalConfig", "State", "SystemPrompt", "SystemPromptConfig", "SystemPromptStrategy", "Task", + "TaskVisibility", "TaskSplit", "Tasks", "Taskset", "TasksetConfig", + "TimeSpan", + "Timing", "SystemMessage", "TextMessage", - "ToolLike", + "ToolCall", "Toolset", "ToolsetConfig", - "Toolsets", + "ToolsetConfigs", "ToolMessage", - "TrajectoryVisibility", - "User", + "Turn", + "TurnTokens", + "TurnUsage", "UserMessage", + "User", "UserConfig", "VisibilityConfig", - "add_metric", - "add_reward", - "add_advantage", + "advantages", "advantage", - "build_signals", "cleanup", - "collect_signals", "discover_sibling_dir", + "default_protocols", "metric", - "get_messages", + "make_runtime_provider", + "resolve_runtime_config", + "load_environment", "load_harness", "load_taskset", "reward", - "score_group", - "score_rollout", "setup", "stop", "teardown", + "resource", + "tool", + "user", "update", ] def __getattr__(name: str): - if name in ("load_harness", "load_taskset"): - module = importlib.import_module("verifiers.utils.env_utils") + if name in ("load_environment", "load_harness", "load_taskset"): + module = importlib.import_module("verifiers.v1.loaders") return getattr(module, name) raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/verifiers/v1/advantages.py b/verifiers/v1/advantages.py new file mode 100644 index 0000000000..a073b846c7 --- /dev/null +++ b/verifiers/v1/advantages.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +import math +from typing import TypeAlias, cast + +from .decorators import advantage +from .state import State +from .task import Task +from .types import Handler +from .utils.scoring_utils import SignalRecord, signal_from_function +from .utils.config_utils import import_config_ref + + +AdvantageConfig: TypeAlias = str | None + + +def resolve_config(value: AdvantageConfig) -> Handler | None: + if value is None: + return None + if not isinstance(value, str) or not value: + raise TypeError("Env advantage must be a non-empty function path.") + ref = value if ":" in value else f"verifiers.v1.advantages:{value}" + resolved = import_config_ref(ref) + if not callable(resolved): + raise TypeError(f"Env advantage {ref!r} must resolve to a callable.") + if not getattr(resolved, "advantage", False): + raise TypeError(f"Env advantage {ref!r} must be decorated with @vf.advantage.") + return cast(Handler, resolved) + + +def signal(fn: Handler) -> SignalRecord: + return signal_from_function(fn) + + +@advantage +def rl(tasks: list[Task], states: list[State]) -> None: + grpo(tasks, states) + + +@advantage +def grpo(tasks: list[Task], states: list[State]) -> None: + _ = tasks + if not states: + return + baseline = sum(state.reward for state in states) / len(states) + values = [float(state.reward - baseline) for state in states] + variance = sum(value * value for value in values) / len(values) + scale = math.sqrt(variance) + if scale == 0.0: + values = [0.0 for _ in values] + else: + values = [float(value / scale) for value in values] + for state, value in zip(states, values, strict=True): + for turn in state.transcript: + if turn.tokens is not None: + turn.tokens.prompt_advantages = [0.0 for _ in turn.tokens.prompt_ids] + turn.tokens.completion_advantages = [ + float(value) for _ in turn.tokens.completion_ids + ] + + +@advantage +def rloo(tasks: list[Task], states: list[State]) -> None: + _ = tasks + if not states: + return + if len(states) == 1: + for turn in states[0].transcript: + if turn.tokens is not None: + turn.tokens.prompt_advantages = [0.0 for _ in turn.tokens.prompt_ids] + turn.tokens.completion_advantages = [ + 0.0 for _ in turn.tokens.completion_ids + ] + return + total = sum(state.reward for state in states) + values = [ + float(state.reward - ((total - state.reward) / (len(states) - 1))) + for state in states + ] + for state, value in zip(states, values, strict=True): + for turn in state.transcript: + if turn.tokens is not None: + turn.tokens.prompt_advantages = [0.0 for _ in turn.tokens.prompt_ids] + turn.tokens.completion_advantages = [ + float(value) for _ in turn.tokens.completion_ids + ] + + +@advantage +def reinforce(tasks: list[Task], states: list[State]) -> None: + _ = tasks + for state in states: + value = float(state.reward) + for turn in state.transcript: + if turn.tokens is not None: + turn.tokens.prompt_advantages = [0.0 for _ in turn.tokens.prompt_ids] + turn.tokens.completion_advantages = [ + value for _ in turn.tokens.completion_ids + ] + + +@advantage +def sft(tasks: list[Task], states: list[State]) -> None: + _ = tasks + for state in states: + for turn in state.transcript: + if turn.tokens is not None: + turn.tokens.prompt_advantages = [1.0 for _ in turn.tokens.prompt_ids] + turn.tokens.completion_advantages = [ + 1.0 for _ in turn.tokens.completion_ids + ] diff --git a/verifiers/v1/artifact.py b/verifiers/v1/artifact.py deleted file mode 100644 index 23eca02f37..0000000000 --- a/verifiers/v1/artifact.py +++ /dev/null @@ -1,86 +0,0 @@ -import json -from typing import Literal, TypeAlias, cast - -from pydantic import ConfigDict, StrictBool, model_validator -from typing_extensions import Self - -from .config import Config, validate_serializable_value -from .types import ConfigData, JsonValue - -ArtifactFormat: TypeAlias = Literal["text", "json"] - - -class ArtifactConfig(Config): - path: str - format: ArtifactFormat = "text" - key: str | None = None - optional: StrictBool = False - - def data(self) -> ConfigData: - data = self.model_dump(exclude_none=True) - if self.optional is False: - data.pop("optional", None) - return cast(ConfigData, data) - - def parse(self, content: str) -> JsonValue: - if self.format == "json": - value = cast(JsonValue, json.loads(content)) - elif self.format == "text": - value = content - else: - raise AssertionError(f"Unsupported artifact format: {self.format!r}") - if self.key is None: - return value - if not isinstance(value, dict): - raise TypeError("Artifact key requires a JSON object artifact.") - return value[self.key] - - -Artifacts: TypeAlias = dict[str, ArtifactConfig] - - -class ArtifactsConfig(Config): - model_config = ConfigDict(extra="allow") - - @model_validator(mode="before") - @classmethod - def validate_mapping_input(cls, value: object) -> object: - if isinstance(value, ArtifactsConfig): - return value - if value is None: - return {} - if not isinstance(value, dict): - raise TypeError("ArtifactsConfig must be a mapping.") - for name, source in value.items(): - if not isinstance(name, str): - raise TypeError("ArtifactsConfig keys must be strings.") - validate_serializable_value(source, f"artifacts.{name}") - return value - - @model_validator(mode="after") - def validate_entries(self) -> Self: - for name, source in self.raw_entries().items(): - if not isinstance(name, str): - raise TypeError("ArtifactsConfig keys must be strings.") - validate_serializable_value(source, f"artifacts.{name}") - ArtifactConfig.model_validate(source) - return self - - def raw_entries(self) -> dict[str, ConfigData | ArtifactConfig]: - return cast( - dict[str, ConfigData | ArtifactConfig], dict(self.model_extra or {}) - ) - - def artifacts(self, field: str = "artifacts") -> Artifacts: - artifacts: Artifacts = {} - for name, source in self.raw_entries().items(): - validate_serializable_value(source, f"{field}.{name}") - artifacts[name] = cast( - ArtifactConfig, ArtifactConfig.model_validate(source) - ) - return artifacts - - def data(self, field: str = "artifacts") -> ConfigData: - return { - name: artifact.data() for name, artifact in self.artifacts(field).items() - } diff --git a/verifiers/v1/config.py b/verifiers/v1/config.py index 941adaeb24..5410440838 100644 --- a/verifiers/v1/config.py +++ b/verifiers/v1/config.py @@ -1,127 +1,16 @@ -from os import PathLike -from typing import Literal, TypeAlias +from collections.abc import Mapping +from typing import TypeAlias -from pydantic import BaseModel, field_validator, model_validator +from pydantic import BaseModel from pydantic_config import BaseConfig -from typing_extensions import Self -from .types import ConfigData from .utils.config_utils import ( - annotation_text, - coerce_config, - default_text, import_config_ref as import_config_ref, resolve_config_object as resolve_config_object, - string_mapping, ) -ConfigSource: TypeAlias = BaseModel | ConfigData | None +ConfigSource: TypeAlias = BaseModel | Mapping[str, object] | None -class Config(BaseConfig): - """Strict serializable v1 config base.""" - - @model_validator(mode="after") - def validate_serializable_config(self) -> Self: - for name in type(self).model_fields: - try: - validate_serializable_value( - self.__dict__[name], f"{type(self).__name__}.{name}" - ) - except TypeError as exc: - raise ValueError(str(exc)) from exc - return self - - @classmethod - def schema_text(cls) -> str: - lines = [cls.__name__] - for name, field in cls.model_fields.items(): - lines.append( - f"- {name}: {annotation_text(field.annotation)} = {default_text(field)}" - ) - return "\n".join(lines) - - -def validate_serializable_value(value: object, field: str) -> None: - if value is None or isinstance(value, str | int | float | bool): - return - if isinstance(value, BaseModel): - return - if callable(value) or isinstance(value, PathLike): - raise TypeError(f"{field} must be serializable; use an import ref string.") - if isinstance(value, dict): - for key, item in value.items(): - if not isinstance(key, str): - raise TypeError(f"{field} mapping keys must be strings.") - validate_serializable_value(item, f"{field}.{key}") - return - if isinstance(value, list | tuple): - for index, item in enumerate(value): - validate_serializable_value(item, f"{field}.{index}") - return - raise TypeError(f"{field} must be serializable; got {type(value).__name__}.") - - -class CallableConfig(Config): - fn: str - priority: int | None = None - stage: Literal["rollout", "group"] | None = None - weight: float | None = None - skip: bool = False - - -CallableEntry: TypeAlias = str | CallableConfig -ToolEntryData: TypeAlias = str | ConfigData -ToolsetData: TypeAlias = str | ConfigData -ToolsetCollectionData: TypeAlias = ( - ToolsetData | list[ToolsetData] | dict[str, ToolsetData] -) - - -class SignalConfig(Config): - stage: Literal["rollout", "group"] | None = None - priority: int | None = None - weight: float | None = None - skip: bool = False - - -def validate_scoring_map(value: object, field: str) -> dict[str, ConfigData]: - if value is None: - return {} - if not isinstance(value, dict): - raise TypeError(f"{field} must be a mapping.") - result: dict[str, ConfigData] = {} - for name, item in value.items(): - if not isinstance(name, str): - raise TypeError(f"{field} keys must be strings.") - if isinstance(item, BaseModel): - data = item.model_dump(exclude_none=True, exclude_unset=True) - elif isinstance(item, dict): - data = string_mapping(item) - else: - raise TypeError(f"{field}.{name} must be a mapping.") - result[name] = coerce_config(SignalConfig, data).model_dump( - exclude_none=True, - exclude_unset=True, - ) - return result - - -class LifecycleConfig(Config): - # Collection fields are configured only here; runtime mutation APIs are separate. - toolsets: ToolsetCollectionData = [] - stops: list[CallableEntry] = [] - setups: list[CallableEntry] = [] - updates: list[CallableEntry] = [] - metrics: list[CallableEntry] = [] - rewards: list[CallableEntry] = [] - advantages: list[CallableEntry] = [] - cleanups: list[CallableEntry] = [] - teardowns: list[CallableEntry] = [] - scoring: dict[str, ConfigData] = {} - - @field_validator("scoring", mode="before") - @classmethod - def validate_scoring(cls, value: object) -> dict[str, ConfigData]: - return validate_scoring_map(value, "scoring") +Config = BaseConfig diff --git a/verifiers/v1/decorators.py b/verifiers/v1/decorators.py new file mode 100644 index 0000000000..23c1f94372 --- /dev/null +++ b/verifiers/v1/decorators.py @@ -0,0 +1,166 @@ +from __future__ import annotations + +import inspect +from collections.abc import Callable +from typing import Literal, TypeVar, overload + +from .types import Handler + +SignalStage = Literal["rollout", "group"] +F = TypeVar("F", bound=Callable[..., object]) + + +def discover_decorated(obj: object, attr: str) -> list[Handler]: + methods = [ + method + for _, method in inspect.getmembers(obj, predicate=inspect.ismethod) + if hasattr(method, attr) and callable(method) + ] + priority_attr = f"{attr}_priority" + return sorted( + methods, + key=lambda method: (-getattr(method, priority_attr, 0), method.__name__), + ) + + +def _mark( + func: F | None, + *, + attr: str, + priority: int, + stage: SignalStage | None = None, + weight: float | None = None, +) -> F | Callable[[F], F]: + def decorator(f: F) -> F: + setattr(f, attr, True) + setattr(f, f"{attr}_priority", priority) + if stage is not None: + setattr(f, f"{attr}_stage", stage) + if weight is not None: + setattr(f, f"{attr}_weight", weight) + return f + + return decorator if func is None else decorator(func) + + +@overload +def stop(func: F, priority: int = 0) -> F: ... + + +@overload +def stop(func: None = None, priority: int = 0) -> Callable[[F], F]: ... + + +def stop(func: F | None = None, priority: int = 0) -> F | Callable[[F], F]: + return _mark(func, attr="stop", priority=priority) + + +@overload +def setup(func: F, priority: int = 0) -> F: ... + + +@overload +def setup(func: None = None, priority: int = 0) -> Callable[[F], F]: ... + + +def setup(func: F | None = None, priority: int = 0) -> F | Callable[[F], F]: + return _mark(func, attr="setup", priority=priority) + + +@overload +def cleanup(func: F, priority: int = 0, stage: SignalStage = "rollout") -> F: ... + + +@overload +def cleanup( + func: None = None, priority: int = 0, stage: SignalStage = "rollout" +) -> Callable[[F], F]: ... + + +def cleanup( + func: F | None = None, priority: int = 0, stage: SignalStage = "rollout" +) -> F | Callable[[F], F]: + return _mark(func, attr="cleanup", priority=priority, stage=stage) + + +@overload +def update(func: F, priority: int = 0, stage: SignalStage = "rollout") -> F: ... + + +@overload +def update( + func: None = None, priority: int = 0, stage: SignalStage = "rollout" +) -> Callable[[F], F]: ... + + +def update( + func: F | None = None, priority: int = 0, stage: SignalStage = "rollout" +) -> F | Callable[[F], F]: + return _mark(func, attr="update", priority=priority, stage=stage) + + +@overload +def metric(func: F, priority: int = 0, stage: SignalStage = "rollout") -> F: ... + + +@overload +def metric( + func: None = None, priority: int = 0, stage: SignalStage = "rollout" +) -> Callable[[F], F]: ... + + +def metric( + func: F | None = None, priority: int = 0, stage: SignalStage = "rollout" +) -> F | Callable[[F], F]: + return _mark(func, attr="metric", priority=priority, stage=stage) + + +@overload +def reward( + func: F, + weight: float = 1.0, + priority: int = 0, + stage: SignalStage = "rollout", +) -> F: ... + + +@overload +def reward( + func: None = None, + weight: float = 1.0, + priority: int = 0, + stage: SignalStage = "rollout", +) -> Callable[[F], F]: ... + + +def reward( + func: F | None = None, + weight: float = 1.0, + priority: int = 0, + stage: SignalStage = "rollout", +) -> F | Callable[[F], F]: + return _mark(func, attr="reward", priority=priority, stage=stage, weight=weight) + + +@overload +def advantage(func: F, priority: int = 0) -> F: ... + + +@overload +def advantage(func: None = None, priority: int = 0) -> Callable[[F], F]: ... + + +def advantage(func: F | None = None, priority: int = 0) -> F | Callable[[F], F]: + return _mark(func, attr="advantage", priority=priority, stage="group") + + +@overload +def teardown(func: F, priority: int = 0) -> F: ... + + +@overload +def teardown(func: None = None, priority: int = 0) -> Callable[[F], F]: ... + + +def teardown(func: F | None = None, priority: int = 0) -> F | Callable[[F], F]: + return _mark(func, attr="teardown", priority=priority) diff --git a/verifiers/v1/env.py b/verifiers/v1/env.py index 9cb62f7c4b..50778d7edf 100644 --- a/verifiers/v1/env.py +++ b/verifiers/v1/env.py @@ -1,203 +1,238 @@ -import asyncio +from __future__ import annotations + import uuid -from typing import TYPE_CHECKING, cast +from typing import TYPE_CHECKING, final -import verifiers as vf -from verifiers.clients import Client -from verifiers.types import ClientConfig -from verifiers.types import RolloutInput, SamplingArgs +from pydantic import Field, field_validator +from pydantic import BaseModel +from verifiers.types import ( + RolloutInput, +) +from . import advantages from .config import Config -from .harness import Harness, HarnessConfig +from .harness import Harness +from .lifecycle import EnvRun +from .runtime import ( + RuntimeConfig, + RuntimeConfigValue, + RuntimeProvider, +) from .state import State -from .taskset import Taskset, TasksetConfig -from .types import JsonData, RuntimeData -from .utils.taskset_utils import task_from_dataset_record +from .task import Task +from .taskset import Taskset +from .types import JsonData, ModelClient, ModelConfig +from .utils.config_utils import explicit_config_data +from .utils.scoring_utils import score_group as score_group_signals if TYPE_CHECKING: from datasets import Dataset +@final class EnvConfig(Config): - taskset: TasksetConfig = TasksetConfig() - harness: HarnessConfig = HarnessConfig() + taskset: dict[str, object] = Field(default_factory=dict) + harness: dict[str, object] = Field(default_factory=dict) + runtime: RuntimeConfig | None = None + advantage: advantages.AdvantageConfig = "rl" + @field_validator("taskset", "harness", mode="before") @classmethod - def __pydantic_init_subclass__(cls, **kwargs: object) -> None: - super().__pydantic_init_subclass__(**kwargs) - extra_fields = set(cls.model_fields) - set(EnvConfig.model_fields) - if extra_fields: - raise TypeError( - f"{cls.__name__} defines unsupported root env config fields: " - f"{', '.join(sorted(extra_fields))}. Put env-specific settings on " - "a TasksetConfig or HarnessConfig instead." - ) - for field_name, expected_type in ( - ("taskset", TasksetConfig), - ("harness", HarnessConfig), - ): - annotation = cls.model_fields[field_name].annotation - if not ( - isinstance(annotation, type) and issubclass(annotation, expected_type) - ): - raise TypeError( - f"{cls.__name__}.{field_name} must be typed as a " - f"{expected_type.__name__} subclass." - ) + def serialize_child_config(cls, value: object) -> object: + if value is None: + return {} + if isinstance(value, BaseModel): + return explicit_config_data(value) + return value -class Env(vf.Environment): +class Env: def __init__( self, *, - taskset: Taskset | None = None, + taskset: Taskset, harness: Harness | None = None, + runtime: RuntimeProvider | RuntimeConfigValue | None = None, + advantage: advantages.AdvantageConfig = "rl", ): - if taskset is None: - raise TypeError("Env requires a taskset.") if not isinstance(taskset, Taskset): raise TypeError("Env taskset must be a Taskset.") if harness is not None and not isinstance(harness, Harness): raise TypeError("Env harness must be a Harness.") self.taskset = taskset - self.harness = harness or Harness(config=HarnessConfig()) + self.harness = harness or Harness() + self.harness.bind(taskset=self.taskset, runtime=runtime) + self.advantage = advantage + self.advantage_function = advantages.resolve_config(advantage) + self.runtime_config = self.harness.runtime_config + self.runtime_provider = self.harness.runtime_provider self.config = EnvConfig( - taskset=cast(TasksetConfig, self.taskset.config), - harness=cast(HarnessConfig, self.harness.config), + taskset=explicit_config_data(self.taskset.config), + harness=explicit_config_data(self.harness.config), + runtime=self.runtime_config, + advantage=self.advantage, ) - self.harness.taskset = self.taskset - self.taskset.runtime_refresh = self.harness.rebuild_runtime - self.harness.rebuild_runtime() - super().__init__( - dataset=self.taskset.get_dataset, - eval_dataset=self.taskset.get_eval_dataset, - rubric=vf.Rubric(), - ) - self._empty_dataset_checked = False - self._empty_eval_dataset_checked = False - - def build_dataset(self) -> "Dataset | None": - if self.dataset is not None: - return self.dataset - if self._empty_dataset_checked: - return None - dataset = self.taskset.get_dataset() - if not len(dataset): - self._empty_dataset_checked = True - return None - self.dataset = self._format_dataset_source(dataset) - return self.dataset - - def build_eval_dataset(self) -> "Dataset | None": - if self.eval_dataset is not None: - return self.eval_dataset - if self._empty_eval_dataset_checked: - return None - eval_dataset = self.taskset.get_eval_dataset() - if not len(eval_dataset): - self._empty_eval_dataset_checked = True - return None - self.eval_dataset = self._format_dataset_source(eval_dataset) - return self.eval_dataset - - @vf.teardown - async def teardown_harness(self) -> None: - await self.harness.teardown() + self.env_id = "" + self.env_args: JsonData = {} + self.pass_threshold = 0.5 @property def requires_group_rollouts(self) -> bool: uses_custom_init_group = type(self.taskset).init_group is not Taskset.init_group - return self.harness.runtime.has_group_stage or uses_custom_init_group + return ( + self.advantage_function is not None + or self.taskset.has_group_signals + or any(signal["stage"] == "group" for signal in self.harness.signals) + or uses_custom_init_group + ) @property def provides_advantages(self) -> bool: - return self.harness.runtime.has_group_advantages + return self.advantage_function is not None - async def rollout( - self, - input: RolloutInput, - client: Client | ClientConfig, - model: str, - sampling_args: SamplingArgs | None = None, - ) -> State: - task = task_from_dataset_record(cast(JsonData, input), self.taskset.taskset_id) - state = State.for_task(task) - self.apply_controls( - [state], - { - "client": client, - "model": model, - "sampling_args": sampling_args or {}, - "score_rollout": self.score_rollouts, - }, - ) - return await self.harness.run(task, state) + def get_dataset(self, n: int = -1, seed: int | None = None) -> "Dataset": + dataset = self.taskset.get_dataset() + if seed is not None: + dataset = dataset.shuffle(seed=seed) + if n > 0: + return dataset.select(range(min(n, len(dataset)))) + return dataset + + def get_eval_dataset(self, n: int = -1, seed: int | None = None) -> "Dataset": + dataset = self.taskset.get_eval_dataset() + if not len(dataset): + dataset = self.taskset.get_dataset() + if seed is not None: + dataset = dataset.shuffle(seed=seed) + if n > 0: + return dataset.select(range(min(n, len(dataset)))) + return dataset - async def _run_rollout_state( + def run(self) -> EnvRun: + return EnvRun(env=self) + + async def run_handlers_for_group( self, - input: RolloutInput, - client: Client, - model: str, - sampling_args: SamplingArgs, - ) -> State: - return await self.rollout(input, client, model, sampling_args) + kind: str, + tasks: list[Task], + states: list[State], + teacher: ModelClient | None = None, + ) -> None: + if not tasks or not states: + return + handlers = [*self.taskset.handlers[kind], *self.harness.handlers[kind]] + for handler in handlers: + if getattr(handler, f"{kind}_stage", "rollout") != "group": + continue + result = await self.harness.call_handler( + handler, + tasks[0], + states[0], + tasks=tasks, + states=states, + teacher=teacher, + teacher_name=teacher.config.model if teacher is not None else None, + ) + if result is not None: + raise TypeError(f"Group {kind} handlers must mutate states in place.") - async def _run_group_states( + async def run_rollout( self, - group_inputs: list[RolloutInput], - client: Client, - model: str, - sampling_args: SamplingArgs, - ) -> list[vf.State]: - base_task = task_from_dataset_record( - cast(JsonData, group_inputs[0]), self.taskset.taskset_id - ) - tasks, states = await self.taskset.init_group(base_task, len(group_inputs)) - if len(tasks) != len(group_inputs) or len(states) != len(group_inputs): - raise ValueError( - "Taskset.init_group must return one task/state per rollout." + input: RolloutInput | Task, + *, + model: ModelConfig, + teacher: ModelConfig | None = None, + state: State | None = None, + max_retries: int = 0, + ) -> State: + async with self.run() as env_run: + return await env_run.run_rollout( + input, + model=model, + teacher=teacher, + state=state, + score=True, + max_retries=max_retries, ) - group_key = uuid.uuid4().hex - for state in states: - state.runtime_state()["group_key"] = group_key - self.apply_controls( - states, - { - "client": client, - "model": model, - "sampling_args": sampling_args, - "score_rollout": self.score_rollouts, - }, - ) - states = await asyncio.gather( - *[self.harness.run(task, state) for task, state in zip(tasks, states)] - ) - try: - if self.score_rollouts: - await self.harness.score_group(tasks, states) - finally: - await self.harness.cleanup_group(tasks, states) - for state in states: - state.strip_runtime_handles() - state.assert_serializable() - return cast(list[vf.State], states) - def apply_controls( - self, states: list[State], controls: RuntimeData | None = None + async def score_group( + self, + tasks: list[Task], + states: list[State], + *, + model: ModelConfig | None = None, + teacher: ModelConfig | None = None, ) -> list[State]: - if controls is None: + if len(tasks) != len(states): + raise ValueError("score_group requires one state per task.") + if not states: return states - serializable_controls = { - key: value for key, value in controls.items() if key != "client" - } + model_config = ( + model if model is not None else (states[0].model if states else None) + ) + teacher_config = ( + teacher if teacher is not None else (states[0].teacher if states else None) + ) + if model_config is not None: + for state in states: + if state.model is None: + state.model = model_config + elif state.model != model_config: + raise ValueError("Group states must use one model config.") + if teacher_config is not None: + for state in states: + if state.teacher is None: + state.teacher = teacher_config + elif state.teacher != teacher_config: + raise ValueError("Group states must use one teacher config.") + group_id = uuid.uuid4().hex for state in states: - runtime_state = state.runtime_state() - client = controls.get("client") - self.harness.runtime.bind_model_client( - state, - cast(Client | ClientConfig | None, client) - if client is not None - else None, + state.group_id = state.group_id or group_id + model_client: ModelClient | None = None + teacher_client: ModelClient | None = None + try: + model_client = ( + self.harness.load_model_client(model_config) + if model_config is not None + else None + ) + teacher_client = ( + self.harness.load_model_client(teacher_config) + if teacher_config is not None + else None + ) + signals = self.harness.owner_signals() + if self.advantage_function is not None: + signals = [ + signal for signal in signals if signal["kind"] != "advantage" + ] + signals.append(advantages.signal(self.advantage_function)) + await score_group_signals( + signals, + tasks, + states, + model_client=model_client, + teacher=teacher_client, ) - runtime_state.update(serializable_controls) + await self.run_handlers_for_group( + "update", tasks, states, teacher=teacher_client + ) + finally: + try: + await self.run_handlers_for_group( + "cleanup", tasks, states, teacher=teacher_client + ) + finally: + try: + if teacher_client is not None: + await self.harness.close_model_client(teacher_client) + finally: + if model_client is not None: + await self.harness.close_model_client(model_client) + for state in states: + self.harness.validate_extras(state) + state.assert_serializable() return states + + async def close(self) -> None: + await self.harness.close() diff --git a/verifiers/v1/eval.py b/verifiers/v1/eval.py new file mode 100644 index 0000000000..60d1a3f74b --- /dev/null +++ b/verifiers/v1/eval.py @@ -0,0 +1,242 @@ +from __future__ import annotations + +import asyncio +from collections import defaultdict +from pathlib import Path +from typing import cast + +from verifiers.types import ( + EvalConfig, + GenerateOutputs, + LogCallback, + ProgressCallback, + RolloutInput, + RolloutOutput, + StartCallback, +) +from verifiers.utils.eval_utils import effective_max_concurrent, filter_inputs +from verifiers.utils.path_utils import is_valid_eval_results_path +from verifiers.utils.save_utils import ( + GenerateOutputsBuilder, + load_outputs, + push_results_to_hf_hub, + save_metadata, + save_outputs, + truncate_malformed_trailing_line, + validate_resume_metadata, +) + +from .env import Env +from .lifecycle import EnvRun +from .types import ModelConfig +from .utils.json_utils import json_data + + +def eval_inputs( + env: Env, + num_examples: int, + rollouts_per_example: int, + seed: int | None = None, +) -> list[RolloutInput]: + rows = [ + cast(RolloutInput, json_data(row)) + for row in env.get_eval_dataset(n=num_examples, seed=seed) + ] + if rollouts_per_example <= 1: + return rows + return [ + cast(RolloutInput, dict(row)) + for _ in range(rollouts_per_example) + for row in rows + ] + + +def progress_inputs( + env: Env, + inputs: list[RolloutInput], + independent_scoring: bool, +) -> list[RolloutInput] | list[list[RolloutInput]]: + if independent_scoring or not env.requires_group_rollouts: + return inputs + groups: dict[object, list[RolloutInput]] = defaultdict(list) + for row in inputs: + groups[row["example_id"]].append(row) + return list(groups.values()) + + +async def run_rollouts( + env_run: EnvRun, + inputs: list[RolloutInput], + model: ModelConfig, + max_concurrent: int, + max_retries: int, + state_columns: list[str], +) -> list[RolloutOutput]: + semaphore = asyncio.Semaphore(max_concurrent) if max_concurrent > 0 else None + + async def run_one(row: RolloutInput) -> RolloutOutput: + if semaphore is None: + state = await env_run.run_rollout(row, model=model, max_retries=max_retries) + else: + async with semaphore: + state = await env_run.run_rollout( + row, model=model, max_retries=max_retries + ) + task = env_run.to_task(row) + return RolloutOutput(state.to_output(task, state_columns)) + + return list(await asyncio.gather(*(run_one(row) for row in inputs))) + + +async def run_rollout_groups( + env_run: EnvRun, + inputs: list[RolloutInput], + model: ModelConfig, + max_concurrent: int, + max_retries: int, + state_columns: list[str], +) -> list[RolloutOutput]: + semaphore = asyncio.Semaphore(max_concurrent) if max_concurrent > 0 else None + groups: dict[object, list[RolloutInput]] = defaultdict(list) + for row in inputs: + groups[row["example_id"]].append(row) + + async def run_rows(rows: list[RolloutInput]) -> list[RolloutOutput]: + group = await env_run.group(rows) + if semaphore is None: + states = await group.run(model=model, max_retries=max_retries) + else: + async with semaphore: + states = await group.run(model=model, max_retries=max_retries) + return [ + RolloutOutput(state.to_output(task, state_columns)) + for task, state in zip(group.tasks, states, strict=True) + ] + + grouped_outputs = await asyncio.gather( + *(run_rows(rows) for rows in groups.values()) + ) + return [output for group_outputs in grouped_outputs for output in group_outputs] + + +async def run_evaluation( + env: Env, + config: EvalConfig, + results_path: Path, + on_start: StartCallback | None, + on_progress: ProgressCallback | list[ProgressCallback] | None, + on_log: LogCallback | None, +) -> GenerateOutputs: + if config.extra_env_kwargs: + raise ValueError("extra_env_kwargs are only supported by legacy environments.") + + raw_inputs = eval_inputs( + env, + config.num_examples, + config.rollouts_per_example, + seed=config.shuffle_seed if config.shuffle else None, + ) + example_ids = {row["example_id"] for row in raw_inputs} + num_examples = len(example_ids) + rollouts_per_example = ( + len(raw_inputs) // num_examples + if num_examples > 0 + else config.rollouts_per_example + ) + model_config = ModelConfig( + client=config.client_config, + model=config.model, + sampling_args=dict(config.sampling_args), + ) + builder = GenerateOutputsBuilder( + env_id=env.env_id, + env_args=env.env_args, + model=config.model, + client=config.client_config, + num_examples=num_examples, + rollouts_per_example=rollouts_per_example, + state_columns=config.state_columns, + sampling_args=dict(config.sampling_args), + results_path=results_path, + pass_threshold=env.pass_threshold, + ) + callbacks: list[ProgressCallback] = [] + if isinstance(on_progress, list): + callbacks = cast(list[ProgressCallback], on_progress) + elif on_progress is not None: + callbacks = [on_progress] + + try: + if is_valid_eval_results_path(results_path): + validate_resume_metadata( + results_path=results_path, + env_id=env.env_id, + model=config.model, + num_examples=num_examples, + rollouts_per_example=rollouts_per_example, + ) + if on_log is not None: + on_log(f"Resuming evaluation from {results_path}") + outputs = load_outputs(results_path) + truncate_malformed_trailing_line(results_path / "results.jsonl") + builder.add_outputs(outputs) + filtered_inputs = filter_inputs(raw_inputs, outputs, rollouts_per_example) + else: + filtered_inputs = raw_inputs + + if on_start is not None: + on_start( + raw_inputs, + progress_inputs(env, filtered_inputs, config.independent_scoring), + ) + if not filtered_inputs: + return builder.build(sort_by_example_id=True) + if config.save_results and on_log is not None: + on_log(f"Saving results to {builder.results_path}") + + if config.independent_scoring or not env.requires_group_rollouts: + state_columns = config.state_columns or [] + async with env.run() as env_run: + new_outputs = await run_rollouts( + env_run, + filtered_inputs, + model_config, + config.max_concurrent, + config.max_retries, + state_columns, + ) + else: + state_columns = config.state_columns or [] + async with env.run() as env_run: + new_outputs = await run_rollout_groups( + env_run, + filtered_inputs, + model_config, + effective_max_concurrent(config), + config.max_retries, + state_columns, + ) + builder.add_outputs(new_outputs) + metadata = builder.build_metadata() + for callback in callbacks: + callback(builder.outputs, new_outputs, metadata) + + results = builder.build(sort_by_example_id=True) + if config.save_results: + await asyncio.to_thread( + save_outputs, + results["outputs"], + builder.results_path, + ) + await asyncio.to_thread( + save_metadata, + results["metadata"], + builder.results_path, + ) + if config.save_to_hf_hub: + push_results_to_hf_hub(results, config.hf_hub_dataset_name) + if on_log is not None: + on_log(f"Saved final results to {results['metadata']['path_to_save']}") + return results + finally: + await env.close() diff --git a/verifiers/v1/harness.py b/verifiers/v1/harness.py index ddbfdf7e92..74b1e5bb74 100644 --- a/verifiers/v1/harness.py +++ b/verifiers/v1/harness.py @@ -1,136 +1,74 @@ +from __future__ import annotations + import asyncio -from collections.abc import Awaitable, Callable -from typing import TYPE_CHECKING, Generic, TypeAlias, TypeVar, cast, final - -from pydantic import AliasChoices, Field - -import verifiers as vf -from verifiers.clients.client import Client -from verifiers.errors import Error, OverlongPromptError, SandboxError -from verifiers.types import ( - ClientConfig, - MessageContent, - Messages, - SamplingArgs, - ToolMessage, -) +import inspect +import time +from contextlib import asynccontextmanager +from copy import deepcopy +from pydantic import TypeAdapter +from typing import TYPE_CHECKING, AsyncIterator, Generic, TypeVar, cast, final + + +from verifiers.clients import resolve_client +from verifiers.errors import Error, OverlongPromptError, ToolError +from verifiers.types import Messages, SamplingArgs, ToolMessage, UserMessage from verifiers.utils.async_utils import maybe_call_with_named_args -from verifiers.utils.message_utils import normalize_messages from verifiers.utils.response_utils import parse_response_message -from verifiers.utils.tool_utils import is_valid_tool_content_parts -from .config import ( - ConfigSource, - LifecycleConfig, - import_config_ref, -) -from .artifact import ArtifactsConfig -from .program import ( - ProgramConfig, -) -from .model import ModelConfig, model_config_from_task -from .sandbox import SandboxConfig -from .user import UserConfig -from .utils.binding_utils import ( - BindingSources, - BindingsConfig, - ObjectsConfig, -) -from .utils.endpoint_utils import ( - Endpoint, - assistant_completion_from_messages, - run_intercepted_program, +from .config import Config, ConfigSource +from .decorators import discover_decorated +from .interception import EndpointProtocol +from .protocols import default_protocols +from .mcp import BoundUpdate, ServerResult, MCPToolRegistry, ServerResponse +from .runtime import ( + RuntimeConfig, + RuntimeConfigValue, + RuntimeProvider, + Runtime, + SubprocessRuntimeConfig, + make_runtime_provider, + resolve_runtime_config, ) +from .state import Extras, State, TimeSpan, Turn, TurnTokens, TurnUsage +from .task import Task +from .types import Handler, JsonData, JsonValue, ModelClient, ModelConfig, Context from .utils.config_utils import ( coerce_config, config_ref_context, config_type_from_class, - qualified_config_ref, registered_config_type, register_config_type, ) -from .utils.runtime_owner_utils import RuntimeOwnerMixin -from .utils.json_utils import json_args -from .utils.mcp_proxy_utils import ( - proxy_program, - proxy_sandbox, -) -from .utils.program_utils import ( - merge_task_program, - merge_task_sandbox, - program_channels, - program_kind, - run_local_command, - validate_program_options, - validate_program_sandbox_scope, -) -from .runtime import Runtime -from .utils.sandbox_utils import run_sandbox_command -from .utils.sandbox_program_utils import ( - python_program_sandbox, - run_sandbox_python_program, -) -from .utils.logging_utils import log_rollout_finish, log_rollout_start +from .utils.json_utils import json_args, json_data from .utils.prompt_utils import ( SystemPrompt, SystemPromptStrategy, - normalize_prompt, normalize_system_prompt, resolve_system_prompt, ) -from .utils.tool_utils import tool_error_content -from .utils.trajectory_utils import has_borrowed_trajectory, sync_trajectory -from .state import State -from .task import Task -from .types import ( - Handler, - ModelClient, - ConfigData, - JsonData, - Objects, -) +from .utils.scoring_utils import SignalRecord, build_signals +from .utils.scoring_utils import score_rollout if TYPE_CHECKING: from .taskset import Taskset -ProgramResult: TypeAlias = State | JsonData | None -ProgramRunner: TypeAlias = Callable[[Task, State], Awaitable[ProgramResult]] +_MESSAGES_ADAPTER = TypeAdapter(Messages) -class HarnessConfig(LifecycleConfig): - # Core fields configure harness-owned runtime behavior. - harness_id: str | None = Field( - default=None, - validation_alias=AliasChoices("harness_id", "id"), - ) - program: ProgramConfig = ProgramConfig() - model: ModelConfig = ModelConfig() - version: str | None = None +class HarnessConfig(Config): + id: str | None = None system_prompt: SystemPrompt = None system_prompt_strategy: SystemPromptStrategy = "HT" - sandbox: SandboxConfig | None = None - user: UserConfig | None = None - bindings: BindingsConfig = BindingsConfig() - objects: ObjectsConfig = ObjectsConfig() - artifacts: ArtifactsConfig = ArtifactsConfig() max_turns: int = -1 - - @classmethod - def __pydantic_init_subclass__(cls, **kwargs: object) -> None: - super().__pydantic_init_subclass__(**kwargs) - field = cls.model_fields.get("harness_id") - if field is not None: - field.validation_alias = AliasChoices("harness_id", "id") - cls.model_rebuild(force=True) + runtime: RuntimeConfig | None = None + extras: Extras | None = None ConfigT = TypeVar("ConfigT", bound=HarnessConfig) -class Harness(RuntimeOwnerMixin[ConfigT], Generic[ConfigT]): +class Harness(Generic[ConfigT]): config: ConfigT - program_config: ProgramConfig - program: ProgramRunner def __init_subclass__(cls, **kwargs: object) -> None: super().__init_subclass__(**kwargs) @@ -144,523 +82,866 @@ def __init_subclass__(cls, **kwargs: object) -> None: register_config_type(cls, config_type) @final - def __init__( - self, - config: ConfigSource = None, - *, - model: str | ModelConfig | None = None, - client: ModelClient | str | None = None, - sampling_args: SamplingArgs | None = None, - ): - model_kwargs_present = ( - model is not None or client is not None or sampling_args is not None - ) - if config is not None and model_kwargs_present: - raise TypeError("Pass either config or model/client/sampling_args kwargs.") - runtime_model_client: Client | None = None + def __init__(self, config: ConfigSource = None): config_type = registered_config_type(type(self), HarnessConfig) - if config is None and model_kwargs_present: - if isinstance(model, ModelConfig): - if client is not None or sampling_args is not None: - raise TypeError( - "ModelConfig cannot be combined with client or sampling_args." - ) - config = config_type(model=model) - else: - model_client: ClientConfig | str | None = None - if client is not None: - if isinstance(client, Client): - runtime_model_client = client - else: - model_client = client - config = config_type( - model=ModelConfig( - name=model, - client=model_client, - sampling_args=sampling_args or {}, - ) - ) self.config = cast(ConfigT, coerce_config(config_type, config)) with config_ref_context(self.config): - self.initialize_runtime_refresh() - resolved_harness_id = self.config.harness_id - if resolved_harness_id is not None and not isinstance( - resolved_harness_id, str - ): - raise TypeError("harness_id must be a string.") - self.harness_id = resolved_harness_id or type(self).__name__ - self.program_config = self.load_program_config(self.config) - system_prompt_value = self.load_system_prompt(self.config) + resolved_id = self.config.id + if resolved_id is not None and not isinstance(resolved_id, str): + raise TypeError("harness id must be a string.") + self.id = resolved_id or type(self).__name__ self.system_prompt = normalize_system_prompt( - system_prompt_value, field_name="harness.system_prompt" + self.load_system_prompt(self.config), field_name="harness.system_prompt" ) self.system_prompt_strategy = self.config.system_prompt_strategy - self.initialize_runtime_user(self.config.user) - self.bindings: BindingSources = self.config.bindings.entries( - "harness.bindings" - ) - self.objects: Objects = self.load_objects(self.config.objects) - self.artifacts = self.load_artifacts(self.config.artifacts) - self.sandbox = self.load_sandbox(self.config.sandbox) - self.model = self.load_model(self.config.model) - self.model_client: ModelClient | None = cast( - ModelClient | None, - self.model.client_object(), - ) - if runtime_model_client is not None: - self.model_client = runtime_model_client - self.initialize_runtime_toolsets(self.config, self.config.toolsets) - self.initialize_runtime_handlers() - self.taskset: "Taskset | None" = None - self.runtime = Runtime(taskset=self.taskset, harness=self) - self.runtime_refresh = self.rebuild_runtime - self.endpoint = self.load_endpoint() - self.program = self.compile_program(self.program_config) + self.handlers = self.load_handlers() + self.protocols = self.load_protocols() + self.signals = build_signals(self) + for signal in self.signals: + if signal["kind"] != "metric": + raise ValueError("Harness signals must be metrics.") + self.taskset: Taskset | None = None + self.runtime_config = resolve_runtime_config(self.config.runtime) + self.runtime_provider: RuntimeProvider | None = None + self._env_toolsets: MCPToolRegistry | None = None + self._env_user: MCPToolRegistry | None = None + self._env_servers_lock = asyncio.Lock() + self._env_scope_count = 0 + self.extras_schema: type[Extras] | None = Extras.schema_for(self.config.extras) + self.extras_defaults: JsonData = Extras.defaults_for(self.config.extras) def load_system_prompt(self, config: ConfigT) -> SystemPrompt: return config.system_prompt - def load_program_config(self, config: ConfigT) -> ProgramConfig: - return config.program.resolve() + def load_handlers(self) -> dict[str, list[Handler]]: + handlers: dict[str, list[Handler]] = { + "stop": [], + "setup": [], + "update": [], + "cleanup": [], + "teardown": [], + } + for kind in handlers: + handlers[kind].extend(discover_decorated(self, kind)) + return handlers - def load_sandbox(self, config: SandboxConfig | None) -> SandboxConfig | None: - sandbox = self.program_config.sandbox - if sandbox is None or sandbox is False: - return config - base = config.data() if config is not None else {} - if sandbox is True: - return SandboxConfig.model_validate(base) - return SandboxConfig.model_validate({**base, **sandbox.data()}) + def load_protocols(self) -> list[EndpointProtocol]: + return default_protocols() - def load_model(self, config: ModelConfig) -> ModelConfig: - return config - - def load_endpoint(self) -> Endpoint: - return Endpoint( - use_tunnel=self.program_sandbox_config(self.program_config) is not None - ) + def load_model_client(self, config: ModelConfig) -> ModelClient: + return ModelClient(config=config, client=resolve_client(config.client)) - def rebuild_runtime(self) -> None: - self.runtime = Runtime(taskset=self.taskset, harness=self) + async def close_model_client(self, model_client: ModelClient) -> None: + await model_client.client.close() - async def run(self, task: Task, state: State | None = None) -> State: - state = await self.init_state(task) if state is None else state - log_rollout_start(state) - timing_recorded = False - completed = False + def bind( + self, + *, + taskset: "Taskset | None" = None, + runtime: RuntimeProvider | RuntimeConfigValue | None = None, + ) -> None: + self.taskset = taskset + taskset_runtime = taskset.config.runtime if taskset is not None else None + taskset_extras = None if taskset is None else taskset.config.extras + self.extras_schema = Extras.realize_schema( + Extras.schema_for(taskset_extras), Extras.schema_for(self.config.extras) + ) + self.extras_defaults = Extras.merge_defaults( + Extras.defaults_for(taskset_extras), Extras.defaults_for(self.config.extras) + ) + if isinstance(runtime, RuntimeProvider): + self.runtime_provider = runtime + self.runtime_config = resolve_runtime_config( + taskset_runtime, self.config.runtime + ) + return + self.runtime_config = resolve_runtime_config( + taskset_runtime, self.config.runtime, runtime + ) + self.runtime_provider = None + + def load_runtime_provider(self, config: RuntimeConfigValue) -> RuntimeProvider: + return make_runtime_provider(config) + + def runtime_for(self, task: Task) -> RuntimeConfigValue: + updates: JsonData = {} + if task.image is not None: + if isinstance(self.runtime_config, SubprocessRuntimeConfig): + raise ValueError( + f"task {task.task_id!r} declares an image; use docker or " + "prime runtime." + ) + if not isinstance(task.image, str) or not task.image: + raise TypeError("task.image must be a non-empty string.") + updates["image"] = task.image + for field, value in task.resources.model_dump(exclude_none=True).items(): + spec = type(self.runtime_config).model_fields.get(field) + if spec is None: + raise ValueError( + f"task {task.task_id!r} declares resource {field!r}; runtime " + f"{self.runtime_config.type!r} does not support it." + ) + if getattr(self.runtime_config, field) == spec.default: + updates[field] = value + if not updates: + return self.runtime_config + return self.runtime_config.model_copy(update=updates) + + def runtime_provider_for(self, task: Task) -> RuntimeProvider: + if self.runtime_provider is not None: + return self.runtime_provider + return self.load_runtime_provider(self.runtime_for(task)) + + @asynccontextmanager + async def open_context( + self, + *, + task: Task, + state: State, + model: ModelConfig, + teacher: ModelConfig | None = None, + runtime: Runtime | None = None, + toolsets: MCPToolRegistry | None = None, + user: MCPToolRegistry | None = None, + parent: Context | None = None, + score: bool = False, + ) -> AsyncIterator[Context]: + model_client = ( + parent.model_client + if parent is not None and parent.model_client.config == model + else self.load_model_client(model) + ) + teacher_client = None + if teacher is not None: + teacher_client = ( + parent.teacher + if parent is not None + and parent.teacher is not None + and parent.teacher.config == teacher + else self.load_model_client(teacher) + ) try: - try: - state = await self.setup_state(task, state) - if not await self.runtime.is_completed(task, state): - state = await self.run_program(task, state) - await self.runtime.is_completed(task, state) - state._set_stop_condition("program_completed") - await self.runtime.collect_artifacts(task, state) - except Error as e: - self.record_error(state, e) - await self.runtime.update_rollout(task, state) - state.record_generation_timing() - timing_recorded = True - if state.runtime_state().get("score_rollout", True): - await self.runtime.score_rollout(task, state) - state._set_completed(True) - completed = True + yield Context( + task=task, + state=state, + model_client=model_client, + teacher=teacher_client, + runtime=runtime, + toolsets=toolsets, + user=user, + parent=parent, + score=score, + ) finally: - if not timing_recorded: - state.record_generation_timing() - await self.runtime.cleanup_rollout(task, state) - if "group_key" not in state.runtime_state(): - await self.runtime.cleanup_group([task], [state]) - if completed: - state.finalize() - else: - state.strip_runtime_handles() - elif completed: - state.serialize_error() - state.assert_serializable() - log_rollout_finish(state) - return state - - def record_error(self, state: State, error: Error) -> None: - if state.get("prompt_too_long"): - state._set_truncated(True) - state._set_stop_condition("prompt_too_long", overwrite=True) - return - if isinstance(error, SandboxError) and isinstance(state.get("error"), Error): - state._set_stop_condition("has_error", overwrite=True) - return - if isinstance(error, OverlongPromptError): - state["prompt_too_long"] = True - state._set_truncated(True) - state._set_stop_condition("prompt_too_long", overwrite=True) - return - state._set_error(error) - state._set_stop_condition("has_error", overwrite=True) + if model_client is not ( + parent.model_client if parent is not None else None + ): + await self.close_model_client(model_client) + if teacher_client is not None and teacher_client is not ( + parent.teacher if parent is not None else None + ): + await self.close_model_client(teacher_client) - async def score_group(self, tasks: list[Task], states: list[State]) -> list[State]: - return await self.runtime.score_group(tasks, states) + async def run( + self, + task: Task | str, + state: State | None = None, + *, + model: ModelConfig | str | None = None, + teacher: ModelConfig | str | None = None, + context: Context | None = None, + score: bool = False, + ) -> State: + if isinstance(task, str): + task = Task(prompt=task) + if isinstance(model, str): + model = ModelConfig(model=model) + if isinstance(teacher, str): + teacher = ModelConfig(model=teacher) + if context is not None and score and context.has_active_scoring(): + raise RuntimeError("Nested scored harness runs are not supported.") + using_context_state = state is None and context is not None + if state is None: + state = ( + context.state if context is not None else State(task_id=task.task_id) + ) + if model is None: + if context is None: + raise TypeError("Harness.run requires model unless context is passed.") + model = context.model_client.config + if teacher is None and context is not None and context.teacher is not None: + teacher = context.teacher.config + if not using_context_state or state.task_id is None: + state.task_id = task.task_id + state.model = state.model or model + state.teacher = state.teacher or teacher + self.initialize_extras(state) + + if context is not None: + async with self.open_context( + task=task, + state=state, + model=model, + teacher=teacher, + runtime=context.runtime, + toolsets=context.toolsets, + user=context.user, + parent=context, + score=score, + ) as child_context: + if child_context.toolsets is not None: + child_context.toolsets.set_visibility( + toolsets=task.toolsets, + tools=task.tools, + ) + await self.run_lifecycle(child_context) + return state - async def cleanup_group(self, tasks: list[Task], states: list[State]) -> None: - await self.runtime.cleanup_group(tasks, states) - for state in states: - state.strip_runtime_handles() + from .lifecycle import EnvRun - async def teardown(self) -> None: - await self.runtime.teardown() - await self.endpoint.teardown() + async with EnvRun(harness=self) as env_run: + return await env_run.run_context( + task, + state, + model=model, + teacher=teacher, + score=score, + ) - async def init_state(self, task: Task) -> State: - return State.for_task(task) + async def run_lifecycle(self, context: Context) -> None: + task = context.task + state = context.state + try: + try: + state.timing.setup.begin() + await self.run_handlers("setup", "rollout", context) + await self.resolve_toolsets(context) + self.validate_extras(state) + state.timing.setup.finish() + state.timing.generation.begin() + await self.run_with_context(context) + state.timing.generation.finish() + await self.run_handlers("update", "rollout", context) + self.validate_extras(state) + if context.score: + context.scoring = True + try: + await score_rollout( + self.owner_signals(), + task, + state, + runtime=context.runtime, + model_client=context.model_client, + teacher=context.teacher, + context=context, + ) + self.validate_extras(state) + finally: + context.scoring = False + except OverlongPromptError as exc: + state.is_truncated = True + state.capture_error(exc) + except Error as exc: + state.capture_error(exc) + except BaseException as exc: + state.capture_error(exc) + finally: + if not state.timing.generation.end: + state.timing.generation.finish() + if not state.timing.cleanup.end: + state.timing.cleanup.begin() + await self.run_handlers("cleanup", "rollout", context) + self.validate_extras(state) + state.timing.cleanup.finish() + if "num_turns" not in state.metrics: + state.metrics["num_turns"] = float(len(state.transcript)) + self.validate_extras(state) + state.assert_serializable() + + def initialize_extras(self, state: State) -> None: + for key, value in self.extras_defaults.items(): + state.extras.setdefault(key, deepcopy(value)) + + def validate_extras(self, state: State) -> None: + if self.extras_schema is None: + return + self.extras_schema.model_validate(state.extras) - @vf.update(priority=-100) - async def render_completion(self, state: State) -> None: - if has_borrowed_trajectory(state): + async def resolve_toolsets(self, context: Context) -> None: + if context.toolsets is None: return - sync_trajectory(state) - - @vf.metric - async def num_turns(self, state: State) -> float: - trajectory = state.get("trajectory") or [] - if not isinstance(trajectory, list): - raise TypeError("state.trajectory must be a list.") - return float(len(trajectory)) - - @vf.stop - async def max_turns_reached(self, state: State) -> bool: - max_turns = state.get_max_turns(self.config.max_turns) - return max_turns > 0 and self.runtime.visible_model_requests(state) >= max_turns - - async def setup_state(self, task: Task, state: State) -> State: - await self.setup_runtime_state(task, state) - await self.setup_model_state(task, state) - await self.resolve_system_prompt(task, state) - await self.setup_tool_state(task, state) - await self.setup_sandbox_state(state) - await self.setup_default_state_fields(state) - return state - - async def setup_runtime_state(self, task: Task, state: State) -> None: - runtime_state = state.runtime_state() - if "max_turns" in task: - runtime_state.setdefault("max_turns", task["max_turns"]) - task.tools_config() - task.toolsets_config() - - async def setup_model_state(self, task: Task, state: State) -> None: - runtime_state = state.runtime_state() - model_handle = self.runtime.model_handle(state) - if model_handle is not None: - if model_handle.model is not None: - runtime_state.setdefault("model", model_handle.model) - if model_handle.client_type is not None: - runtime_state.setdefault("client_type", model_handle.client_type) - if model_handle.sampling_args is not None: - runtime_state.setdefault("sampling_args", model_handle.sampling_args) - task_model = model_config_from_task(task) - model_name = task_model.name or self.model.name - if model_name is not None: - runtime_state.setdefault("model", model_name) - sampling_args = dict(self.model.sampling_args) - sampling_args.update(task_model.sampling_args) - state_sampling_args = runtime_state.get("sampling_args") - if state_sampling_args is not None: - if not isinstance(state_sampling_args, dict): - raise TypeError("state.runtime.sampling_args must be a mapping.") - sampling_args.update(state_sampling_args) - if sampling_args: - runtime_state["sampling_args"] = sampling_args - task_model_client = ( - cast(ModelClient | None, task_model.client_object()) - if task_model.client is not None - else None + task = context.task + state = context.state + await context.toolsets.resolve( + context=self.binding_context(task, state), + resolution_key=f"{state.id}:{task.task_id}", + apply_updates=lambda updates: self.apply_bound_updates(state, updates), ) - model_client = ( - task_model_client if task_model_client is not None else self.model_client + self.validate_extras(state) + state.assert_serializable() + + async def run_with_context(self, context: Context) -> None: + task = context.task + state = context.state + toolsets = context.toolsets + user = context.user + messages = self.initial_messages(task) + if not self.has_model_prompt(messages): + bootstrap_messages = await self.call_user(user, task, state) + if not bootstrap_messages and not state.is_completed: + raise ValueError( + "Task prompt is empty and no user server produced an initial " + "message." + ) + messages = [*messages, *bootstrap_messages] + if state.is_completed: + return + max_turns = self.max_turns(task) + turns = 0 + while max_turns <= 0 or turns < max_turns: + if await self.is_completed(context): + return + sampling = self.sampling_args(task, context.sampling_args) + start = time.time() + response = await context.model_client.get_response( + prompt=messages, + model=context.model, + sampling_args=sampling, + tools=toolsets.tools() if toolsets is not None else None, + state=state, + ) + end = time.time() + turn = Turn( + prompt=messages, + completion=await parse_response_message(response), + tool_calls=list(response.message.tool_calls or []), + response_id=response.id, + model=response.model, + created=response.created, + finish_reason=response.message.finish_reason, + usage=TurnUsage.from_usage(response.usage), + tokens=TurnTokens.from_response( + response.message.tokens, + is_truncated=bool(response.message.is_truncated), + ), + is_truncated=bool(response.message.is_truncated), + timing=TimeSpan(start=start, end=end), + ) + state.transcript.append(turn) + if turn.is_truncated: + state.is_truncated = True + messages = [*messages, *turn.completion] + turns += 1 + if turn.tool_calls: + if toolsets is None: + raise RuntimeError("Model requested tools but no tools are loaded.") + tool_messages, server_messages = await self.call_tools( + task, state, toolsets, turn.tool_calls + ) + turn.tool_results = tool_messages + messages = [*messages, *tool_messages, *server_messages] + continue + user_messages = await self.call_user(user, task, state) + if user_messages: + messages = [*messages, *user_messages] + continue + state.stop("assistant_completed") + return + state.stop("max_turns") + + async def start_env_scope(self) -> None: + async with self._env_servers_lock: + if self._env_scope_count > 0: + self._env_scope_count += 1 + return + taskset = self.taskset + started_toolsets = False + try: + if self._env_toolsets is None: + toolsets = ( + {} + if taskset is None + else { + name: toolset + for name, toolset in taskset.toolsets.items() + if toolset.scope == "env" + } + ) + env_toolsets = MCPToolRegistry(toolsets) + await env_toolsets.__aenter__() + self._env_toolsets = env_toolsets + started_toolsets = True + if self._env_user is None: + user = None if taskset is None else taskset.user + user_toolsets = ( + {} if user is None or user.scope != "env" else {"user": user} + ) + env_user = MCPToolRegistry(user_toolsets) + await env_user.__aenter__() + self._env_user = env_user + self._env_scope_count = 1 + except BaseException: + self._env_scope_count = 0 + try: + if self._env_user is not None: + await self._env_user.__aexit__(None, None, None) + finally: + self._env_user = None + if started_toolsets and self._env_toolsets is not None: + try: + await self._env_toolsets.__aexit__(None, None, None) + finally: + self._env_toolsets = None + raise + + async def stop_env_scope(self, *, force: bool = False) -> None: + async with self._env_servers_lock: + if not force and self._env_scope_count > 1: + self._env_scope_count -= 1 + return + self._env_scope_count = 0 + try: + if self._env_user is not None: + await self._env_user.__aexit__(None, None, None) + finally: + self._env_user = None + try: + if self._env_toolsets is not None: + await self._env_toolsets.__aexit__(None, None, None) + finally: + self._env_toolsets = None + + def rollout_toolsets( + self, runtime: Runtime, user: MCPToolRegistry | None = None + ) -> MCPToolRegistry: + _ = runtime + taskset = self.taskset + parents = [ + parent for parent in (self._env_toolsets, user) if parent is not None + ] + toolsets = ( + {} + if taskset is None + else { + name: toolset + for name, toolset in taskset.toolsets.items() + if toolset.scope == "rollout" + } ) - if ( - model_handle is None - and "client_key" not in runtime_state - and model_client is not None - ): - self.runtime.bind_model_client(state, model_client) - - async def setup_tool_state(self, task: Task, state: State) -> None: - self.runtime.prepare_state(task, state) - self.runtime.validate_bindings(state, allow_unresolved_tool_bindings=True) - await self.runtime.ensure_mcp_tools(state) - self.runtime.refresh_tools(state, validate=False) - - async def setup_sandbox_state(self, state: State) -> None: - await self.runtime.ensure_global_sandboxes(state) - self.runtime.bind_global_sandboxes(state) - - async def setup_default_state_fields(self, state: State) -> None: - state.setdefault("artifacts", {}) - state.setdefault("metrics", {}) - state.setdefault("reward", 0.0) - state.ensure_timing() - - async def resolve_system_prompt(self, task: Task, state: State) -> None: - taskset_system_prompt = ( - self.taskset.system_prompt if self.taskset is not None else [] + return MCPToolRegistry(toolsets, runtime=runtime, parents=parents) + + def rollout_user(self, runtime: Runtime, task: Task) -> MCPToolRegistry: + _ = runtime + if task.user is False: + return MCPToolRegistry({}, expose_tools=False) + taskset = self.taskset + parents = [self._env_user] if self._env_user is not None else [] + user = None if taskset is None else taskset.user + user_toolsets = ( + {} if user is None or user.scope != "rollout" else {"user": user} ) - state["system_prompt"] = resolve_system_prompt( + return MCPToolRegistry(user_toolsets, runtime=runtime, parents=parents) + + @property + def stop_handlers(self) -> list[Handler]: + return self.owner_handlers("stop") + + def owner_handlers(self, kind: str) -> list[Handler]: + taskset_handlers: list[Handler] = [] + if self.taskset is not None: + taskset_handlers = self.taskset.handlers[kind] + return [*taskset_handlers, *self.handlers[kind]] + + def owner_signals(self) -> list[SignalRecord]: + signals = list(getattr(self.taskset, "signals", [])) if self.taskset else [] + seen = {str(signal["name"]) for signal in signals} + for signal in self.signals: + if signal["name"] in seen: + raise ValueError(f"Signal {signal['name']!r} is defined twice.") + signals.append(signal) + return sorted(signals, key=lambda signal: (-signal["priority"], signal["name"])) + + def initial_messages(self, task: Task) -> Messages: + taskset_system_prompt = [] + if self.taskset is not None: + taskset_system_prompt = getattr(self.taskset, "system_prompt", []) + system_prompt = resolve_system_prompt( task=task, taskset_system_prompt=taskset_system_prompt, harness_system_prompt=self.system_prompt, strategy=self.system_prompt_strategy, ) - - async def run_program(self, task: Task, state: State) -> State: - endpoint = self.resolved_endpoint(state) - result = await run_intercepted_program( - self.program, endpoint, self.runtime, task, state - ) - if result is None: - return state - if isinstance(result, State): - return result - if isinstance(result, dict): - state.update(result) - return state - raise TypeError("Harness program must return None, State, or a mapping.") - - def resolved_endpoint(self, state: State) -> Endpoint: - handle = self.runtime.endpoint_handle(state) - if handle is None: - return self.endpoint - runtime = self.runtime.resolved_runtime(handle) - harness = runtime.harness - if harness is None: - raise RuntimeError("Resolved endpoint handle has no live harness.") - endpoint = harness.endpoint - if not isinstance(endpoint, Endpoint): - raise RuntimeError("Resolved endpoint handle has no live endpoint.") - return endpoint - - def compile_program(self, program: ProgramConfig) -> ProgramRunner: - program_data = program.data() - if not isinstance(program_data, dict): - raise TypeError("program must materialize to a mapping.") - kind = program_kind(program_data) - if kind == "base": - sandbox_config = self.program_sandbox_config(program) - validate_program_options(program_data, kind, sandbox_config) - if sandbox_config is not None: - return self.sandbox_base_program(program_data, sandbox_config) - return self.base_program - if kind == "fn": - sandbox_config = self.program_sandbox_config(program) - validate_program_options(program_data, kind, sandbox_config) - fn_ref = program_data["fn"] - if not isinstance(fn_ref, str): - raise TypeError("program.fn must be a string ref.") - if sandbox_config is not None: - return self.sandbox_fn_program( - program_data, sandbox_config, qualified_config_ref(fn_ref) + return [*_MESSAGES_ADAPTER.validate_python(system_prompt), *task.prompt] + + @staticmethod + def has_model_prompt(messages: Messages) -> bool: + return any(getattr(message, "role", None) != "system" for message in messages) + + def max_turns(self, task: Task) -> int: + value = task.max_turns + if value is None: + return self.config.max_turns + return value + + def sampling_args(self, task: Task, sampling_args: SamplingArgs) -> SamplingArgs: + _ = task + return dict(sampling_args) + + async def call_tools( + self, task: Task, state: State, toolsets: MCPToolRegistry, tool_calls + ) -> tuple[list[ToolMessage], Messages]: + toolsets.set_context(self.binding_context(task, state)) + tool_results: list[ToolMessage] = [] + server_messages: Messages = [] + updates: list[BoundUpdate] = [] + for tool_call in tool_calls: + try: + arguments = json_args(tool_call.arguments or "{}") + result = await toolsets.call(tool_call.name, arguments) + updates.extend(result.updates) + call_tool_results, call_messages = self.tool_response_messages( + tool_call.id, result.response ) - fn = import_config_ref(fn_ref) - if not callable(fn): - raise TypeError("program.fn did not resolve to a callable.") - return self.local_callable_program(cast(Handler, fn)) - if kind == "command": - sandbox_config = self.program_sandbox_config(program) - validate_program_options(program_data, kind, sandbox_config) - return self.command_program(program_data, sandbox_config) - raise AssertionError(f"Unhandled program kind: {kind}") - - def local_callable_program(self, fn: Handler) -> ProgramRunner: - async def run(task: Task, state: State) -> ProgramResult: - await self.runtime.setup_rollout(task, state) - result = await maybe_call_with_named_args( - fn, task=task, state=state, runtime=self.runtime, harness=self + except Exception as exc: + if isinstance(exc, ToolError): + content = f"Tool error: {exc}" + else: + content = f"Tool error: {type(exc).__name__}: {exc}" + call_tool_results = [ + ToolMessage(tool_call_id=tool_call.id, content=content) + ] + call_messages = [] + tool_results.extend(call_tool_results) + server_messages.extend(call_messages) + try: + self.apply_bound_updates(state, updates) + except Exception as exc: + message = f"Tool error: {type(exc).__name__}: {exc}" + return ( + [ + ToolMessage(tool_call_id=tool_call.id, content=message) + for tool_call in tool_calls + ], + [], ) - if result is None or isinstance(result, State | dict): - return cast(ProgramResult, result) - raise TypeError("program.fn must return None, State, or a mapping.") - - return run - - async def base_program(self, task: Task, state: State) -> State: - await self.runtime.setup_rollout(task, state) - prompt = normalize_messages( - cast( - Messages, - normalize_prompt(state.get("prompt", []), field_name="state.prompt"), - ), - field_name="state.prompt", + return tool_results, server_messages + + @staticmethod + def tool_response_messages( + tool_call_id: str, response: ServerResponse + ) -> tuple[list[ToolMessage], Messages]: + if response.content is not None: + return [ + ToolMessage(tool_call_id=tool_call_id, content=response.content) + ], list(response.messages) + if not response.messages: + return [ToolMessage(tool_call_id=tool_call_id, content="")], [] + tool_results = [ + message for message in response.messages if isinstance(message, ToolMessage) + ] + extra_messages: Messages = [] + for message in response.messages: + if not isinstance(message, ToolMessage): + extra_messages.append(message) + if not tool_results: + tool_results = [ToolMessage(tool_call_id=tool_call_id, content="")] + return tool_results, extra_messages + + @staticmethod + def binding_context(task: Task, state: State) -> JsonData: + state_data = state.model_dump( + mode="json", + exclude_none=True, + exclude_computed_fields=True, ) - system_prompt = normalize_messages( - state.get("system_prompt", []), field_name="state.system_prompt" + state_data["prompt"] = State.serialized_messages(state.prompt) + state_data["completion"] = State.serialized_messages(state.completion) + state_data["messages"] = State.serialized_messages(state.messages) + return json_data( + { + "task": task.model_dump( + mode="json", exclude_none=True, exclude_defaults=True + ), + "state": state_data, + "extras": state.extras, + }, + context="binding context", ) - messages = [*system_prompt, *prompt] - prompt_messages = [ - message.model_dump(exclude_none=True) for message in messages - ] - def sync_completion() -> list[JsonData]: - rendered_messages = [ - message.model_dump(exclude_none=True) for message in messages - ] - state["completion"] = assistant_completion_from_messages( - prompt_messages, rendered_messages + def apply_tool_result(self, state: State, result: ServerResult) -> ServerResponse: + self.apply_bound_updates(state, list(result.updates)) + return result.response + + def apply_bound_updates(self, state: State, updates: list[BoundUpdate]) -> None: + assignments: list[tuple[str, JsonValue, str]] = [] + for update in updates: + assignments.extend(self.bound_assignments(update)) + for index, (target, _, mode) in enumerate(assignments): + for existing, _, existing_mode in assignments[:index]: + if self.assignment_conflicts(existing, existing_mode, target, mode): + raise ValueError( + f"Conflicting bound state updates: {existing!r} and {target!r}." + ) + for target, value, mode in assignments: + self.apply_assignment(state, target, value, mode) + if assignments: + state.assert_serializable() + self.validate_extras(state) + + @staticmethod + def bound_assignments(update: BoundUpdate) -> list[tuple[str, JsonValue, str]]: + target = update.target + if target.startswith("extras."): + target = f"state.{target}" + if not target.startswith("state."): + raise ValueError( + f"Bound return target {target!r} must start with state. or extras." ) - return rendered_messages - - turn = 0 - max_turns = state.get_max_turns(self.config.max_turns) - while max_turns <= 0 or turn < max_turns: - if await self.runtime.is_completed(task, state): - return state - response = await self.runtime.submit_model_request( - messages, + parts = target.split(".") + if len(parts) < 2: + raise ValueError(f"Bound return target {target!r} is incomplete.") + if update.mode == "extend": + if not isinstance(update.value, list): + raise TypeError(f"Extend target {target!r} requires a list.") + return [(target, deepcopy(update.value), "extend")] + if update.mode == "set": + return [(target, deepcopy(update.value), "set")] + raise ValueError(f"Unknown bound return mode {update.mode!r}.") + + @staticmethod + def apply_assignment( + state: State, target: str, value: JsonValue, mode: str + ) -> None: + parts = target.split(".") + field = parts[1] + if field in {"extras", "metadata", "artifacts"}: + if len(parts) < 3: + raise ValueError(f"Bound return target {target!r} needs a key.") + if field == "extras": + container = state.extras + elif field == "metadata": + container = state.metadata + else: + container = state.artifacts + if mode == "extend": + Harness.extend_mapping_path(container, parts[2:], value) + else: + Harness.set_mapping_path(container, parts[2:], value) + return + if mode != "set": + raise ValueError(f"Bound return target {target!r} only supports set.") + if field == "transcript": + if parts != ["state", "transcript", "last", "reward"]: + raise ValueError( + "state.transcript only supports state.transcript.last.reward." + ) + if not state.transcript: + raise RuntimeError("state.transcript.last.reward requires a turn.") + if isinstance(value, bool) or not isinstance(value, int | float): + raise TypeError("state.transcript.last.reward requires a number.") + state.transcript[-1].reward = float(value) + return + if field == "metrics": + if len(parts) != 3: + raise ValueError(f"Metric target {target!r} must name one metric.") + if isinstance(value, bool) or not isinstance(value, int | float): + raise TypeError(f"Metric target {target!r} requires a number.") + state.metrics[parts[2]] = float(value) + return + if field == "reward": + if len(parts) != 2: + raise ValueError("state.reward does not support nested targets.") + if isinstance(value, bool) or not isinstance(value, int | float): + raise TypeError("state.reward requires a number.") + state.reward = float(value) + return + if field == "is_completed": + if len(parts) != 2 or not isinstance(value, bool): + raise TypeError("state.is_completed requires a boolean.") + state.is_completed = value + return + if field == "is_truncated": + if len(parts) != 2 or not isinstance(value, bool): + raise TypeError("state.is_truncated requires a boolean.") + state.is_truncated = value + return + if field == "stop_condition": + if len(parts) != 2 or not isinstance(value, str): + raise TypeError("state.stop_condition requires a string.") + state.stop(value) + return + raise ValueError(f"Bound return target {target!r} is not writable.") + + @staticmethod + def assignment_conflicts( + left: str, left_mode: str, right: str, right_mode: str + ) -> bool: + if left_mode == right_mode == "extend" and left == right: + return False + return Harness.paths_conflict(left, right) + + @staticmethod + def paths_conflict(left: str, right: str) -> bool: + return ( + left == right + or left.startswith(f"{right}.") + or right.startswith(f"{left}.") + ) + + @staticmethod + def set_mapping_path( + container: JsonData, path: list[str], value: JsonValue + ) -> None: + target = Harness.nested_mapping(container, path[:-1]) + target[path[-1]] = deepcopy(value) + + @staticmethod + def extend_mapping_path( + container: JsonData, path: list[str], value: JsonValue + ) -> None: + if not isinstance(value, list): + raise TypeError(f"Bound extend target {'.'.join(path)!r} requires a list.") + target = Harness.nested_mapping(container, path[:-1]) + existing = target.get(path[-1]) + if existing is None: + target[path[-1]] = deepcopy(value) + return + if not isinstance(existing, list): + raise TypeError(f"Bound extend target {'.'.join(path)!r} is not a list.") + existing.extend(deepcopy(value)) + + @staticmethod + def nested_mapping(container: JsonData, path: list[str]) -> JsonData: + current = container + for part in path: + value = current.get(part) + if value is None: + child: JsonData = {} + current[part] = child + current = child + continue + if not isinstance(value, dict): + raise TypeError(f"Bound return path {part!r} traverses a non-object.") + current = value + return current + + async def call_user( + self, user: MCPToolRegistry | None, task: Task, state: State + ) -> Messages: + if task.user is False: + return [] + if user is None or not user.has_hidden("respond"): + if task.user is True: + raise ValueError("Task requires a user server, but none is loaded.") + return [] + user.set_context(self.binding_context(task, state)) + result = await user.call_hidden("respond", {}) + self.apply_bound_updates(state, list(result.updates)) + state.assert_serializable() + messages = list(result.response.messages) + if result.response.content is not None: + messages.insert(0, UserMessage(content=result.response.content)) + return messages + + async def is_completed(self, context: Context) -> bool: + task = context.task + state = context.state + if state.is_completed: + return True + for handler in self.stop_handlers: + if await self.call_handler( + handler, task, state, - tool_defs=self.runtime.tool_defs(state), - ) - turn += 1 - messages.extend(await parse_response_message(response)) - rendered_messages = sync_completion() - tool_calls = list(response.message.tool_calls or []) - if not tool_calls: - user_messages = await self.runtime.user_messages( - task, state, transcript=rendered_messages - ) - if user_messages: - messages.extend( - normalize_messages( - cast(Messages, user_messages), - field_name="user_messages", - ) - ) - sync_completion() - continue - state._set_stop_condition("no_tools") - return state - callable_tools = state.get_tools() - - async def call_tool(tool_call) -> ToolMessage: - content: MessageContent - try: - name = tool_call.name - result = await maybe_call_with_named_args( - callable_tools[name], **json_args(tool_call.arguments) - ) - content = ( - cast(MessageContent, result) - if is_valid_tool_content_parts(result) - else str(result) - ) - except Exception as e: - content = tool_error_content(e) - return ToolMessage(tool_call_id=tool_call.id, content=content) + context=context, + runtime=context.runtime, + toolsets=context.toolsets, + user=context.user, + model=context.model_client, + model_name=context.model, + teacher=context.teacher, + teacher_name=context.teacher.config.model + if context.teacher is not None + else None, + ): + state.stop(getattr(handler, "__name__", "stop")) + return True + return False - messages.extend( - await asyncio.gather( - *(call_tool(tool_call) for tool_call in tool_calls) - ) + async def run_handlers( + self, + kind: str, + stage: str, + context: Context, + ) -> None: + task = context.task + state = context.state + for handler in self.owner_handlers(kind): + handler_stage = getattr(handler, f"{kind}_stage", "rollout") + if handler_stage != stage: + continue + binding_context = self.binding_context(task, state) + if context.toolsets is not None: + context.toolsets.set_context(binding_context) + if context.user is not None: + context.user.set_context(binding_context) + result = await self.call_handler( + handler, + task, + state, + context=context, + runtime=context.runtime, + toolsets=context.toolsets, + user=context.user, + model=context.model_client, + model_name=context.model, + teacher=context.teacher, + teacher_name=context.teacher.config.model + if context.teacher is not None + else None, ) - sync_completion() - if await self.runtime.is_completed(task, state): - return state - return state - - def command_program( - self, program: ConfigData, sandbox_config: SandboxConfig | None - ) -> ProgramRunner: - async def run(task: Task, state: State) -> State: - runtime = self.runtime - merged_program = merge_task_program(program, task, kind="command") - if sandbox_config is not None: - return await run_sandbox_command( - self.prepare_sandbox_program(merged_program, state), - self.prepare_sandbox_config( - merge_task_sandbox(sandbox_config, task), program - ), - task, - state, - runtime, - ) - await runtime.setup_rollout(task, state) - return await run_local_command(merged_program, task, state, runtime) - - return run - - def sandbox_base_program( - self, program: ConfigData, sandbox_config: SandboxConfig - ) -> ProgramRunner: - async def run(task: Task, state: State) -> State: - merged_program = merge_task_program(program, task, kind="base") - return await run_sandbox_python_program( - program=self.prepare_sandbox_program(merged_program, state), - sandbox_config=self.prepare_sandbox_config( - merge_task_sandbox(sandbox_config, task), merged_program - ), - task=task, - state=state, - runtime=self.runtime, - mode="base", - fn_ref=None, - max_turns=state.get_max_turns(self.config.max_turns), + if result is None: + continue + raise TypeError( + f"{kind} handler {getattr(handler, '__name__', handler)!r} must mutate " + f"state in place and return None, not {type(result).__name__}." ) - return run - - def sandbox_fn_program( + async def call_handler( self, - program: ConfigData, - sandbox_config: SandboxConfig, - fn_ref: str, - ) -> ProgramRunner: - async def run(task: Task, state: State) -> State: - merged_program = merge_task_program(program, task, kind="fn") - return await run_sandbox_python_program( - program=self.prepare_sandbox_program(merged_program, state), - sandbox_config=self.prepare_sandbox_config( - merge_task_sandbox(sandbox_config, task), merged_program - ), - task=task, - state=state, - runtime=self.runtime, - mode="fn", - fn_ref=fn_ref, - max_turns=state.get_max_turns(self.config.max_turns), - ) + handler: Handler, + task: Task, + state: State, + runtime: Runtime | None = None, + **extra: object, + ) -> object: + return await maybe_call_with_named_args( + handler, + task=task, + state=state, + extras=state.extras, + transcript=state.transcript, + completion=state.completion, + metrics=state.metrics, + reward=state.reward, + prompt=state.prompt if state.transcript else task.prompt, + example_id=task.row_id, + harness=self, + context=extra.pop("context", None), + runtime=runtime, + toolsets=extra.pop("toolsets", None), + user=extra.pop("user", None), + **extra, + ) - return run - - def program_sandbox_config(self, program: ProgramConfig) -> SandboxConfig | None: - sandbox = program.resolve().sandbox - if sandbox is False: - return None - if self.sandbox is None: - return None - if sandbox is None and self.config.sandbox is None: - return None - validate_program_sandbox_scope(self.sandbox) - return self.sandbox - - def prepare_sandbox_program(self, program: ConfigData, state: State) -> ConfigData: - if "mcp" in program_channels(program): - endpoint_root_url = state.get("endpoint_root_url") - if not isinstance(endpoint_root_url, str): - raise RuntimeError("MCP program tools require an active endpoint.") - api_key_var = state.get("endpoint_api_key_var") - if not isinstance(api_key_var, str): - api_key_var = "OPENAI_API_KEY" - return proxy_program( - program, - tool_base_url=f"{endpoint_root_url.rstrip('/')}/vf/tools", - tool_auth_var=api_key_var, - ) - return program - - def prepare_sandbox_config( - self, sandbox_config: SandboxConfig, program: ConfigData - ) -> SandboxConfig: - config = sandbox_config.data() - if "mcp" in program_channels(program): - config = proxy_sandbox(config) - if program_kind(program) in {"base", "fn"}: - config = python_program_sandbox(config) - return SandboxConfig.model_validate(config) + async def teardown(self) -> None: + for handler in self.handlers["teardown"]: + result = handler() + if inspect.isawaitable(result): + await result + + async def close(self) -> None: + try: + await self.stop_env_scope(force=True) + finally: + await self.teardown() diff --git a/verifiers/v1/interception.py b/verifiers/v1/interception.py new file mode 100644 index 0000000000..5e7244e65c --- /dev/null +++ b/verifiers/v1/interception.py @@ -0,0 +1,164 @@ +from __future__ import annotations + +import secrets +from abc import ABC, abstractmethod +from collections.abc import Awaitable, Callable, Sequence +from dataclasses import dataclass + +from aiohttp import web +from pydantic import BaseModel, Field + +from verifiers.types import Messages, Response, Tool +from verifiers.utils.response_utils import parse_response_message + +from .state import State, Turn, TurnTokens, TurnUsage +from .task import Task +from .types import JsonData, JsonValue, Context +from .utils.json_utils import json_data + +StopCheck = Callable[[], Awaitable[str | None]] + + +@dataclass(frozen=True) +class ProtocolRoute: + method: str + path: str + + +class InterceptedRequest(BaseModel, extra="forbid"): + protocol: str + prompt: Messages + model: str | None = None + sampling_args: dict[str, JsonValue] = Field(default_factory=dict) + tools: list[Tool] | None = None + body: JsonData = Field(default_factory=dict) + + +class EndpointProtocol(ABC): + name: str + routes: Sequence[ProtocolRoute] + + def env(self, *, base_url: str, api_key: str, model: str) -> dict[str, str]: + _ = base_url, api_key, model + return {} + + @abstractmethod + async def parse( + self, request: web.Request, body: JsonData + ) -> InterceptedRequest: ... + + @abstractmethod + def serialize( + self, response: Response, request: InterceptedRequest + ) -> JsonData: ... + + def serialize_error(self, error: BaseException) -> tuple[int, JsonData]: + return 502, {"error": str(error)} + + +class InterceptionServer: + def __init__( + self, + context: Context, + task: Task, + state: State, + *, + protocols: Sequence[EndpointProtocol] | None = None, + stop_check: StopCheck | None = None, + ): + if protocols is None: + from .protocols import default_protocols + + protocols = default_protocols() + self.context = context + self.task = task + self.state = state + self.protocols = list(protocols) + self.stop_check = stop_check + self.secret = secrets.token_urlsafe(16) + self.port = 0 + self.runner: web.AppRunner | None = None + + async def __aenter__(self) -> "InterceptionServer": + app = web.Application() + for protocol in self.protocols: + for route in protocol.routes: + app.router.add_route(route.method, route.path, self.handler(protocol)) + self.runner = web.AppRunner(app) + await self.runner.setup() + site = web.TCPSite(self.runner, "127.0.0.1", 0) + await site.start() + sockets = getattr(site, "_server").sockets + self.port = int(sockets[0].getsockname()[1]) + return self + + async def __aexit__(self, *exc: object) -> None: + if self.runner is not None: + await self.runner.cleanup() + + def env(self, *, base_url: str, model: str) -> dict[str, str]: + env: dict[str, str] = {} + for protocol in self.protocols: + env.update( + protocol.env(base_url=base_url, api_key=self.secret, model=model) + ) + return env + + def handler(self, protocol: EndpointProtocol): + async def handle(request: web.Request) -> web.Response: + if request.headers.get("Authorization") != f"Bearer {self.secret}": + return web.json_response({"error": "unauthorized"}, status=401) + try: + body = await json_body(request) + intercepted = await protocol.parse(request, body) + stop_condition = await self.check_stop() + if stop_condition is not None: + self.state.stop(stop_condition) + return web.json_response( + {"error": f"rollout stopped: {stop_condition}"}, + status=400, + ) + response = await self.context.model_client.get_response( + prompt=intercepted.prompt, + model=intercepted.model or self.context.model, + sampling_args={ + **self.context.sampling_args, + **intercepted.sampling_args, + }, + tools=intercepted.tools, + state=self.state, + ) + turn = Turn( + prompt=intercepted.prompt, + completion=await parse_response_message(response), + tool_calls=list(response.message.tool_calls or []), + response_id=response.id, + model=response.model, + created=response.created, + finish_reason=response.message.finish_reason, + usage=TurnUsage.from_usage(response.usage), + tokens=TurnTokens.from_response( + response.message.tokens, + is_truncated=bool(response.message.is_truncated), + ), + is_truncated=bool(response.message.is_truncated), + ) + self.state.transcript.append(turn) + if turn.is_truncated: + self.state.is_truncated = True + return web.json_response(protocol.serialize(response, intercepted)) + except BaseException as exc: + status, payload = protocol.serialize_error(exc) + return web.json_response(payload, status=status) + + return handle + + async def check_stop(self) -> str | None: + if self.stop_check is None: + return None + return await self.stop_check() + + +async def json_body(request: web.Request) -> JsonData: + body = await request.json() + return json_data(body, context="Protocol request body") diff --git a/verifiers/v1/lifecycle.py b/verifiers/v1/lifecycle.py new file mode 100644 index 0000000000..36f124c127 --- /dev/null +++ b/verifiers/v1/lifecycle.py @@ -0,0 +1,217 @@ +from __future__ import annotations + +import asyncio +import uuid +from contextlib import AsyncExitStack +from typing import TYPE_CHECKING + +from verifiers.types import RolloutInput +from verifiers.utils.async_utils import maybe_retry + +from .state import State +from .task import Task +from .types import ModelConfig +from .utils.json_utils import json_data + +if TYPE_CHECKING: + from .env import Env + from .harness import Harness + from .taskset import Taskset + + +class EnvRun: + def __init__( + self, + *, + env: Env | None = None, + harness: Harness | None = None, + ) -> None: + if env is None and harness is None: + raise TypeError("EnvRun requires an env or harness.") + if env is not None and harness is not None: + raise TypeError("EnvRun accepts env or harness, not both.") + if harness is None: + assert env is not None + harness = env.harness + self.env = env + self.harness = harness + self.taskset: Taskset | None = None if env is None else env.taskset + self._entered = False + + async def __aenter__(self) -> "EnvRun": + await self.harness.start_env_scope() + self._entered = True + return self + + async def __aexit__(self, *exc: object) -> None: + self._entered = False + await self.harness.stop_env_scope() + + def to_task(self, input: RolloutInput | Task | str) -> Task: + if isinstance(input, str): + return Task(prompt=input) + if isinstance(input, Task): + if self.taskset is None: + return input + return self.taskset.to_task(input) + if isinstance(input, dict): + row = json_data(input) + if self.taskset is None: + return Task.model_validate(row) + return self.taskset.to_task(row) + raise TypeError("Rollout input must be a Task, string prompt, or mapping.") + + async def run_rollout( + self, + input: RolloutInput | Task | str, + *, + model: ModelConfig | str, + teacher: ModelConfig | str | None = None, + state: State | None = None, + score: bool = True, + max_retries: int = 0, + ) -> State: + task = self.to_task(input) + model_config = type(self).normalize_model(model) + teacher_config = ( + type(self).normalize_model(teacher) if teacher is not None else None + ) + + async def attempt() -> State: + if state is None: + rollout_state = State(task_id=task.task_id) + elif max_retries > 0: + rollout_state = state.model_copy(deep=True) + else: + rollout_state = state + return await self.run_context( + task, + rollout_state, + model=model_config, + teacher=teacher_config, + score=score, + ) + + return await maybe_retry(attempt, max_retries=max_retries)() + + async def run_context( + self, + task: Task | str, + state: State | None = None, + *, + model: ModelConfig | str, + teacher: ModelConfig | str | None = None, + score: bool = False, + ) -> State: + if not self._entered: + raise RuntimeError("EnvRun must be entered before running a context.") + task = self.to_task(task) + model_config = type(self).normalize_model(model) + teacher_config = ( + type(self).normalize_model(teacher) if teacher is not None else None + ) + state = state or State(task_id=task.task_id) + state.task_id = task.task_id + state.model = state.model or model_config + state.teacher = state.teacher or teacher_config + self.harness.initialize_extras(state) + + async with self.harness.runtime_provider_for(task).create_runtime() as runtime: + async with AsyncExitStack() as stack: + user = await stack.enter_async_context( + self.harness.rollout_user(runtime, task) + ) + toolsets = await stack.enter_async_context( + self.harness.rollout_toolsets(runtime, user) + ) + async with self.harness.open_context( + task=task, + state=state, + model=model_config, + teacher=teacher_config, + runtime=runtime, + toolsets=toolsets, + user=user, + score=score, + ) as context: + if toolsets is not None: + toolsets.set_visibility( + toolsets=task.toolsets, + tools=task.tools, + ) + await self.harness.run_lifecycle(context) + return state + + async def group(self, rows: list[RolloutInput]) -> "Group": + if self.env is None or self.taskset is None: + raise RuntimeError("Grouped rollouts require an Env.") + if not rows: + raise ValueError("Group requires at least one row.") + base_task = self.taskset.to_task(json_data(rows[0])) + tasks, states = await self.taskset.init_group(base_task, len(rows)) + group_id = str(rows[0].get("example_id") or uuid.uuid4().hex) + for state in states: + state.group_id = state.group_id or group_id + return Group(env_run=self, tasks=tasks, states=states) + + @staticmethod + def normalize_model(model: ModelConfig | str) -> ModelConfig: + if isinstance(model, str): + return ModelConfig(model=model) + return model + + +class Group: + def __init__( + self, + *, + env_run: EnvRun, + tasks: list[Task], + states: list[State], + ) -> None: + self.env_run = env_run + self.tasks = tasks + self.states = states + + async def run( + self, + *, + model: ModelConfig | str, + teacher: ModelConfig | str | None = None, + max_retries: int = 0, + ) -> list[State]: + model_config = EnvRun.normalize_model(model) + teacher_config = ( + EnvRun.normalize_model(teacher) if teacher is not None else None + ) + self.states = list( + await asyncio.gather( + *[ + self.env_run.run_rollout( + task, + model=model_config, + teacher=teacher_config, + state=state, + max_retries=max_retries, + ) + for task, state in zip(self.tasks, self.states, strict=True) + ] + ) + ) + return await self.score(model=model_config, teacher=teacher_config) + + async def score( + self, + *, + model: ModelConfig | None = None, + teacher: ModelConfig | None = None, + ) -> list[State]: + if self.env_run.env is None: + raise RuntimeError("Grouped scoring requires an Env.") + self.states = await self.env_run.env.score_group( + self.tasks, + self.states, + model=model, + teacher=teacher, + ) + return self.states diff --git a/verifiers/v1/loaders.py b/verifiers/v1/loaders.py new file mode 100644 index 0000000000..88331d2cf3 --- /dev/null +++ b/verifiers/v1/loaders.py @@ -0,0 +1,300 @@ +from __future__ import annotations + +import importlib +import importlib.util +import inspect +import sys +from collections.abc import Mapping +from types import ModuleType, UnionType +from typing import TypeAlias, TypeVar, Union, cast, get_args, get_origin, get_type_hints + +from pydantic import BaseModel + +from .env import Env, EnvConfig +from .harness import Harness, HarnessConfig +from .taskset import Taskset, TasksetConfig +from .utils.config_utils import coerce_config, explicit_config_data + +ConfigMapping: TypeAlias = Mapping[str, object] +EnvConfigLoadData: TypeAlias = dict[str, object] +EnvConfigChildInput: TypeAlias = ConfigMapping | EnvConfigLoadData +EnvConfigInput: TypeAlias = BaseModel | ConfigMapping +ConfigBaseT = TypeVar("ConfigBaseT", bound=BaseModel) + +FACTORY_MODULES = { + "load_taskset": "taskset", + "load_harness": "harness", +} + + +def env_module_name(env_id: str) -> str: + return env_id.replace("-", "_").split("/")[-1] + + +def import_env_module(env_id: str) -> ModuleType: + return importlib.import_module(env_module_name(env_id)) + + +def caller_module() -> ModuleType: + frame = inspect.currentframe() + try: + if frame is None or frame.f_back is None or frame.f_back.f_back is None: + raise RuntimeError("Could not resolve caller module.") + module_name = frame.f_back.f_back.f_globals.get("__name__") + if not isinstance(module_name, str): + raise RuntimeError("Caller module has no __name__.") + module = sys.modules.get(module_name) + if not isinstance(module, ModuleType): + raise RuntimeError(f"Caller module {module_name!r} is not loaded.") + return module + finally: + del frame + + +def load_taskset( + env_id: str | None = None, + *, + config: TasksetConfig | ConfigMapping | None = None, +) -> Taskset: + module = caller_module() if env_id is None else import_env_module(env_id) + return load_taskset_from_module(module, config=config) + + +def load_harness( + env_id: str | None = None, + *, + config: HarnessConfig | ConfigMapping | None = None, +) -> Harness: + module = caller_module() if env_id is None else import_env_module(env_id) + return load_harness_from_module(module, config=config) + + +def load_environment(env_id: str, **env_args: object) -> Env: + return load_environment_from_components(import_env_module(env_id), env_args) + + +def load_taskset_from_module( + module: ModuleType, + *, + config: TasksetConfig | ConfigMapping | None = None, +) -> Taskset: + source_module_name = module.__name__ + module = factory_module(module, "load_taskset") + factory = getattr(module, "load_taskset", None) + if factory is None: + loader_id = child_loader_id(config) + if loader_id is not None and not matches_loader(source_module_name, loader_id): + return load_taskset(loader_id, config=config) + raise AttributeError( + f"Module '{module.__name__}' does not expose load_taskset, and " + "config.id is not set to a taskset loader package." + ) + config_type = factory_config_type(module, "load_taskset", TasksetConfig) + if config_type is None: + raise TypeError(f"{module.__name__}.load_taskset must accept config.") + taskset = factory(config=coerce_config(config_type, config)) + if not isinstance(taskset, Taskset): + raise TypeError(f"{module.__name__}.load_taskset must return a Taskset.") + return taskset + + +def load_harness_from_module( + module: ModuleType, + *, + config: HarnessConfig | ConfigMapping | None = None, +) -> Harness: + source_module_name = module.__name__ + module = factory_module(module, "load_harness") + factory = getattr(module, "load_harness", None) + if factory is None: + loader_id = child_loader_id(config) + if loader_id is not None: + if matches_loader(source_module_name, loader_id): + raise AttributeError( + f"Module '{module.__name__}' does not expose load_harness." + ) + return load_harness(loader_id, config=config) + return Harness(config=coerce_config(HarnessConfig, config)) + config_type = factory_config_type(module, "load_harness", HarnessConfig) + if config_type is None: + raise TypeError(f"{module.__name__}.load_harness must accept config.") + harness = factory(config=coerce_config(config_type, config)) + if not isinstance(harness, Harness): + raise TypeError(f"{module.__name__}.load_harness must return a Harness.") + return harness + + +def load_environment_from_components( + module: ModuleType, + env_args: dict[str, object], +) -> Env: + extra_args = set(env_args) - {"config"} + if extra_args: + raise TypeError( + "Default Taskset/Harness environment loading only accepts config; " + f"got {sorted(extra_args)}." + ) + config_input = env_args.get("config", {}) + if not isinstance(config_input, BaseModel | Mapping): + raise TypeError("config must be a mapping or EnvConfig.") + config = load_env_config(module, EnvConfig, cast(EnvConfigInput, config_input)) + return Env( + taskset=load_taskset_from_module(module, config=config.taskset), + harness=load_harness_from_module(module, config=config.harness), + runtime=config.runtime, + advantage=config.advantage, + ) + + +def load_env_config( + module: ModuleType, + config_type: type[EnvConfig], + value: EnvConfigInput, + *, + child_types: Mapping[str, type[BaseModel]] | None = None, +) -> EnvConfig: + data: EnvConfigLoadData + if isinstance(value, config_type): + data = dict(explicit_config_data(value)) + elif isinstance(value, BaseModel): + raise TypeError( + f"load_environment config must be {config_type.__name__}; " + f"got {type(value).__name__}." + ) + elif not isinstance(value, Mapping): + raise TypeError("load_environment config must be a mapping or EnvConfig.") + else: + data = dict(value) + resolved_child_types = ( + env_config_child_types(module, config_type, data) + if child_types is None + else child_types + ) + for field_name, child_type in resolved_child_types.items(): + if field_name not in data: + data[field_name] = child_type() + continue + child = data[field_name] + if isinstance(child, child_type): + continue + if child is None: + raise TypeError(f"config.{field_name} cannot be None.") + if not isinstance(child, BaseModel | Mapping): + raise TypeError(f"config.{field_name} must be a mapping or config object.") + data[field_name] = child_type.model_validate( + explicit_config_data(cast(EnvConfigInput, child)) + ) + return config_type.model_validate(data) + + +def env_config_child_types( + module: ModuleType, + config_type: type[EnvConfig], + value: EnvConfigChildInput | None = None, +) -> dict[str, type[BaseModel]]: + child_types: dict[str, type[BaseModel]] = {} + for field_name, factory_name, base_type in ( + ("taskset", "load_taskset", TasksetConfig), + ("harness", "load_harness", HarnessConfig), + ): + factory_type = factory_config_type(module, factory_name, base_type) + child_config = value.get(field_name) if value is not None else None + if factory_type is None and child_config_requires_loader_type( + child_config, base_type + ): + loader_id = child_loader_id(child_config) + if loader_id is not None and not matches_loader(module.__name__, loader_id): + factory_type = factory_config_type( + import_env_module(loader_id), factory_name, base_type + ) + if factory_type is not None: + child_types[field_name] = factory_type + else: + child_types[field_name] = base_type + return child_types + + +def child_config_requires_loader_type( + config: object, + base_type: type[BaseModel], +) -> bool: + if not isinstance(config, Mapping): + return False + base_fields = set(base_type.model_fields) + return bool(set(config) - base_fields) + + +def child_loader_id(config: object) -> str | None: + if isinstance(config, BaseModel): + value = config.__dict__.get("id") + elif isinstance(config, Mapping): + value = dict(config).get("id") + else: + return None + if value is None: + return None + if not isinstance(value, str) or not value: + raise TypeError("config.id must be a non-empty string.") + return value + + +def matches_loader(module_name: str, loader_id: str) -> bool: + loader_module_name = env_module_name(loader_id) + return module_name == loader_module_name or module_name.startswith( + f"{loader_module_name}." + ) + + +def factory_config_type( + module: ModuleType, + factory_name: str, + base_type: type[ConfigBaseT], +) -> type[ConfigBaseT] | None: + module = factory_module(module, factory_name) + factory = getattr(module, factory_name, None) + if factory is None: + return None + signature = inspect.signature(factory) + if "config" not in signature.parameters: + raise TypeError(f"{module.__name__}.{factory_name} must accept config.") + try: + annotation = get_type_hints(factory).get( + "config", signature.parameters["config"].annotation + ) + except Exception: + annotation = signature.parameters["config"].annotation + return config_type_from_annotation( + annotation, + base_type, + f"{module.__name__}.{factory_name}.config", + ) + + +def factory_module(module: ModuleType, factory_name: str) -> ModuleType: + if getattr(module, factory_name, None) is not None: + return module + child_name = FACTORY_MODULES.get(factory_name) + if child_name is None or not hasattr(module, "__path__"): + return module + module_name = f"{module.__name__}.{child_name}" + spec = importlib.util.find_spec(module_name) + if spec is None: + return module + return importlib.import_module(module_name) + + +def config_type_from_annotation( + annotation: object, + base_type: type[ConfigBaseT], + context: str, +) -> type[ConfigBaseT]: + if annotation is inspect.Parameter.empty: + raise TypeError(f"{context} must be annotated.") + origin = get_origin(annotation) + if origin in (Union, UnionType): + args = [arg for arg in get_args(annotation) if arg is not type(None)] + if len(args) == 1: + annotation = args[0] + if isinstance(annotation, type) and issubclass(annotation, base_type): + return annotation + raise TypeError(f"{context} must be a {base_type.__name__} subclass.") diff --git a/verifiers/v1/mcp.py b/verifiers/v1/mcp.py new file mode 100644 index 0000000000..c4803a4509 --- /dev/null +++ b/verifiers/v1/mcp.py @@ -0,0 +1,883 @@ +from __future__ import annotations + +import asyncio +import contextlib +from collections.abc import Mapping +from contextlib import AsyncExitStack +from copy import deepcopy +from dataclasses import dataclass +import inspect +import json +import os +import random +import socket +import sys +from collections.abc import Callable +from typing import TYPE_CHECKING, cast +import urllib.error +import urllib.request + +from pydantic import Field +from pydantic import ValidationError +from verifiers.errors import ToolError +from verifiers.types import MessageContent, Messages, Tool + +from .config import Config +from .runtime import Runtime, SubprocessRuntime, make_runtime_provider +from .toolset import ( + ServerConfig, + ToolBinding, + ToolSpec, + Toolset, +) +from .types import JsonData, JsonValue +from .utils.config_utils import import_config_ref, registered_config_type +from .utils.json_utils import json_data, json_value + +if TYPE_CHECKING: + from mcp.client.session import ClientSession + from .task import TaskVisibility + +_BINDINGS_TOOL = "__vf_bindings" + + +@dataclass(frozen=True) +class BoundUpdate: + target: str + value: JsonValue + mode: str = "set" + + +class ServerResponse(Config): + content: MessageContent | None = None + messages: Messages = Field(default_factory=list) + + +@dataclass(frozen=True) +class ServerResult: + response: ServerResponse + updates: tuple[BoundUpdate, ...] = () + value: JsonValue = "" + + +@dataclass(frozen=True) +class ToolDispatch: + session: "ClientSession" + toolset_name: str + raw_name: str + binding: ToolBinding + dynamic_name: str | None = None + name_arg: str = "name" + input_arg: str = "input" + + +class MCPToolRegistry: + def __init__( + self, + servers: Mapping[str, ServerConfig], + *, + runtime: Runtime | None = None, + parents: list["MCPToolRegistry"] | None = None, + expose_tools: bool = True, + ) -> None: + self.servers = dict(servers) + self.runtime = runtime + self.parents = parents or [] + self.expose_tools = expose_tools + self._stack = AsyncExitStack() + self._dispatch: dict[str, ToolDispatch] = {} + self._dynamic_dispatch: dict[str, ToolDispatch] = {} + self._tools: list[Tool] = [] + self._dynamic_tools: list[Tool] = [] + self._context: JsonData = {} + self._toolsets_visibility: TaskVisibility | None = None + self._tools_visibility: TaskVisibility | None = None + self._resolution_key: str | None = None + + def set_context(self, context: JsonData) -> None: + self._context = context + for parent in self.parents: + parent.set_context(context) + + def set_visibility( + self, + *, + toolsets: "TaskVisibility | None", + tools: "TaskVisibility | None", + ) -> None: + self._toolsets_visibility = toolsets + self._tools_visibility = tools + for parent in self.parents: + parent.set_visibility(toolsets=toolsets, tools=tools) + + async def __aenter__(self) -> "MCPToolRegistry": + if not self.servers: + return self + from mcp.client.session import ClientSession + + try: + for toolset_name, server in self.servers.items(): + if not isinstance(toolset_name, str) or not toolset_name: + raise TypeError("MCP server names must be non-empty strings.") + seen_tools: set[str] = set() + read, write = await self._stack.enter_async_context( + self.open_server(toolset_name, server) + ) + session = await self._stack.enter_async_context( + ClientSession(read, write) + ) + await session.initialize() + bindings = await server_bindings(session) + for raw_tool in (await session.list_tools()).tools: + tool_name = getattr(raw_tool, "name", "") + if not isinstance(tool_name, str) or not tool_name: + raise TypeError("MCP tools require a non-empty name.") + if tool_name == _BINDINGS_TOOL: + continue + seen_tools.add(tool_name) + exposed_name = f"{toolset_name}_{tool_name}" + if exposed_name in self._dispatch: + raise ValueError(f"MCP tool {exposed_name!r} is defined twice.") + raw_schema = getattr(raw_tool, "inputSchema", None) + schema: dict[str, object] = ( + dict(raw_schema) + if isinstance(raw_schema, dict) + else {"type": "object", "properties": {}} + ) + binding = bindings.get( + tool_name, + ToolBinding(args={}, sets={}, extends={}, hidden=False), + ) + visible_schema = model_visible_schema(tool_name, schema, binding) + if tool_visible(server, tool_name) and not binding.hidden: + self._tools.append( + Tool( + name=exposed_name, + description=str( + getattr(raw_tool, "description", "") or "" + ), + parameters=visible_schema, + ) + ) + self._dispatch[exposed_name] = ToolDispatch( + session=session, + toolset_name=toolset_name, + raw_name=tool_name, + binding=binding, + ) + missing = sorted(set(bindings) - seen_tools) + if missing: + raise ValueError( + f"Toolset {toolset_name!r} binds unknown tools: " + f"{', '.join(missing)}." + ) + except BaseException: + await self._stack.aclose() + self._dispatch.clear() + self._dynamic_dispatch.clear() + self._tools.clear() + self._dynamic_tools.clear() + raise + return self + + @contextlib.asynccontextmanager + async def open_server(self, name: str, server: ServerConfig): + if server.placement == "remote": + if server.url is None: + raise ValueError("Remote server requires url.") + async with connect_streamable_http(server.url, server.headers) as ( + read, + write, + ): + yield read, write + return + + owns_runtime = False + runtime = self.runtime + if server.placement == "dedicated": + runtime_config = server.runtime + if runtime_config is None: + raise ValueError("Dedicated server placement requires runtime config.") + runtime = make_runtime_provider(runtime_config).create_runtime() + owns_runtime = True + await runtime.start() + elif runtime is None: + raise ValueError("Colocated server placement requires a running runtime.") + + assert runtime is not None + try: + port = free_port() + env = _server_env(server) + env["MCP_PORT"] = str(port) + log = f"vf_server_{name}.log" + await runtime.run_background( + _server_command( + server, + name, + in_runtime=not isinstance(runtime, SubprocessRuntime), + ), + env=env, + log=log, + ) + base_url = await runtime.public_url(port) + if base_url is None: + base_url = f"http://127.0.0.1:{port}" + url = f"{base_url.rstrip('/')}/mcp" + await wait_for_http(url, timeout_seconds=server.startup_timeout_seconds) + async with connect_streamable_http(url, server.headers) as (read, write): + yield read, write + finally: + if owns_runtime: + await runtime.stop() + + async def __aexit__(self, *exc: object) -> None: + await self._stack.aclose() + + def tools(self, *, include_hidden: bool = False) -> list[Tool] | None: + tools: list[Tool] = [] + for parent in self.parents: + tools.extend(parent.tools(include_hidden=include_hidden) or []) + if self.expose_tools or include_hidden: + tools.extend( + tool + for tool in self._tools + if include_hidden or self.tool_allowed(tool.name) + ) + tools.extend( + tool + for tool in self._dynamic_tools + if include_hidden or self.tool_allowed(tool.name) + ) + return tools or None + + def has_tool(self, name: str) -> bool: + return ( + name in self._dispatch + or name in self._dynamic_dispatch + or any(parent.has_tool(name) for parent in self.parents) + ) + + def hidden_matches(self, raw_name: str) -> list[str]: + matches = [ + exposed_name + for exposed_name, dispatch in self._dispatch.items() + if dispatch.raw_name == raw_name + ] + for parent in self.parents: + matches.extend(parent.hidden_matches(raw_name)) + return matches + + async def call(self, name: str, arguments: JsonData) -> ServerResult: + if name not in self._dispatch and name not in self._dynamic_dispatch: + for parent in self.parents: + if parent.has_tool(name): + return await parent.call(name, arguments) + raise ToolError(f"Unknown MCP tool {name!r}.") + if not self.tool_allowed(name): + raise ToolError(f"MCP tool {name!r} is disabled for this task.") + return await self._call(name, arguments) + + async def _call(self, name: str, arguments: JsonData) -> ServerResult: + dispatch = self._dynamic_dispatch.get(name) or self._dispatch[name] + payload: JsonData + if dispatch.dynamic_name is None: + payload = dict(arguments) + else: + payload = { + dispatch.name_arg: dispatch.dynamic_name, + dispatch.input_arg: dict(arguments), + } + for arg_name, source in dispatch.binding.args.items(): + if source.startswith("resources."): + continue + if arg_name in payload: + raise ToolError( + f"MCP tool {name!r} argument {arg_name!r} is bound and cannot " + "be provided by the model." + ) + payload[arg_name] = resolve_binding(self._context, source) + result = await dispatch.session.call_tool(dispatch.raw_name, payload) + if bool(getattr(result, "isError", False)): + raise ToolError(str(mcp_content_value(getattr(result, "content", [])))) + content = mcp_content_value(getattr(result, "content", [])) + return split_result(content, dispatch.binding) + + async def call_hidden(self, raw_name: str, arguments: JsonData) -> ServerResult: + local_matches = [ + exposed_name + for exposed_name, dispatch in self._dispatch.items() + if dispatch.raw_name == raw_name + ] + parent_matches = [ + parent for parent in self.parents if parent.hidden_matches(raw_name) + ] + matches = [*local_matches, *parent_matches] + if not matches: + raise ToolError(f"Unknown hidden MCP tool {raw_name!r}.") + if len(matches) > 1: + raise ToolError(f"Hidden MCP tool {raw_name!r} is ambiguous.") + match = matches[0] + if isinstance(match, MCPToolRegistry): + return await match.call_hidden(raw_name, arguments) + return await self._call(match, arguments) + + def has_hidden(self, raw_name: str) -> bool: + if any(dispatch.raw_name == raw_name for dispatch in self._dispatch.values()): + return True + return any(parent.has_hidden(raw_name) for parent in self.parents) + + def tool_allowed(self, name: str) -> bool: + dispatch = self._dynamic_dispatch.get(name) + if dispatch is not None: + return visibility_allows( + dispatch.toolset_name, self._toolsets_visibility + ) and visibility_allows(name, self._tools_visibility) + dispatch = self._dispatch.get(name) + if dispatch is not None: + return ( + not dispatch.binding.hidden + and visibility_allows(dispatch.toolset_name, self._toolsets_visibility) + and visibility_allows(name, self._tools_visibility) + ) + for parent in self.parents: + if parent.has_tool(name): + return parent.tool_allowed(name) + return False + + async def resolve( + self, + *, + context: JsonData, + resolution_key: str, + apply_updates: Callable[[list[BoundUpdate]], None] | None = None, + ) -> None: + if self._resolution_key == resolution_key: + return + self._dynamic_dispatch.clear() + self._dynamic_tools.clear() + self.set_context(context) + updates: list[BoundUpdate] = [] + for setup in self.setup_dispatches(): + result = await self.call_dispatch(setup, {}) + if result.updates: + updates.extend(result.updates) + self.register_setup_tools(setup, result.value) + if updates: + if apply_updates is None: + raise RuntimeError("Toolset setup returned state updates.") + apply_updates(updates) + self._resolution_key = resolution_key + + def setup_dispatches(self) -> list[ToolDispatch]: + dispatches: list[ToolDispatch] = [] + for parent in self.parents: + dispatches.extend(parent.setup_dispatches()) + for dispatch in self._dispatch.values(): + if dispatch.raw_name != "setup": + continue + if not dispatch.binding.hidden: + raise ValueError( + f"Toolset {dispatch.toolset_name!r} setup tool must be hidden." + ) + if visibility_allows(dispatch.toolset_name, self._toolsets_visibility): + dispatches.append(dispatch) + return dispatches + + async def call_dispatch( + self, dispatch: ToolDispatch, arguments: JsonData + ) -> ServerResult: + payload: JsonData = dict(arguments) + for arg_name, source in dispatch.binding.args.items(): + if source.startswith("resources."): + continue + if arg_name in payload: + raise ToolError( + f"MCP tool {dispatch.raw_name!r} argument {arg_name!r} is bound " + "and cannot be provided by the model." + ) + payload[arg_name] = resolve_binding(self._context, source) + result = await dispatch.session.call_tool(dispatch.raw_name, payload) + if bool(getattr(result, "isError", False)): + raise ToolError(str(mcp_content_value(getattr(result, "content", [])))) + content = mcp_content_value(getattr(result, "content", [])) + return split_result(content, dispatch.binding) + + def register_setup_tools(self, setup: ToolDispatch, value: JsonValue) -> None: + if not isinstance(value, dict): + raise TypeError( + f"Toolset {setup.toolset_name!r} setup must return a JSON object." + ) + if "messages" in value: + raise ValueError("Toolset setup cannot return messages.") + raw_tools = value.get("tools") + if raw_tools is None: + return + if not isinstance(raw_tools, list): + raise TypeError("Toolset setup tools must be a list.") + through = string_field(value, "through", default="call_tool") + name_arg = string_field(value, "name_arg", default="name") + input_arg = string_field(value, "input_arg", default="input") + route = self.dispatch_for(setup.toolset_name, through) + if not route.binding.hidden: + raise ValueError( + f"Dynamic tool route {setup.toolset_name}.{through} must be hidden." + ) + server = self.server_config(route.toolset_name) + for raw_tool in raw_tools: + tool = Tool.model_validate(raw_tool) + if server is not None and not tool_visible(server, tool.name): + continue + if self.has_tool(tool.name): + raise ValueError(f"Dynamic tool {tool.name!r} is defined twice.") + self._dynamic_tools.append(tool) + self._dynamic_dispatch[tool.name] = ToolDispatch( + session=route.session, + toolset_name=route.toolset_name, + raw_name=route.raw_name, + binding=route.binding, + dynamic_name=tool.name, + name_arg=name_arg, + input_arg=input_arg, + ) + + def server_config(self, toolset_name: str) -> ServerConfig | None: + server = self.servers.get(toolset_name) + if server is not None: + return server + for parent in self.parents: + server = parent.server_config(toolset_name) + if server is not None: + return server + return None + + def dispatch_for(self, toolset_name: str, raw_name: str) -> ToolDispatch: + for dispatch in self._dispatch.values(): + if dispatch.toolset_name == toolset_name and dispatch.raw_name == raw_name: + return dispatch + for parent in self.parents: + try: + return parent.dispatch_for(toolset_name, raw_name) + except KeyError: + pass + raise KeyError(f"Toolset {toolset_name!r} has no hidden tool {raw_name!r}.") + + +async def server_bindings(session: "ClientSession") -> dict[str, ToolBinding]: + result = await session.call_tool(_BINDINGS_TOOL, {}) + if bool(getattr(result, "isError", False)): + raise ToolError(str(mcp_content_value(getattr(result, "content", [])))) + content = mcp_content_value(getattr(result, "content", [])) + data = mapping_result(content) + tools = data.get("tools", {}) + if not isinstance(tools, dict): + raise TypeError("Server bindings metadata must contain a tools object.") + bindings: dict[str, ToolBinding] = {} + for tool_name, value in tools.items(): + if not isinstance(tool_name, str): + raise TypeError("Server binding tool names must be strings.") + if not isinstance(value, dict): + raise TypeError(f"Server binding {tool_name!r} must be an object.") + bindings[tool_name] = ToolBinding( + args=dict_mapping(value.get("args", {}), field=f"{tool_name}.args"), + sets=dict_mapping(value.get("sets", {}), field=f"{tool_name}.sets"), + extends=dict_mapping( + value.get("extends", {}), field=f"{tool_name}.extends" + ), + hidden=bool(value.get("hidden", False)), + ) + return bindings + + +def dict_mapping(value: object, *, field: str) -> dict[str, str]: + if not isinstance(value, Mapping): + raise TypeError(f"Server binding {field} must be an object.") + result: dict[str, str] = {} + for key, item in value.items(): + if not isinstance(key, str) or not isinstance(item, str): + raise TypeError(f"Server binding {field} must map strings to strings.") + result[key] = item + return result + + +def string_field(data: Mapping[str, JsonValue], field: str, *, default: str) -> str: + value = data.get(field, default) + if not isinstance(value, str) or not value: + raise TypeError(f"Toolset setup {field} must be a non-empty string.") + return value + + +def free_port() -> int: + for _ in range(50): + port = random.randint(3000, 8999) + probe = socket.socket() + try: + probe.bind(("127.0.0.1", port)) + return port + except OSError: + continue + finally: + probe.close() + raise RuntimeError("Could not find a free port in [3000, 9000).") + + +async def wait_for_http(url: str, *, timeout_seconds: float = 18.0) -> None: + deadline = asyncio.get_running_loop().time() + timeout_seconds + while asyncio.get_running_loop().time() < deadline: + if await asyncio.to_thread(http_serves, url): + return + await asyncio.sleep(0.1) + raise RuntimeError(f"MCP server did not start at {url}.") + + +@contextlib.asynccontextmanager +async def connect_streamable_http(url: str, headers: dict[str, str]): + from mcp.client.streamable_http import streamablehttp_client + + async with streamablehttp_client(url, headers=headers or None) as ( + read, + write, + *_, + ): + yield read, write + + +def http_serves(url: str) -> bool: + try: + urllib.request.urlopen(url, timeout=2) + return True + except urllib.error.HTTPError: + return True + except Exception: + return False + + +def mcp_content_value(content: object) -> JsonValue: + if not isinstance(content, list): + return serializable_content(content) + values = [serializable_content(item) for item in content] + if len(values) == 1: + return values[0] + return values + + +def serializable_content(item: object) -> JsonValue: + item_type = getattr(item, "type", None) + text = getattr(item, "text", None) + if item_type == "text" and isinstance(text, str): + return text + model_dump = getattr(item, "model_dump", None) + if callable(model_dump): + return json_value(model_dump(mode="json", exclude_none=True)) + return json_value(item) + + +def tool_visible(server: ServerConfig, tool_name: str) -> bool: + if server.show is not None: + return tool_name in server.show + if server.hide is not None: + return tool_name not in server.hide + return True + + +def visibility_allows(name: str, visibility: "TaskVisibility | None") -> bool: + if visibility is None: + return True + if visibility.show is not None: + return name in visibility.show + if visibility.hide is not None: + return name not in visibility.hide + return True + + +def _server_command( + server: ServerConfig, name: str, *, in_runtime: bool = False +) -> list[str]: + python = "python" if in_runtime else sys.executable + payload = { + "name": name, + "server": server.implementation_ref(), + "config": server.model_dump(mode="json"), + } + return [ + python, + "-m", + "verifiers.v1.mcp", + json.dumps(payload), + "streamable-http", + ] + + +def _server_env(server: ServerConfig) -> dict[str, str]: + resolved = dict(server.env) + pythonpath = os.pathsep.join(path for path in sys.path if path) + if pythonpath: + resolved.setdefault("PYTHONPATH", pythonpath) + return resolved + + +def build_fastmcp(toolset: Toolset): + from mcp.server.fastmcp import FastMCP + + mcp = FastMCP(toolset.name) + bindings = { + name: { + "args": dict(spec.args), + "sets": dict(spec.sets), + "extends": dict(spec.extends), + "hidden": spec.hidden, + } + for name, spec in type(toolset).tool_specs().items() + } + + @mcp.tool(name=_BINDINGS_TOOL) + def vf_bindings() -> dict[str, object]: + return {"tools": bindings} + + for method_name, method in inspect.getmembers(toolset, predicate=callable): + spec = getattr(getattr(type(toolset), method_name, None), "__vf_tool__", None) + if not isinstance(spec, ToolSpec): + continue + tool_name = spec.name or method_name + mcp.tool(name=tool_name)(server_tool_wrapper(toolset, method, spec)) + return mcp + + +def server_tool_wrapper( + toolset: Toolset, method: Callable[..., object], spec: ToolSpec +): + resource_args = { + arg_name: source + for arg_name, source in spec.args.items() + if source.startswith("resources.") + } + if not resource_args: + return method + signature = inspect.signature(method) + parameters = [ + parameter + for parameter in signature.parameters.values() + if parameter.name not in resource_args + ] + + async def invoke(**kwargs: object) -> object: + for arg_name, source in resource_args.items(): + kwargs[arg_name] = resolve_resource(toolset, source) + result = method(**kwargs) + if inspect.isawaitable(result): + return await result + return result + + invoke.__name__ = getattr(method, "__name__", "tool") + invoke.__doc__ = getattr(method, "__doc__", None) + setattr(invoke, "__signature__", signature.replace(parameters=parameters)) + annotations = dict(getattr(method, "__annotations__", {})) + for arg_name in resource_args: + annotations.pop(arg_name, None) + invoke.__annotations__ = annotations + return invoke + + +def resolve_resource(toolset: Toolset, source: str) -> object: + parts = source.split(".") + if len(parts) < 2 or parts[0] != "resources": + raise ValueError(f"Resource binding {source!r} must start with resources.") + value: object = toolset.resources + for part in parts[1:]: + if isinstance(value, dict): + value = cast(dict[str, object], value)[part] + else: + value = getattr(value, part) + return value + + +def _run_toolset_server(config_json: str, transport: str = "stdio") -> None: + payload = json.loads(config_json) + if not isinstance(payload, dict): + raise TypeError("Toolset runner payload must decode to an object.") + server = payload.get("server") + if not isinstance(server, str) or not server: + raise TypeError("Toolset runner payload requires a server.") + name = payload.get("name") + if not isinstance(name, str) or not name: + raise TypeError("Toolset runner payload requires a name.") + data = payload.get("config") + if not isinstance(data, dict): + raise TypeError("Toolset runner payload requires a config object.") + config_data = json_data(data, context="Toolset runner config") + config_cls = config_type_for_server(server) + config = config_cls.model_validate(config_data) + toolset = Toolset.load_ref(server, config) + toolset.name = name + toolset.start() + toolset.load_resources() + try: + mcp = build_fastmcp(toolset) + if transport == "streamable-http": + port = int(os.environ["MCP_PORT"]) + mcp.settings.host = "127.0.0.1" + mcp.settings.port = port + mcp.run(transport="streamable-http") + else: + mcp.run(transport="stdio") + finally: + toolset.stop() + + +def config_type_for_server(server: str) -> type[ServerConfig]: + obj = import_config_ref(server) + if isinstance(obj, type) and issubclass(obj, Toolset): + return registered_config_type(obj, ServerConfig) + return ServerConfig + + +def model_visible_schema( + tool_name: str, + schema: dict[str, object], + binding: ToolBinding, +) -> dict[str, object]: + if schema.get("type") != "object": + if binding.args: + raise ValueError(f"MCP tool {tool_name!r} has bindings but no args schema.") + return dict(schema) + properties = schema.get("properties") + if not isinstance(properties, dict): + if binding.args: + raise ValueError(f"MCP tool {tool_name!r} has bindings but no properties.") + return dict(schema) + json_bound_args = { + arg_name + for arg_name, source in binding.args.items() + if not source.startswith("resources.") + } + missing = sorted(json_bound_args - {str(key) for key in properties}) + if missing: + raise ValueError( + f"MCP tool {tool_name!r} binds unknown args: {', '.join(missing)}." + ) + visible = dict(schema) + visible["properties"] = { + str(key): value + for key, value in properties.items() + if str(key) not in binding.args + } + required = schema.get("required") + if isinstance(required, list): + visible["required"] = [item for item in required if item not in binding.args] + return visible + + +def resolve_binding(context: JsonData, source: str) -> JsonValue: + if not source: + raise ValueError("Binding source must be non-empty.") + value: JsonValue = context + for part in source.split("."): + if not part: + raise ValueError(f"Binding source {source!r} has an empty path segment.") + if isinstance(value, dict): + if part not in value: + raise KeyError(f"Binding source {source!r} is missing {part!r}.") + value = value[part] + elif isinstance(value, list): + if not part.isdigit(): + raise TypeError( + f"Binding source {source!r} indexes a list with {part!r}." + ) + index = int(part) + try: + value = value[index] + except IndexError as exc: + raise IndexError( + f"Binding source {source!r} list index {index} is out of range." + ) from exc + else: + raise TypeError( + f"Binding source {source!r} cannot traverse {type(value).__name__}." + ) + return deepcopy(value) + + +def split_result(content: JsonValue, binding: ToolBinding) -> ServerResult: + if not binding.sets and not binding.extends: + return ServerResult(response=server_response(content), value=deepcopy(content)) + result = mapping_result(content) + visible = dict(result) + updates: list[BoundUpdate] = [] + for result_key, target in binding.sets.items(): + if result_key not in result: + continue + updates.append( + BoundUpdate( + target=target, + value=deepcopy(result[result_key]), + mode="set", + ) + ) + visible.pop(result_key, None) + for result_key, target in binding.extends.items(): + if result_key not in result: + continue + updates.append( + BoundUpdate( + target=target, + value=deepcopy(result[result_key]), + mode="extend", + ) + ) + visible.pop(result_key, None) + if not visible: + model_content: JsonValue = "" + elif set(visible) == {"content"}: + model_content = visible["content"] + else: + model_content = dict(visible) + return ServerResult( + response=server_response(model_content), + updates=tuple(updates), + value=deepcopy(dict(result)), + ) + + +def mapping_result(content: JsonValue) -> JsonData: + if isinstance(content, dict): + return content + if isinstance(content, str): + parsed = json.loads(content) + if isinstance(parsed, dict): + return json_data(parsed, context="Bound MCP tool return") + raise TypeError("Bound MCP tool returns must be JSON objects.") + + +def server_response(content: JsonValue) -> ServerResponse: + if content == "": + return ServerResponse() + if isinstance(content, str): + try: + parsed = json.loads(content) + except json.JSONDecodeError: + return ServerResponse(content=content) + if isinstance(parsed, Mapping) and ( + "content" in parsed or "messages" in parsed + ): + return ServerResponse.model_validate(parsed) + return ServerResponse(content=content) + if isinstance(content, Mapping): + if "content" in content or "messages" in content: + return ServerResponse.model_validate(content) + return ServerResponse(content=json.dumps(content)) + if isinstance(content, list): + try: + return ServerResponse(content=cast(MessageContent, content)) + except ValidationError: + return ServerResponse(content=json.dumps(content)) + return ServerResponse(content=json.dumps(content)) + + +def main() -> None: + if len(sys.argv) not in (2, 3): + raise SystemExit("usage: python -m verifiers.v1.mcp CONFIG_JSON [TRANSPORT]") + transport = sys.argv[2] if len(sys.argv) == 3 else "stdio" + _run_toolset_server(sys.argv[1], transport) + + +if __name__ == "__main__": + main() diff --git a/verifiers/v1/model.py b/verifiers/v1/model.py deleted file mode 100644 index c947e0b813..0000000000 --- a/verifiers/v1/model.py +++ /dev/null @@ -1,54 +0,0 @@ -from typing import TYPE_CHECKING, TypeAlias, cast - -from verifiers.types import ClientConfig, SamplingArgs - -from .config import Config -from .types import ConfigData, ModelClient -from .utils.config_utils import resolve_config_object, string_mapping - -if TYPE_CHECKING: - from .task import Task - - -class ModelConfig(Config): - name: str | None = None - client: ClientConfig | str | None = None - sampling_args: SamplingArgs = {} - - def client_object(self) -> ModelClient | None: - if self.client is None: - return None - client = resolve_config_object(self.client) - if isinstance(client, ClientConfig): - return client - from verifiers.clients import Client - - if isinstance(client, Client): - return client - raise TypeError("model.client must resolve to a Client or ClientConfig.") - - -ModelConfigSource: TypeAlias = ModelConfig | str | ConfigData | None - - -def model_config_from_value(value: ModelConfigSource = None) -> ModelConfig: - if isinstance(value, ModelConfig): - return value - if isinstance(value, str): - return ModelConfig(name=value) - if isinstance(value, dict): - return ModelConfig.model_validate(string_mapping(value)) - if value is None: - return ModelConfig() - raise TypeError("model must be a string or mapping.") - - -def model_config_data(value: ModelConfigSource = None) -> ConfigData: - return cast( - ConfigData, model_config_from_value(value).model_dump(exclude_none=True) - ) - - -def model_config_from_task(task: "Task") -> ModelConfig: - value = task.get("model") - return model_config_from_value(value) diff --git a/verifiers/v1/program.py b/verifiers/v1/program.py deleted file mode 100644 index b636feb00f..0000000000 --- a/verifiers/v1/program.py +++ /dev/null @@ -1,301 +0,0 @@ -from typing import Literal, TypeAlias, cast - -from pydantic import field_validator, model_validator -from typing_extensions import TypeAliasType - -from .artifact import ArtifactsConfig -from .config import Config -from .sandbox import SandboxConfig -from .types import ConfigData, ConfigValue -from .utils.binding_utils import BindingsConfig -from .utils.config_utils import explicit_config_data, string_mapping -from .utils.mcp_proxy_utils import validate_program_channels - -ProgramCallableRef: TypeAlias = str -ProgramScalar: TypeAlias = str | int | float | bool | None -ProgramValue = TypeAliasType( - "ProgramValue", - ProgramScalar | list["ProgramValue"] | dict[str, "ProgramValue"], -) -ProgramCommand: TypeAlias = str | list[ProgramValue] -ProgramFiles: TypeAlias = dict[str, ProgramValue] -ProgramDirs: TypeAlias = dict[str, ProgramValue] -ProgramEnv: TypeAlias = dict[str, ProgramValue] -ProgramArtifacts: TypeAlias = ArtifactsConfig -ProgramSetup: TypeAlias = ProgramValue | list[ProgramValue] -ProgramArgs: TypeAlias = list[ProgramValue] -ProgramChannel: TypeAlias = Literal["callable", "mcp"] -ProgramChannelConfig: TypeAlias = ProgramChannel | dict[str, ProgramValue] -ProgramChannels: TypeAlias = ProgramChannelConfig | list[ProgramChannelConfig] -COMMAND_SANDBOX_DEFAULTS: ConfigData = { - "image": "python:3.11-slim", - "workdir": "/app", - "scope": "rollout", - "timeout_minutes": 120, - "command_timeout": 900, - "network_access": True, -} -COMMAND_PROGRAM_PATCH_KEYS = { - "sandbox", - "files", - "dirs", - "setup", - "setup_timeout", - "bindings", - "env", - "artifacts", - "args", -} -COMMAND_PROGRAM_MAP_PATCH_KEYS = {"files", "dirs", "bindings", "env", "artifacts"} -COMMAND_PROGRAM_LIST_PATCH_KEYS = {"setup", "args"} - -__all__ = [ - "ProgramArgs", - "ProgramArtifacts", - "ProgramCallableRef", - "ProgramChannels", - "ProgramCommand", - "ProgramConfig", - "ProgramDirs", - "ProgramEnv", - "ProgramFiles", - "ProgramSetup", - "ProgramValue", - "program_config_data", -] - - -class ProgramConfig(Config): - base: bool = False - fn: ProgramCallableRef | None = None - command: ProgramCommand | None = None - sandbox: bool | SandboxConfig | None = None - files: ProgramFiles = {} - dirs: ProgramDirs = {} - setup: ProgramSetup = [] - setup_timeout: int = 300 - bindings: BindingsConfig = BindingsConfig() - env: ProgramEnv = {} - artifacts: ArtifactsConfig = ArtifactsConfig() - channels: ProgramChannels | None = None - args: ProgramArgs = [] - - @field_validator("fn") - @classmethod - def validate_fn(cls, value: object) -> object: - validate_program_callable_ref(value, "program.fn") - return value - - @field_validator("channels") - @classmethod - def validate_channels(cls, value: object) -> object: - validate_program_channels(value) - return value - - @field_validator("env", mode="before") - @classmethod - def validate_env(cls, value: object) -> object: - if isinstance(value, dict): - return {str(key): item for key, item in value.items()} - return value - - @field_validator("env") - @classmethod - def normalize_env_values(cls, value: ProgramEnv) -> ProgramEnv: - return { - key: item if isinstance(item, dict) else str(item) - for key, item in value.items() - } - - @field_validator("bindings", mode="before") - @classmethod - def validate_bindings(cls, value: object) -> BindingsConfig: - bindings = BindingsConfig.model_validate(value or {}) - bindings.entries("program.bindings", allow_objects=False) - return bindings - - @model_validator(mode="after") - def validate_program_callable_refs(self) -> "ProgramConfig": - for name, value in ( - ("command", self.command), - ("files", self.files), - ("dirs", self.dirs), - ("setup", self.setup), - ("env", self.env), - ("artifacts", self.artifacts), - ("channels", self.channels), - ("args", self.args), - ): - validate_program_value_refs(value, f"program.{name}") - return self - - def data(self) -> ConfigData: - resolved = self.resolve() - if resolved is not self: - return resolved.data() - data = program_config_data(self) - return data if data else {"base": True} - - def resolve(self) -> "ProgramConfig": - return self - - def resolve_command( - self, - *, - command: ProgramCommand, - sandbox: bool | SandboxConfig | None = None, - default_sandbox: bool | SandboxConfig | None = True, - sandbox_defaults: ConfigData | None = None, - files: ProgramFiles | None = None, - dirs: ProgramDirs | None = None, - setup: ProgramSetup | None = None, - setup_timeout: int | None = None, - bindings: ConfigData | None = None, - env: ProgramEnv | None = None, - artifacts: ArtifactsConfig | ConfigData | None = None, - channels: ProgramChannels | None = None, - args: ProgramArgs | None = None, - ) -> "ProgramConfig": - sandbox_value = ( - sandbox - if sandbox is not None - else self.sandbox - if self.sandbox is not None - else default_sandbox - if default_sandbox is not None - else True - ) - resolved_sandbox = command_sandbox_config( - sandbox_value, defaults=sandbox_defaults - ) - data: ConfigData = { - "command": cast(ConfigValue, command), - "sandbox": resolved_sandbox.model_dump(exclude_none=True) - if resolved_sandbox is not None - else False, - } - if files is not None: - data["files"] = dict(files) - if dirs is not None: - data["dirs"] = dict(dirs) - if setup is not None: - data["setup"] = setup - if setup_timeout is not None: - data["setup_timeout"] = setup_timeout - if bindings is not None: - data["bindings"] = dict(bindings) - if env is not None: - data["env"] = dict(env) - if artifacts is not None: - data["artifacts"] = ArtifactsConfig.model_validate(artifacts).data( - "program.artifacts" - ) - if channels is not None: - data["channels"] = channels - if args is not None: - data["args"] = list(args) - return ProgramConfig.model_validate(merge_command_program_config(data, self)) - - -def validate_program_callable_ref(value: object, field_name: str) -> None: - if value is None: - return - if not isinstance(value, str) or not value: - raise ValueError(f"{field_name} must be a non-empty import ref string.") - - -def validate_program_value_refs(value: object, field_name: str) -> None: - if isinstance(value, dict): - mapping = string_mapping(value) - if "fn" in mapping: - validate_program_callable_ref(mapping["fn"], f"{field_name}.fn") - for key, item in mapping.items(): - validate_program_value_refs(item, f"{field_name}.{key}") - return - if isinstance(value, list | tuple): - for index, item in enumerate(value): - validate_program_value_refs(item, f"{field_name}.{index}") - - -def command_sandbox_config( - sandbox: bool | SandboxConfig, - *, - defaults: ConfigData | None = None, -) -> SandboxConfig | None: - if sandbox is False: - return None - base = {**COMMAND_SANDBOX_DEFAULTS, **dict(defaults or {})} - if sandbox is True: - return SandboxConfig.model_validate(base) - return SandboxConfig.model_validate({**base, **sandbox.data()}) - - -def merge_command_program_config( - program: ConfigData, - patch_config: ProgramConfig, -) -> ConfigData: - patch = program_config_data(patch_config) - unknown = sorted(set(patch) - COMMAND_PROGRAM_PATCH_KEYS) - if unknown: - allowed = ", ".join(sorted(COMMAND_PROGRAM_PATCH_KEYS)) - raise ValueError( - f"Command ProgramConfig can only define {allowed}; got {unknown}." - ) - merged: ConfigData = dict(program) - for key, value in patch.items(): - if key == "sandbox" and "sandbox" in merged: - continue - if key in COMMAND_PROGRAM_MAP_PATCH_KEYS: - if not isinstance(value, dict): - raise TypeError(f"program.{key} must be a mapping.") - base = merged.get(key, {}) - if base is None: - base = {} - if not isinstance(base, dict): - raise TypeError(f"command program {key} must be a mapping.") - merged[key] = {**dict(base), **dict(value)} - elif key in COMMAND_PROGRAM_LIST_PATCH_KEYS: - merged[key] = [ - *program_list_items( - cast(ProgramSetup | None, merged.get(key)), - f"command program {key}", - ), - *program_list_items( - cast(ProgramSetup | None, value), - f"program.{key}", - ), - ] - else: - merged[key] = value - return merged - - -def program_list_items( - value: ProgramSetup | None, field_name: str -) -> list[ProgramValue]: - if value is None: - return [] - if isinstance(value, list): - return cast(list[ProgramValue], list(value)) - if isinstance(value, str) or isinstance(value, dict): - return [cast(ProgramValue, value)] - raise TypeError(f"{field_name} must be a string, mapping, or list.") - - -PROGRAM_DEFAULT_DUMP_DATA = ProgramConfig().model_dump(exclude_none=True) -PROGRAM_DEFAULT_DUMP_KEYS = set(PROGRAM_DEFAULT_DUMP_DATA) - - -def program_config_data(config: ProgramConfig) -> ConfigData: - data = { - key: value - for key, value in explicit_config_data(config).items() - if key in ProgramConfig.model_fields - } - if PROGRAM_DEFAULT_DUMP_KEYS.issubset(config.model_fields_set): - data = { - key: value - for key, value in data.items() - if value != PROGRAM_DEFAULT_DUMP_DATA.get(key) - } - return data diff --git a/verifiers/v1/protocols.py b/verifiers/v1/protocols.py new file mode 100644 index 0000000000..311a5e922d --- /dev/null +++ b/verifiers/v1/protocols.py @@ -0,0 +1,572 @@ +from __future__ import annotations + +import json +from typing import cast + +from aiohttp import web +from pydantic import TypeAdapter + +from verifiers.types import ( + AssistantMessage, + Message, + MessageContent, + Messages, + Response, + SystemMessage, + TextMessage, + Tool, + ToolCall, + ToolMessage, + UserMessage, +) + +from .interception import EndpointProtocol, InterceptedRequest, ProtocolRoute +from .types import JsonData, JsonValue +from .utils.json_utils import json_data, json_value + +_MESSAGES_ADAPTER = TypeAdapter(Messages) + + +class OpenAIProtocol: + def env(self, *, base_url: str, api_key: str, model: str) -> dict[str, str]: + return { + "OPENAI_BASE_URL": f"{base_url.rstrip('/')}/v1", + "OPENAI_API_KEY": api_key, + "OPENAI_MODEL": model, + } + + def usage(self, response: Response) -> JsonData | None: + if response.usage is None: + return None + return { + "prompt_tokens": response.usage.prompt_tokens, + "completion_tokens": response.usage.completion_tokens, + "total_tokens": response.usage.total_tokens, + } + + def tool_calls(self, response: Response) -> list[JsonValue] | None: + calls = response.message.tool_calls + if not calls: + return None + return [ + { + "id": call.id, + "type": "function", + "function": {"name": call.name, "arguments": call.arguments}, + } + for call in calls + ] + + +class OpenAIChatCompletionsProtocol(OpenAIProtocol, EndpointProtocol): + name = "openai_chat_completions" + routes = (ProtocolRoute("POST", "/v1/chat/completions"),) + + async def parse(self, request: web.Request, body: JsonData) -> InterceptedRequest: + _ = request + return InterceptedRequest( + protocol=self.name, + prompt=parse_openai_messages(body.get("messages")), + model=string_value(body.get("model")), + sampling_args=openai_sampling_args(body), + tools=parse_openai_tools(body.get("tools")), + body=body, + ) + + def serialize(self, response: Response, request: InterceptedRequest) -> JsonData: + message: JsonData = { + "role": "assistant", + "content": message_content(response.message.content), + } + tool_calls = self.tool_calls(response) + if tool_calls: + message["tool_calls"] = tool_calls + return { + "id": response.id or "vf-v1-intercept", + "object": "chat.completion", + "created": response.created, + "model": response.model or request.model or "", + "choices": [ + { + "index": 0, + "message": message, + "finish_reason": response.message.finish_reason or "stop", + } + ], + "usage": self.usage(response), + } + + +class OpenAICompletionsProtocol(OpenAIProtocol, EndpointProtocol): + name = "openai_completions" + routes = (ProtocolRoute("POST", "/v1/completions"),) + + async def parse(self, request: web.Request, body: JsonData) -> InterceptedRequest: + _ = request + raw_prompt = body.get("prompt") + return InterceptedRequest( + protocol=self.name, + prompt=( + [TextMessage(content=raw_prompt)] + if isinstance(raw_prompt, str) + else _MESSAGES_ADAPTER.validate_python(raw_prompt or []) + ), + model=string_value(body.get("model")), + sampling_args=openai_sampling_args(body), + body=body, + ) + + def serialize(self, response: Response, request: InterceptedRequest) -> JsonData: + return { + "id": response.id or "vf-v1-intercept", + "object": "text_completion", + "created": response.created, + "model": response.model or request.model or "", + "choices": [ + { + "index": 0, + "text": content_text(response.message.content), + "finish_reason": response.message.finish_reason or "stop", + } + ], + "usage": self.usage(response), + } + + +class OpenAIResponsesProtocol(OpenAIProtocol, EndpointProtocol): + name = "openai_responses" + routes = (ProtocolRoute("POST", "/v1/responses"),) + + async def parse(self, request: web.Request, body: JsonData) -> InterceptedRequest: + _ = request + return InterceptedRequest( + protocol=self.name, + prompt=parse_openai_responses_input(body.get("input")), + model=string_value(body.get("model")), + sampling_args=openai_sampling_args(body), + tools=parse_responses_tools(body.get("tools")), + body=body, + ) + + def serialize(self, response: Response, request: InterceptedRequest) -> JsonData: + output: list[JsonValue] = [] + tool_calls = response.message.tool_calls or [] + for call in tool_calls: + output.append( + { + "type": "function_call", + "id": call.id, + "call_id": call.id, + "name": call.name, + "arguments": call.arguments, + } + ) + content = content_text(response.message.content) + if content: + output.append( + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": content}], + } + ) + return json_data( + { + "id": response.id or "vf-v1-intercept", + "object": "response", + "created_at": response.created, + "model": response.model or request.model or "", + "output": output, + "output_text": content, + "usage": self.usage(response), + }, + context="OpenAI responses serialization", + ) + + +class AnthropicMessagesProtocol(EndpointProtocol): + name = "anthropic_messages" + routes = (ProtocolRoute("POST", "/v1/messages"),) + + def env(self, *, base_url: str, api_key: str, model: str) -> dict[str, str]: + return { + "ANTHROPIC_BASE_URL": base_url.rstrip("/"), + "ANTHROPIC_API_KEY": api_key, + "ANTHROPIC_MODEL": model, + } + + async def parse(self, request: web.Request, body: JsonData) -> InterceptedRequest: + _ = request + return InterceptedRequest( + protocol=self.name, + prompt=parse_anthropic_messages(body), + model=string_value(body.get("model")), + sampling_args=anthropic_sampling_args(body), + tools=parse_anthropic_tools(body.get("tools")), + body=body, + ) + + def serialize(self, response: Response, request: InterceptedRequest) -> JsonData: + content: list[JsonValue] = [] + text = content_text(response.message.content) + if text: + content.append({"type": "text", "text": text}) + for call in response.message.tool_calls or []: + content.append( + { + "type": "tool_use", + "id": call.id, + "name": call.name, + "input": json.loads(call.arguments or "{}"), + } + ) + usage = None + if response.usage is not None: + usage = { + "input_tokens": response.usage.prompt_tokens, + "output_tokens": response.usage.completion_tokens, + } + return json_data( + { + "id": response.id or "vf-v1-intercept", + "type": "message", + "role": "assistant", + "model": response.model or request.model or "", + "content": content, + "stop_reason": response.message.finish_reason or "end_turn", + "usage": usage, + }, + context="Anthropic messages serialization", + ) + + +def default_protocols() -> list[EndpointProtocol]: + return [ + OpenAIChatCompletionsProtocol(), + OpenAICompletionsProtocol(), + OpenAIResponsesProtocol(), + AnthropicMessagesProtocol(), + ] + + +def parse_openai_messages(raw: JsonValue | None) -> Messages: + if not isinstance(raw, list): + raise TypeError("OpenAI messages must be a list.") + messages: Messages = [] + for item in raw: + if not isinstance(item, dict): + raise TypeError("OpenAI message entries must be objects.") + messages.append(parse_openai_message(json_data(item))) + return messages + + +def parse_openai_message(raw: JsonData) -> Message: + role = raw.get("role") + content = raw.get("content") + if role == "system": + return SystemMessage(content=content_text(content)) + if role == "tool": + return ToolMessage( + tool_call_id=str(raw.get("tool_call_id") or ""), + content=content_text(content), + ) + if role == "assistant": + return AssistantMessage( + content=content_text(content) if content is not None else None, + tool_calls=parse_openai_tool_calls(raw.get("tool_calls")), + ) + return UserMessage(content=content_text(content)) + + +def parse_openai_tool_calls(raw: JsonValue | None) -> list[ToolCall] | None: + if raw is None: + return None + if not isinstance(raw, list): + raise TypeError("OpenAI tool_calls must be a list.") + calls: list[ToolCall] = [] + for item in raw: + if not isinstance(item, dict): + continue + data = json_data(item) + function = data.get("function") + if not isinstance(function, dict): + continue + function_data = json_data(function) + calls.append( + ToolCall( + id=str(data.get("id") or ""), + name=str(function_data.get("name") or ""), + arguments=str(function_data.get("arguments") or "{}"), + ) + ) + return calls or None + + +def parse_openai_tools(raw: JsonValue | None) -> list[Tool] | None: + if raw is None: + return None + if not isinstance(raw, list): + raise TypeError("OpenAI tools must be a list.") + tools: list[Tool] = [] + for item in raw: + if not isinstance(item, dict): + continue + data = json_data(item) + function = data.get("function") + if data.get("type") != "function" or not isinstance(function, dict): + continue + function_data = json_data(function) + tools.append( + Tool( + name=str(function_data.get("name") or ""), + description=str(function_data.get("description") or ""), + parameters=tool_parameters(function_data.get("parameters")), + strict=bool_value(function_data.get("strict")), + ) + ) + return tools or None + + +def parse_responses_tools(raw: JsonValue | None) -> list[Tool] | None: + if raw is None: + return None + if not isinstance(raw, list): + raise TypeError("Responses tools must be a list.") + tools: list[Tool] = [] + for item in raw: + if not isinstance(item, dict): + continue + data = json_data(item) + tools.append( + Tool( + name=str(data.get("name") or ""), + description=str(data.get("description") or ""), + parameters=tool_parameters(data.get("parameters")), + strict=bool_value(data.get("strict")), + ) + ) + return tools or None + + +def parse_anthropic_tools(raw: JsonValue | None) -> list[Tool] | None: + if raw is None: + return None + if not isinstance(raw, list): + raise TypeError("Anthropic tools must be a list.") + tools: list[Tool] = [] + for item in raw: + if not isinstance(item, dict): + continue + data = json_data(item) + tools.append( + Tool( + name=str(data.get("name") or ""), + description=str(data.get("description") or ""), + parameters=tool_parameters(data.get("input_schema")), + ) + ) + return tools or None + + +def parse_openai_responses_input(raw: JsonValue | None) -> Messages: + if isinstance(raw, str): + return [UserMessage(content=raw)] + if not isinstance(raw, list): + raise TypeError("Responses input must be a string or list.") + messages: Messages = [] + for item in raw: + if not isinstance(item, dict): + raise TypeError("Responses input entries must be objects.") + data = json_data(item) + item_type = data.get("type") + if item_type == "function_call": + call_id = data.get("call_id") or data.get("id") + name = data.get("name") + arguments = data.get("arguments") + if isinstance(call_id, str) and isinstance(name, str): + messages.append( + AssistantMessage( + tool_calls=[ + ToolCall( + id=call_id, + name=name, + arguments=str(arguments or "{}"), + ) + ] + ) + ) + continue + if item_type == "function_call_output": + call_id = data.get("call_id") + if isinstance(call_id, str): + messages.append( + ToolMessage( + tool_call_id=call_id, + content=content_text(data.get("output")), + ) + ) + continue + role = data.get("role") + content = content_text(data.get("content")) + if role in {"system", "developer"}: + messages.append(SystemMessage(content=content)) + elif role == "assistant": + messages.append(AssistantMessage(content=content)) + else: + messages.append(UserMessage(content=content)) + return messages + + +def parse_anthropic_messages(body: JsonData) -> Messages: + messages: Messages = [] + system = body.get("system") + if isinstance(system, str) and system: + messages.append(SystemMessage(content=system)) + raw_messages = body.get("messages") + if not isinstance(raw_messages, list): + raise TypeError("Anthropic messages must be a list.") + for item in raw_messages: + if not isinstance(item, dict): + raise TypeError("Anthropic message entries must be objects.") + data = json_data(item) + role = data.get("role") + if role == "assistant": + messages.append(parse_anthropic_assistant_message(data.get("content"))) + elif role == "user": + messages.extend(parse_anthropic_user_messages(data.get("content"))) + else: + raise ValueError(f"Unsupported Anthropic role: {role!r}.") + return messages + + +def parse_anthropic_assistant_message(content: JsonValue | None) -> AssistantMessage: + if isinstance(content, str): + return AssistantMessage(content=content) + if not isinstance(content, list): + return AssistantMessage(content=content_text(content)) + text_parts: list[str] = [] + tool_calls: list[ToolCall] = [] + for block in content: + if not isinstance(block, dict): + continue + data = json_data(block) + if data.get("type") == "text": + text_parts.append(content_text(data.get("text"))) + elif data.get("type") == "tool_use": + tool_id = data.get("id") + name = data.get("name") + if isinstance(tool_id, str) and isinstance(name, str): + tool_calls.append( + ToolCall( + id=tool_id, + name=name, + arguments=json.dumps(data.get("input") or {}), + ) + ) + return AssistantMessage( + content="\n".join(text_parts) if text_parts else None, + tool_calls=tool_calls or None, + ) + + +def parse_anthropic_user_messages(content: JsonValue | None) -> Messages: + if isinstance(content, str): + return [UserMessage(content=content)] + if not isinstance(content, list): + return [UserMessage(content=content_text(content))] + messages: Messages = [] + content_parts: list[JsonData] = [] + for block in content: + if not isinstance(block, dict): + raise TypeError("Anthropic user content blocks must be objects.") + data = json_data(block) + if data.get("type") == "text": + content_parts.append( + {"type": "text", "text": content_text(data.get("text"))} + ) + elif data.get("type") == "tool_result": + tool_use_id = data.get("tool_use_id") + if isinstance(tool_use_id, str): + messages.append( + ToolMessage( + tool_call_id=tool_use_id, + content=content_text(data.get("content")), + ) + ) + else: + content_parts.append(data) + if content_parts: + if all(part.get("type") == "text" for part in content_parts): + messages.insert( + 0, + UserMessage( + content="\n".join( + content_text(part.get("text")) for part in content_parts + ) + ), + ) + else: + messages.insert(0, UserMessage(content=cast(MessageContent, content_parts))) + if not messages: + raise ValueError("Anthropic user message contained no supported content.") + return messages + + +def openai_sampling_args(body: JsonData) -> dict[str, JsonValue]: + keys = ( + "temperature", + "top_p", + "max_tokens", + "frequency_penalty", + "presence_penalty", + "seed", + "stop", + ) + return {key: body[key] for key in keys if key in body} + + +def anthropic_sampling_args(body: JsonData) -> dict[str, JsonValue]: + keys = ("temperature", "top_p", "top_k", "max_tokens", "stop_sequences") + return {key: body[key] for key in keys if key in body} + + +def message_content(content: object) -> JsonValue | None: + if content is None: + return None + return json_value(content, context="message content") + + +def content_text(content: object) -> str: + value = message_content(content) + if value is None: + return "" + if isinstance(value, str): + return value + if isinstance(value, list): + text_parts: list[str] = [] + for item in value: + if isinstance(item, dict): + data = json_data(item) + text = data.get("text") + if isinstance(text, str): + text_parts.append(text) + else: + text_parts.append(str(item)) + return "\n".join(text_parts) + return str(content) + + +def tool_parameters(value: JsonValue | None) -> dict[str, object]: + if value is None: + return {"type": "object", "properties": {}} + if not isinstance(value, dict): + raise TypeError("Tool parameters must be an object.") + return {str(key): item for key, item in value.items()} + + +def string_value(value: JsonValue | None) -> str | None: + return value if isinstance(value, str) and value else None + + +def bool_value(value: JsonValue | None) -> bool | None: + return value if isinstance(value, bool) else None diff --git a/verifiers/v1/runtime.py b/verifiers/v1/runtime.py index 3fb4775e82..ea7b73000f 100644 --- a/verifiers/v1/runtime.py +++ b/verifiers/v1/runtime.py @@ -1,2812 +1,777 @@ +from __future__ import annotations + import asyncio -import glob -import hashlib -import inspect -import logging -import time +import base64 +import contextlib +import os +import signal +import shlex +import shutil +import tempfile import uuid -from collections.abc import Awaitable, Callable, Iterable, Sequence -from contextlib import AsyncExitStack -from dataclasses import dataclass, field -from importlib.abc import Traversable -from pathlib import Path -from typing import ( - TYPE_CHECKING, - Literal, - Protocol, - TypeAlias, - cast, - get_args, - runtime_checkable, +from abc import ABC, abstractmethod +from pathlib import Path, PurePosixPath +from typing import Annotated, Literal, Protocol + +from pydantic import Field + +from .config import Config + +_SUBPROCESS_ENV = ( + "PATH", + "HOME", + "USER", + "LOGNAME", + "SHELL", + "LANG", + "LC_ALL", + "TERM", + "TMPDIR", ) -from verifiers.clients import Client, resolve_client -from verifiers.types import Messages, Response, ResponseMessage, Tool -from verifiers.types import ClientConfig, ClientType, SamplingArgs -from verifiers.utils.async_utils import maybe_call_with_named_args -from verifiers.utils.client_utils import resolve_client_config -from verifiers.utils.message_utils import normalize_messages -from verifiers.utils.response_utils import parse_response_message, parse_response_tokens -from verifiers.utils.tool_utils import convert_func_to_tool_def - -from .utils.binding_utils import ( - BindingSource, - GROUP_FRAMEWORK_ARGS, - ROLLOUT_FRAMEWORK_ARGS, - binding_key_parts, - binding_object_name, - binding_source_root, - function_name, - owner_object_name, - read_path, - same_callable, - validate_binding_source, - validate_bound_arg, - validate_callable_source, -) -from .utils.config_callable_utils import CallableKind -from .utils.config_utils import resolve_config_object -from .utils.lifecycle_utils import collect_handlers, handler_is_marked, handler_stage -from .utils.lifecycle_utils import run_handlers, sort_handlers -from .utils.lifecycle_utils import state_done, unique_handlers, validate_handler_args -from .utils.object_utils import close_object, resolve_object_factory -from .utils.runtime_registry import load_runtime, register_runtime, unregister_runtime -from .utils.runtime_owner_utils import RuntimeOwnerMixin -from .utils.scoring_utils import SignalRecord, build_signals, collect_signals -from .utils.scoring_utils import group_framework_kwargs, rollout_framework_kwargs -from .utils.scoring_utils import score_group as score_group_signals -from .utils.scoring_utils import score_rollout as score_rollout_signals -from .utils.serialization_utils import serializable -from .artifact import ArtifactConfig -from .sandbox import SandboxConfig -from .runtime_handles import ( - ModelRuntimeHandleConfig, - ResolvedRuntimeHandlesConfig, - RuntimeHandleConfig, - SandboxRuntimeHandleConfig, - SandboxRuntimeStateConfig, -) -from .utils.tool_utils import schema_callable, tool_schema, tool_visible -from .utils.tool_utils import toolset_object_scope -from .utils.usage_utils import record_response_usage -from .state import State -from .task import Task -from .toolset import ( - MCPTool, - ToolEntry, - Toolset, - VisibilityConfig, - iter_toolsets, - tool_name, -) -from .user import User, state_messages -from .types import ( - ConfigData, - Handler, - JsonData, - PromptMessage, - RuntimeCallable, - RuntimeCallableResult, - RuntimeData, - RuntimeObject, -) -logger = logging.getLogger(__name__) - -if TYPE_CHECKING: - from .harness import Harness - from .taskset import Taskset - from .utils.mcp_utils import MCPToolHandle - from .utils.sandbox_utils import SandboxClient, SandboxLease - -BindingOwner = Toolset | RuntimeOwnerMixin | None -BindingEntry = tuple[str, BindingSource, BindingOwner] -ArtifactOwner = RuntimeOwnerMixin | Toolset | User | None -TrajectoryVisibility = Literal["append", "hidden"] -RuntimeObjectOwner = Literal["toolset", "user", "taskset", "harness"] -RuntimeObjectKey = tuple[int, str, str] -RuntimeObjectStore = dict[RuntimeObjectKey, RuntimeObject] -RUNTIME_OBJECT_OWNERS: tuple[RuntimeObjectOwner, ...] = ( - "toolset", - "user", - "taskset", - "harness", -) +class RuntimeTunnel(Protocol): + def sync_stop(self) -> None: ... -@runtime_checkable -class ToolDefinitionProvider(Protocol): - @property - def tool_def(self) -> Tool: ... - - def __call__(self, **kwargs: object) -> RuntimeCallableResult: ... - - -RuntimeTool: TypeAlias = RuntimeCallable | Tool | ToolDefinitionProvider -RuntimeTools: TypeAlias = dict[str, RuntimeTool] - - -def lifecycle_handlers( - owner: RuntimeOwnerMixin | Toolset, kind: CallableKind -) -> Iterable[Handler]: - if kind == "stop": - return owner.stops - if kind == "setup": - return owner.setups - if kind == "update": - return owner.updates - if kind == "cleanup": - return owner.cleanups - if kind == "teardown": - return owner.teardowns - if isinstance(owner, Toolset): - return () - if kind == "metric": - return owner.metrics - if kind == "reward": - return owner.rewards - if kind == "advantage": - return owner.advantages - raise ValueError(f"Unknown lifecycle kind: {kind!r}.") - - -@dataclass(frozen=True) -class ModelRequestContext: - source: Literal["direct", "endpoint"] = "direct" - endpoint_request_id: str | None = None - headers: dict[str, str] = field(default_factory=dict) - trajectory_visibility: TrajectoryVisibility = "append" - - def extras(self) -> ConfigData: - data: ConfigData = {} - if self.source == "endpoint": - data["endpoint"] = True - if self.endpoint_request_id is not None: - data["endpoint_request_id"] = self.endpoint_request_id - if self.headers: - data["headers"] = self.headers - if self.trajectory_visibility != "append": - data["trajectory_visibility"] = self.trajectory_visibility - return data - - -@dataclass(frozen=True) -class RuntimeArtifact: - name: str - config: ArtifactConfig - owner: ArtifactOwner = None - - -class BorrowedTool: - def __init__(self, runtime: "Runtime", handle_id: str, name: str): - self.runtime = runtime - self.handle_id = handle_id - self.name = name - self.__name__ = name +class CommandResult(Config): + returncode: int + stdout: str = "" + stderr: str = "" @property - def tool_def(self) -> Tool: - return self.runtime.borrowed_tool_def(self.handle_id, self.name) - - async def __call__(self, **kwargs: object) -> object: - return await self.runtime.call_borrowed_tool( - self.handle_id, self.name, **kwargs - ) + def exit_code(self) -> int: + return self.returncode + + +class SubprocessRuntimeConfig(Config): + type: Literal["subprocess"] = "subprocess" + + +class DockerRuntimeConfig(Config): + type: Literal["docker"] = "docker" + image: str = "python:3.11-slim" + workdir: str = "/app" + cpu_cores: float | None = None + memory_gb: float | None = None + gpu_count: int | None = None + disk_gb: float | None = None + + +class PrimeRuntimeConfig(Config): + type: Literal["prime"] = "prime" + image: str = "python:3.11-slim" + workdir: str = "/app" + network_access: bool = True + vm: bool = False + guaranteed: bool = False + region: str | None = None + gpu_type: str | None = None + timeout_minutes: int | Literal["auto"] = 360 + idle_timeout_minutes: int | None = None + cpu_cores: float = 1.0 + memory_gb: float = 2.0 + gpu_count: int = 0 + disk_gb: float = 5.0 + labels: list[str] = Field(default_factory=list) + + +class ModalRuntimeConfig(Config): + type: Literal["modal"] = "modal" + image: str = "python:3.11-slim" + + +class DaytonaRuntimeConfig(Config): + type: Literal["daytona"] = "daytona" + image: str = "python:3.11-slim" + + +RuntimeConfig = Annotated[ + SubprocessRuntimeConfig + | DockerRuntimeConfig + | PrimeRuntimeConfig + | ModalRuntimeConfig + | DaytonaRuntimeConfig, + Field(discriminator="type"), +] +RuntimeConfigValue = ( + SubprocessRuntimeConfig + | DockerRuntimeConfig + | PrimeRuntimeConfig + | ModalRuntimeConfig + | DaytonaRuntimeConfig +) +RUNTIME_CONFIG_TYPES = ( + SubprocessRuntimeConfig, + DockerRuntimeConfig, + PrimeRuntimeConfig, + ModalRuntimeConfig, + DaytonaRuntimeConfig, +) -class AsyncRateLimiter: - def __init__(self, rate_per_second: float | None): - self.interval = 0.0 if rate_per_second is None else 1.0 / rate_per_second - self.next_at = 0.0 - self.lock = asyncio.Lock() +class Runtime(ABC): + @abstractmethod + async def start(self) -> None: ... - async def wait(self) -> None: - if not self.interval: - return - async with self.lock: - now = time.monotonic() - delay = self.next_at - now - if delay > 0: - await asyncio.sleep(delay) - now = time.monotonic() - self.next_at = max(now, self.next_at) + self.interval - - -class Runtime: - def __init__( - self, taskset: "Taskset | None" = None, harness: "Harness | None" = None - ): - self.runtime_id = uuid.uuid4().hex - register_runtime(self.runtime_id, self) - self.taskset = taskset - self.harness = harness - owners = (self.taskset, self.harness) - self.toolsets = [] - for owner in owners: - if owner is not None: - self.toolsets.extend(iter_toolsets(owner.toolsets)) - self.named_toolsets = self._collect_named_toolsets() - - self.stop_conditions = collect_handlers( - self._handler_owners(), - "stop", - self._extra_handlers("stop", builtins=[state_done]), - ) - validate_handler_args( - self.stop_conditions, {"task", "state"}, "stop", "rollout" - ) - self.rollout_setup = collect_handlers( - self._handler_owners(), - "setup", - self._extra_handlers("setup"), - ) - validate_handler_args(self.rollout_setup, {"task", "state"}, "setup", "rollout") - self.rollout_update = collect_handlers( - self._handler_owners(), - "update", - self._extra_handlers("update"), - stage="rollout", - ) - self.group_update = collect_handlers( - self._handler_owners(), - "update", - self._extra_handlers("update"), - stage="group", - ) - validate_handler_args( - self.rollout_update, {"task", "state"}, "update", "rollout" - ) - validate_handler_args(self.group_update, {"tasks", "states"}, "update", "group") - signals = collect_signals( - self._owner_signals(self.taskset), - self._owner_signals(self.harness), - ) - self.rollout_signals = [ - signal for signal in signals if signal["stage"] == "rollout" - ] - self.group_signals = [ - signal for signal in signals if signal["stage"] == "group" - ] - self.rollout_cleanup = collect_handlers( - self._handler_owners(), - "cleanup", - self._extra_handlers("cleanup"), - stage="rollout", - ) - self.group_cleanup = collect_handlers( - self._handler_owners(), - "cleanup", - self._extra_handlers("cleanup"), - stage="group", - ) - validate_handler_args( - self.rollout_cleanup, {"task", "state"}, "cleanup", "rollout" - ) - validate_handler_args( - self.group_cleanup, {"tasks", "states"}, "cleanup", "group" - ) - self.teardown_handlers = collect_handlers( - (self.taskset, self.harness, *self.toolsets), - "teardown", - self._extra_handlers( - "teardown", owners=(self.taskset, self.harness, *self.toolsets) - ), - ) + @abstractmethod + async def stop(self) -> None: ... - self.trajectories: dict[str, list[JsonData]] = {} - self.model_clients: dict[str, Client] = {} - self.owned_model_clients: set[str] = set() - self._model_request_locks: dict[str, asyncio.Lock] = {} - self._inflight_visible_model_requests: dict[str, int] = {} - - self.scoped_tools: dict[tuple[int, str, str], list[ToolEntry]] = {} - self.tool_handles: dict[str, tuple[Task, State, tuple[str, ...]]] = {} - self.runtime_objects: dict[RuntimeObjectOwner, RuntimeObjectStore] = { - "toolset": {}, - "user": {}, - "taskset": {}, - "harness": {}, - } + @abstractmethod + async def expose(self, port: int) -> str: ... - self._sandbox_client = None - self.sandbox_leases: dict[tuple[str, str], SandboxLease] = {} - self.sandbox_creation_tasks: dict[ - tuple[str, str], asyncio.Task[SandboxLease] - ] = {} - self.upload_archive_tasks: dict[tuple[str, str, str], asyncio.Task[Path]] = {} - self.sandbox_lock = asyncio.Lock() - sandbox_config = ( - self.harness.sandbox - if self.harness is not None and self.harness.sandbox is not None - else SandboxConfig() - ) - create_concurrency = sandbox_config.create_concurrency - create_rate = sandbox_config.create_rate_per_second - delete_concurrency = sandbox_config.delete_concurrency - delete_rate = sandbox_config.delete_rate_per_second - self.sandbox_create_semaphore = asyncio.Semaphore(create_concurrency) - self.sandbox_create_rate_limiter = AsyncRateLimiter(create_rate) - self.sandbox_delete_semaphore = asyncio.Semaphore(delete_concurrency) - self.sandbox_delete_rate_limiter = AsyncRateLimiter(delete_rate) - - self.mcp_exit_stacks: dict[str, AsyncExitStack] = {} - self.mcp_tools: dict[str, RuntimeTools] = {} - self.mcp_tool_parents: dict[str, dict[str, tuple[Toolset, ...]]] = {} + async def public_url(self, port: int) -> str | None: + _ = port + return None - @property - def has_group_signals(self) -> bool: - return bool(self.group_signals) + @abstractmethod + async def run( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> CommandResult: ... - @property - def has_group_stage(self) -> bool: - return bool(self.group_update or self.group_signals or self.group_cleanup) + @abstractmethod + async def read(self, path: str) -> bytes: ... - @property - def has_group_rewards(self) -> bool: - return any(signal["kind"] == "reward" for signal in self.group_signals) + @abstractmethod + async def write(self, path: str, data: bytes) -> None: ... - @property - def has_group_advantages(self) -> bool: - return any(signal["kind"] == "advantage" for signal in self.group_signals) - - def prepare_state(self, task: Task, state: State) -> None: - state["task"] = task - state.runtime_state()["runtime_id"] = self.runtime_id - self.resolve_trajectory(state) - self.refresh_tools(state, validate=False) - - def refresh_tools(self, state: State, *, validate: bool = True) -> None: - state["tools"] = sorted(self.all_exposed_tools(state, validate=validate)) - - def task_for_state(self, state: State) -> Task: - return cast(Task, state["task"]) - - def register_tool_handle(self, state: State, names: Sequence[str]) -> str: - task = self.task_for_state(state) - available = self.all_exposed_tools(state) - unknown = sorted(set(names) - set(available)) - if unknown: - raise KeyError(f"Unknown borrowed tools: {unknown}.") - handle_id = uuid.uuid4().hex - self.tool_handles[handle_id] = (task, state, tuple(names)) - return handle_id - - def _borrowed_tool_handle( - self, handle_id: str - ) -> tuple[Task, State, tuple[str, ...]]: - handle = self.tool_handles.get(handle_id) - if handle is None: - raise RuntimeError(f"No live tool handle registered for {handle_id!r}.") - return handle - - def borrowed_tool_def(self, handle_id: str, name: str) -> Tool: - _, source_state, names = self._borrowed_tool_handle(handle_id) - if name not in names: - raise KeyError(f"Tool handle does not expose {name!r}.") - source_tool = self.all_exposed_tools(source_state)[name] - return self.tool_def(name, source_tool, source_state) - - async def call_borrowed_tool( - self, handle_id: str, name: str, **kwargs: object - ) -> object: - source_task, source_state, names = self._borrowed_tool_handle(handle_id) - if name not in names: - raise KeyError(f"Tool handle does not expose {name!r}.") - return await self._call_tool( - name, source_task, source_state, exposed=True, **kwargs + async def run_background( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + log: str | None = None, + ) -> None: + _ = command, cwd, env, log + raise NotImplementedError( + f"{type(self).__name__} does not support background commands." ) - def release_tool_handles(self, state: State) -> None: - for handle_id, (_, source_state, _) in list(self.tool_handles.items()): - if source_state is state: - del self.tool_handles[handle_id] - - def add_tool(self, toolset: str, tool: ToolEntry, state: State) -> None: - if toolset not in self.named_toolsets: - raise KeyError(f"Unknown toolset {toolset!r}.") - if isinstance(tool, Toolset): - raise TypeError("State.add_tool accepts a tool, not a Toolset.") - toolset_value = self.named_toolsets[toolset] - scope = toolset_object_scope(toolset_value) - key = (id(toolset_value), scope, self.scope_key(scope, state)) - tools = self.scoped_tools.setdefault(key, []) - if not isinstance(tool, MCPTool): - name = tool_name(tool) - existing = self.tools_for_toolsets( - [toolset_value], apply_visibility=False, state=state - ) - if name in existing: - raise ValueError(f"Tool {name!r} is defined twice.") - tools.append(tool) - - def release_scoped_tools(self, scope: str, state: State) -> None: - scope_key = self.scope_key(scope, state) - for key in list(self.scoped_tools): - _, tool_scope, tool_scope_key = key - if tool_scope == scope and tool_scope_key == scope_key: - del self.scoped_tools[key] - - def scoped_tool_entries( - self, toolset: Toolset, state: State | None - ) -> list[ToolEntry]: - if state is None: - return [] - scope = toolset_object_scope(toolset) - key = (id(toolset), scope, self.scope_key(scope, state)) - return list(self.scoped_tools.get(key, ())) - - def register_trajectory(self, state: State) -> None: - trajectory = state.get("trajectory") - if trajectory is None: - return - if not isinstance(trajectory, list): - raise TypeError("state.trajectory must be a list.") - self.trajectories[str(state["trajectory_id"])] = cast( - list[JsonData], trajectory - ) + async def __aenter__(self) -> "Runtime": + await self.start() + return self - def visible_model_requests(self, state: State) -> int: - trajectory = state.get("trajectory") or [] - if not isinstance(trajectory, list): - raise TypeError("state.trajectory must be a list.") - trajectory_id = str(state["trajectory_id"]) - completed = sum( - 1 - for step in trajectory - if isinstance(step, dict) - and str(step.get("trajectory_id")) == trajectory_id - ) - return completed + self._inflight_visible_model_requests.get(trajectory_id, 0) + async def __aexit__(self, *exc: object) -> None: + await self.stop() - def resolved_handles(self, state: State) -> ResolvedRuntimeHandlesConfig: - runtime = state.runtime_state() - resolved = runtime.get("resolved") or {} - return ResolvedRuntimeHandlesConfig.model_validate(resolved) - def resolved_runtime(self, handle: RuntimeHandleConfig) -> "Runtime": - return load_runtime(handle.runtime_id) +class RuntimeProvider(ABC): + @abstractmethod + def create_runtime(self) -> Runtime: ... - def model_handle(self, state: State) -> ModelRuntimeHandleConfig | None: - return self.resolved_handles(state).model - def endpoint_handle(self, state: State) -> RuntimeHandleConfig | None: - return self.resolved_handles(state).endpoint +class SubprocessRuntimeProvider(RuntimeProvider): + def __init__(self, config: SubprocessRuntimeConfig | None = None) -> None: + self.config = config or SubprocessRuntimeConfig() - def sandbox_handle(self, state: State) -> SandboxRuntimeHandleConfig | None: - return self.resolved_handles(state).sandbox + def create_runtime(self) -> Runtime: + return SubprocessRuntime(self.config) - def resolve_trajectory(self, state: State) -> None: - handle = self.resolved_handles(state).trajectory - if handle is None: - state.setdefault("trajectory", []) - self.register_trajectory(state) - return - if handle.mode != "append": - raise ValueError("state.runtime.resolved.trajectory.mode must be 'append'.") - runtime = self.resolved_runtime(handle) - trajectory = runtime.trajectories.get(handle.trajectory_id) - if trajectory is None: - raise RuntimeError( - f"No live trajectory registered for {handle.trajectory_id!r}." - ) - state["trajectory"] = trajectory - def bind_model_client( - self, state: State, client: Client | ClientConfig | None - ) -> None: - if client is None: - return - owns_client = False - if isinstance(client, ClientConfig): - resolved_config = resolve_client_config(client) - client = resolve_client(resolved_config) - client_type: ClientType = resolved_config.client_type - owns_client = True - elif isinstance(client, Client): - config = client.config - if isinstance(config, ClientConfig): - client_type = config.client_type - else: - client_type = "openai_chat_completions" - else: - client_type = "openai_chat_completions" - key = str( - state.runtime_state().get("client_key") - or state.get("trajectory_id") - or f"client_{uuid.uuid4().hex}" - ) - self.model_clients[key] = client - if owns_client: - self.owned_model_clients.add(key) - runtime = state.runtime_state() - runtime["client_key"] = key - runtime["client_type"] = client_type - - def model_client(self, state: State) -> Client: - handle = self.model_handle(state) - runtime = self - if handle is None: - key = str(state.runtime_state().get("client_key") or "default") - else: - runtime = self.resolved_runtime(handle) - key = handle.client_key - client = runtime.model_clients.get(key) - if client is None: - raise RuntimeError("Harness has no model client for intercepted requests.") - return client - - def client_type(self, state: State) -> ClientType: - raw_client_type = state.runtime_state().get("client_type") - if raw_client_type is None: - handle = self.model_handle(state) - if handle is not None: - raw_client_type = handle.client_type - if raw_client_type is None: - return "openai_chat_completions" - if raw_client_type not in get_args(ClientType): - raise ValueError(f"Unsupported client type: {raw_client_type!r}") - return cast(ClientType, raw_client_type) - - def model(self, state: State) -> str: - model = state.runtime_state().get("model") - if model is None: - handle = self.model_handle(state) - if handle is not None: - model = handle.model - if not isinstance(model, str) or not model: - raise RuntimeError("Harness has no model for intercepted requests.") - return model - - def sampling_args(self, state: State) -> SamplingArgs: - sampling = state.runtime_state().get("sampling_args") or {} - if not sampling: - handle = self.model_handle(state) - if handle is not None: - sampling = handle.sampling_args or {} - if not isinstance(sampling, dict): - raise TypeError("state.runtime.sampling_args must be a mapping.") - return cast(SamplingArgs, dict(cast(ConfigData, sampling))) - - def tool_defs(self, state: State) -> list[Tool] | None: - defs: list[Tool] = [] - for name, tool in self.all_exposed_tools(state).items(): - if ( - isinstance(tool, Tool) - or isinstance(tool, ToolDefinitionProvider) - or callable(tool) - ): - defs.append(self.tool_def(name, tool, state)) - return defs or None - - async def user_messages( - self, - task: Task, - state: State, - transcript: Sequence[PromptMessage] | None = None, - ) -> list[JsonData]: - user = self.active_user() - if user is None: - return [] - kwargs: RuntimeData = {} - if user.sandbox is not None: - kwargs["sandbox"] = await self.resolve_user_sandbox(user, task, state) - for name, source in user.bindings.items(): - validate_bound_arg(user.get_response, name, f"User binding {name!r}") - validate_binding_source(source, f"User binding {name!r}") - kwargs[name] = await self.resolve_user_binding( - user, source, task, state, transcript - ) - raw_messages = await maybe_call_with_named_args( - user.get_response, - task=task, - state=state, - messages=state_messages(state, transcript), - **kwargs, - ) - if raw_messages is None: - return [] - messages = normalize_messages(raw_messages, field_name="user") - return [message.model_dump(exclude_none=True) for message in messages] - - def active_user(self) -> User | None: - users = [] - if self.taskset is not None and self.taskset.user is not None: - users.append(self.taskset.user) - if self.harness is not None and self.harness.user is not None: - users.append(self.harness.user) - if len(users) > 1: - raise ValueError("Taskset and harness cannot both define user.") - if not users: - return None - return users[0] - - def tool_def(self, name: str, tool: object, state: State) -> Tool: - hidden_args = self.hidden_tool_args(name, state) - if isinstance(tool, Tool): - return tool_schema(tool, hidden_args) - if isinstance(tool, ToolDefinitionProvider): - return tool_schema(tool.tool_def, hidden_args) - schema_tool = tool - filtered_signature = self._tool_signature(name, tool, state) - if hidden_args and filtered_signature is not None: - schema_tool = schema_callable(tool, filtered_signature) - return tool_schema(convert_func_to_tool_def(schema_tool), hidden_args) - - def hidden_tool_args(self, name: str, state: State) -> set[str]: - hidden_args = {"runtime", "task", "state"} - owner = self.tool_owner(name, state) - if owner is not None and owner.sandbox is not None: - hidden_args.add("sandbox") - if owner is not None: - for binding_key in owner.bindings: - tool_name_prefix, arg_name = binding_key_parts(binding_key) - if tool_name_prefix == name: - tool = self.all_tools(state)[name] - target = owner.handler if isinstance(tool, Tool) else tool - if target is None: - raise TypeError( - f"Schema-backed tool {name!r} requires a Toolset handler." - ) - validate_bound_arg( - target, - arg_name, - f"Tool binding {binding_key!r}", - ) - hidden_args.add(arg_name) - return hidden_args - - async def call_tool( - self, tool_name: str, task: Task, state: State, **kwargs: object - ) -> object: - return await self._call_tool(tool_name, task, state, True, **kwargs) - - async def is_completed(self, task: Task, state: State) -> bool: - conditions = unique_handlers( - [*self.stop_conditions, *self._rollout_handlers("stop", state)] - ) - for condition in conditions: - framework_kwargs = rollout_framework_kwargs(task, state) - extra_kwargs = await self.binding_kwargs( - condition, task, state, set(framework_kwargs) - ) - completed = await maybe_call_with_named_args( - condition, **extra_kwargs, **framework_kwargs - ) - if completed: - state._set_completed(True) - state._set_truncated( - any( - step.get("is_truncated", False) - for step in state.get("trajectory", []) - if isinstance(step, dict) - ) - ) - state._set_stop_condition(function_name(condition)) - return True - return False - - def tool_calls(self, task: Task, state: State) -> dict[str, RuntimeCallable]: - return { - name: self._tool_call(name, task, state, exposed=True) - for name in self.all_exposed_tools(state) - } +class DockerRuntimeProvider(RuntimeProvider): + def __init__(self, config: DockerRuntimeConfig) -> None: + self.config = config - def _tool_call( - self, tool_name: str, task: Task, state: State, exposed: bool - ) -> RuntimeCallable: - async def call(**kwargs: object) -> object: - return await self._call_tool(tool_name, task, state, exposed, **kwargs) - - tools = self.all_exposed_tools(state) if exposed else self.all_tools(state) - tool = tools[tool_name] - tool_def = self.tool_def(tool_name, tool, state) - call.__name__ = tool_def.name - call.__doc__ = tool_def.description - signature = self._tool_signature(tool_name, tool, state) - if signature is not None: - setattr(call, "__signature__", signature) - return call - - def _tool_signature( - self, tool_name: str, tool: object, state: State - ) -> inspect.Signature | None: - if not callable(tool): - return None - try: - signature = inspect.signature(tool) - except (TypeError, ValueError): - return None - hidden_args = self.hidden_tool_args(tool_name, state) - parameters = [ - parameter - for parameter in signature.parameters.values() - if parameter.name not in hidden_args - ] - return signature.replace(parameters=parameters) + def create_runtime(self) -> Runtime: + return DockerRuntime(self.config) - async def _call_tool_callable( - self, - tool: RuntimeCallable, - tool_name: str, - task: Task, - state: State, - visible_kwargs: RuntimeData, - hidden_kwargs: RuntimeData, - ) -> object: - call_kwargs = dict(visible_kwargs) - try: - signature = inspect.signature(tool) - except (TypeError, ValueError): - if hidden_kwargs: - raise TypeError( - f"Tool {tool_name!r} uses hidden args, but its signature " - "cannot be inspected." - ) - result = tool(**call_kwargs) - if inspect.isawaitable(result): - return await result - return result - parameters = signature.parameters - hidden_values: RuntimeData = { - "task": task, - "state": state, - **hidden_kwargs, - } - for arg_name, value in hidden_values.items(): - if arg_name in parameters: - call_kwargs[arg_name] = value - elif arg_name in hidden_kwargs: - raise TypeError( - f"Tool {tool_name!r} has hidden arg {arg_name!r}, but does " - "not declare it in its signature." - ) - result = tool(**call_kwargs) - if inspect.isawaitable(result): - return await result - return result - async def _call_schema_tool( - self, - tool: Tool, - handler: RuntimeCallable, - tool_name: str, - task: Task, - state: State, - visible_kwargs: RuntimeData, - hidden_kwargs: RuntimeData, - ) -> object: - call_kwargs: RuntimeData = { - "tool": tool, - "tool_name": tool_name, - "arguments": dict(visible_kwargs), - "task": task, - "state": state, - "runtime": self, - **hidden_kwargs, - } - try: - signature = inspect.signature(handler) - except (TypeError, ValueError): - if hidden_kwargs: - raise TypeError( - f"Toolset handler for {tool_name!r} uses hidden args, but its " - "signature cannot be inspected." - ) - result = handler(tool=tool, arguments=dict(visible_kwargs)) - if inspect.isawaitable(result): - return await result - return result - if not any( - parameter.kind == parameter.VAR_KEYWORD - for parameter in signature.parameters.values() - ): - for arg_name in hidden_kwargs: - if arg_name not in signature.parameters: - raise TypeError( - f"Toolset handler for {tool_name!r} has hidden arg " - f"{arg_name!r}, but does not declare it in its signature." - ) - return await maybe_call_with_named_args(handler, **call_kwargs) - - async def _call_tool( - self, - tool_name: str, - task: Task, - state: State, - exposed: bool, - **kwargs: object, - ) -> object: - tools = self.all_exposed_tools(state) if exposed else self.all_tools(state) - if tool_name not in tools: - kind = "exposed tool" if exposed else "tool" - raise KeyError(f"Unknown {kind} {tool_name!r}.") - visible_kwargs = dict(kwargs) - hidden_kwargs: RuntimeData = {} - owner = self.tool_owner(tool_name, state) - for hidden_arg in ("runtime", "task", "state"): - if hidden_arg in visible_kwargs: - raise ValueError(f"Tool arg {tool_name}.{hidden_arg} is reserved.") - if owner is not None and owner.sandbox is not None: - if "sandbox" in visible_kwargs: - raise ValueError(f"Tool arg {tool_name}.sandbox is reserved.") - hidden_kwargs["sandbox"] = await self.resolve_tool_sandbox( - owner, task, state - ) - for binding_key, source in ( - owner.bindings if owner is not None else {} - ).items(): - tool_name_prefix, arg_name = binding_key_parts(binding_key) - if tool_name_prefix != tool_name: - continue - if arg_name in visible_kwargs or arg_name in hidden_kwargs: - raise ValueError(f"Tool arg {tool_name}.{arg_name} is already set.") - hidden_kwargs[arg_name] = await self.resolve_tool_binding( - owner, source, task, state - ) - tool = tools[tool_name] - if isinstance(tool, Tool): - if owner is None or owner.handler is None: - raise TypeError( - f"Schema-backed tool {tool_name!r} requires a Toolset handler." - ) - return await self._call_schema_tool( - tool, - owner.handler, - tool_name, - task=task, - state=state, - visible_kwargs=visible_kwargs, - hidden_kwargs=hidden_kwargs, - ) - if not callable(tool): - raise TypeError(f"Tool {tool_name!r} must be callable or schema-backed.") - return await self._call_tool_callable( - cast(RuntimeCallable, tool), - tool_name, - task=task, - state=state, - visible_kwargs=visible_kwargs, - hidden_kwargs=hidden_kwargs, - ) +class PrimeRuntimeProvider(RuntimeProvider): + def __init__(self, config: PrimeRuntimeConfig) -> None: + self.config = config - async def submit_model_request( - self, - prompt: Messages, - task: Task, - state: State, - tool_defs: list[Tool] | None = None, - context: ModelRequestContext | None = None, - ) -> Response: - context = context or ModelRequestContext() - reserved = await self._reserve_model_request(task, state, context) - if not reserved: - return self._completed_model_response(state) - released = False - try: - client = self.model_client(state) - request_start = time.time() - response = await client.get_response( - prompt=prompt, - model=self.model(state), - tools=tool_defs, - sampling_args=self.sampling_args(state), - state=state, - ) - request_end = time.time() - state.record_model_timing(request_start, request_end) - record_response_usage(state, response) - completion = await parse_response_message(response) - tokens = await parse_response_tokens(response) - is_truncated = response.message.is_truncated or ( - tokens is not None and bool(tokens.get("is_truncated")) - ) - step = { - "prompt": serializable(prompt), - "completion": serializable(completion), - "response": serializable(response), - "tokens": serializable(tokens), - "reward": None, - "advantage": None, - "is_truncated": bool(is_truncated), - "trajectory_id": str(state["trajectory_id"]), - "extras": context.extras(), - } - if context.trajectory_visibility == "append": - state["trajectory"].append(step) - self._release_model_request(state, context) - released = True - await self.is_completed(task, state) - elif context.trajectory_visibility != "hidden": - raise AssertionError( - f"Unknown trajectory visibility: {context.trajectory_visibility!r}" - ) - return response - finally: - if not released: - self._release_model_request(state, context) - - async def _reserve_model_request( - self, task: Task, state: State, context: ModelRequestContext - ) -> bool: - key = str(state["trajectory_id"]) - lock = self._model_request_locks.setdefault(key, asyncio.Lock()) - async with lock: - if await self.is_completed(task, state): - return False - if context.trajectory_visibility != "append": - return True - self._inflight_visible_model_requests[key] = ( - self._inflight_visible_model_requests.get(key, 0) + 1 - ) - return True + def create_runtime(self) -> Runtime: + return PrimeRuntime(self.config) - def _release_model_request( - self, state: State, context: ModelRequestContext - ) -> None: - if context.trajectory_visibility != "append": - return - key = str(state["trajectory_id"]) - count = self._inflight_visible_model_requests.get(key, 0) - if count <= 1: - self._inflight_visible_model_requests.pop(key, None) - else: - self._inflight_visible_model_requests[key] = count - 1 - - def _completed_model_response(self, state: State) -> Response: - return Response( - id=f"completed_{uuid.uuid4().hex}", - created=int(time.time()), - model=self.model(state), - message=ResponseMessage( - role="assistant", - content="", - finish_reason="stop", - is_truncated=False, - ), - ) - async def setup_rollout( - self, - task: Task, - state: State, - setup_handlers: Iterable[Handler] = (), - **kwargs: object, - ) -> State: - handlers = sort_handlers( - unique_handlers( - [ - *setup_handlers, - *self.rollout_setup, - *self._rollout_handlers("setup", state), - ] - ), - "setup", - ) - validate_handler_args(handlers, {"task", "state"}, "setup", "rollout") - await self.run_rollout_handlers(handlers, task=task, state=state, **kwargs) - await self.ensure_mcp_tools(state) - self.validate_bindings(state) - self.refresh_tools(state) - return state - - async def update_rollout(self, task: Task, state: State) -> State: - handlers = sort_handlers( - unique_handlers( - [ - *self.rollout_update, - *self._rollout_handlers("update", state, stage="rollout"), - ] - ), - "update", - ) - validate_handler_args(handlers, {"task", "state"}, "update", "rollout") - await self.run_rollout_handlers(handlers, task=task, state=state) - return state - - async def update_group(self, tasks: list[Task], states: list[State]) -> list[State]: - handlers = sort_handlers( - unique_handlers( - [ - *self.group_update, - *self._group_handlers("update", states, stage="group"), - ] - ), - "update", - ) - validate_handler_args(handlers, {"tasks", "states"}, "update", "group") - await self.run_group_handlers(handlers, tasks=tasks, states=states) - return states - - async def score_rollout(self, task: Task, state: State) -> State: - await score_rollout_signals( - self.rollout_signals, - task, - state, - resolve_kwargs=self.binding_kwargs, - ) - return state - - async def score_group(self, tasks: list[Task], states: list[State]) -> list[State]: - await self.update_group(tasks, states) - await score_group_signals( - self.group_signals, - tasks, - states, - resolve_kwargs=self.group_binding_kwargs, - ) - return states +class ModalRuntimeProvider(RuntimeProvider): + def __init__(self, config: ModalRuntimeConfig) -> None: + self.config = config - async def cleanup_rollout(self, task: Task, state: State) -> None: - handlers = unique_handlers( - [ - *self.rollout_cleanup, - *self._rollout_handlers("cleanup", state, stage="rollout"), - ] - ) - await self.run_rollout_handlers(handlers, task=task, state=state) - await self.release_runtime_objects("rollout", state) - await self.release_sandboxes(scope="rollout", state=state) - await self.close_mcp_tools(state) - self.release_scoped_tools("rollout", state) - await self.release_model_client(state) - key = str(state["trajectory_id"]) - self._model_request_locks.pop(key, None) - self._inflight_visible_model_requests.pop(key, None) - self.release_tool_handles(state) - - async def cleanup_group(self, tasks: list[Task], states: list[State]) -> None: - handlers = unique_handlers( - [ - *self.group_cleanup, - *self._group_handlers("cleanup", states, stage="group"), - ] - ) - await self.run_group_handlers(handlers, tasks=tasks, states=states) - for state in states: - await self.release_runtime_objects("group", state) - await self.release_sandboxes(scope="group", state=state) - await self.close_mcp_tools(state, scope="group") - self.release_scoped_tools("group", state) - await self.release_model_client(state, group=True) - - async def collect_artifacts(self, task: Task, state: State) -> None: - artifacts = self.runtime_artifacts(task, state) - if not artifacts: - return - state_artifacts = state.setdefault("artifacts", {}) - if not isinstance(state_artifacts, dict): - raise TypeError("state.artifacts must be a mapping.") - for name in artifacts: - if name in state_artifacts: - raise ValueError(f"Artifact {name!r} is already present on state.") - values = await asyncio.gather( - *( - self.collect_runtime_artifact(artifact, task, state) - for artifact in artifacts.values() - ) - ) - state_artifacts.update(dict(zip(artifacts, values, strict=True))) - - def runtime_artifacts(self, task: Task, state: State) -> dict[str, RuntimeArtifact]: - artifacts: dict[str, RuntimeArtifact] = {} - - sources: list[tuple[ArtifactOwner, dict[str, ArtifactConfig]]] = [] - if self.taskset is not None: - sources.append((self.taskset, self.taskset.artifacts)) - if self.harness is not None: - sources.append((self.harness, self.harness.artifacts)) - sources.append( - ( - None, - self.harness.program_config.artifacts.artifacts( - "harness.program.artifacts" - ), - ) - ) - for toolset in iter_toolsets(self.active_toolsets(state)): - sources.append((toolset, toolset.artifacts)) - user = self.active_user() - if user is not None: - sources.append((user, user.artifacts)) - - for owner, source in sources: - for name, artifact in source.items(): - if name in artifacts: - raise ValueError(f"Artifact {name!r} is defined twice.") - artifacts[name] = RuntimeArtifact( - name=name, config=artifact, owner=owner - ) - for name, artifact in ( - task.artifacts_config().artifacts("task.artifacts").items() - ): - if name in artifacts: - raise ValueError(f"Artifact {name!r} is defined twice.") - artifacts[name] = RuntimeArtifact(name=name, config=artifact) - return artifacts - - async def teardown(self) -> None: - await run_handlers(self.teardown_handlers) - await self.release_runtime_objects() - failed_sandbox_deletions = [] - try: - if self.sandbox_creation_tasks: - await self.clear_sandbox_creation_tasks( - list(self.sandbox_creation_tasks.items()) - ) - for key, handle in list(self.sandbox_leases.items()): - try: - await self.close_sandbox_lease(handle) - except Exception as exc: - logger.warning( - "Failed to delete sandbox %s during teardown: %s", - handle.id, - exc, - exc_info=True, - ) - failed_sandbox_deletions.append(handle) - else: - async with self.sandbox_lock: - if self.sandbox_leases.get(key) is handle: - del self.sandbox_leases[key] - finally: - if not any( - handle.client is self._sandbox_client - for handle in failed_sandbox_deletions - ): - await self.teardown_sandbox_client() - await self.cleanup_upload_archives() - self.tool_handles.clear() - self.scoped_tools.clear() - await self.close_all_mcp_tools() - await self.release_all_model_clients() - if not failed_sandbox_deletions: - unregister_runtime(self.runtime_id) - - def sandbox_client(self) -> "SandboxClient": - if self._sandbox_client is None: - from verifiers.utils.threaded_sandbox_client import ( - ThreadedAsyncSandboxClient, - ) + def create_runtime(self) -> Runtime: + return ModalRuntime(self.config) - from .utils.sandbox_utils import SandboxClient - self._sandbox_client = cast( - SandboxClient, - ThreadedAsyncSandboxClient(), - ) - return self._sandbox_client +class DaytonaRuntimeProvider(RuntimeProvider): + def __init__(self, config: DaytonaRuntimeConfig) -> None: + self.config = config - async def teardown_sandbox_client(self) -> None: - if self._sandbox_client is None: - return - from .utils.sandbox_utils import close_sandbox_client - - await close_sandbox_client(self._sandbox_client) - self._sandbox_client = None - - async def cached_upload_archive( - self, local_source: Path | Traversable, remote_path: str - ) -> Path: - from .utils.sandbox_utils import UPLOAD_IGNORE_PARTS, build_dir_archive - - if isinstance(local_source, Path): - root = local_source.resolve() - digest = hashlib.sha256() - paths = [root] - if root.is_dir(): - paths = [] - directories = [root] - while directories: - for path in sorted(directories.pop().iterdir()): - if path.name in UPLOAD_IGNORE_PARTS: - continue - paths.append(path) - if path.is_dir(): - directories.append(path) - for path in paths: - relative = ( - path.relative_to(root).as_posix() if path != root else path.name - ) - stat = path.stat() - kind = "d" if path.is_dir() else "f" - digest.update( - f"{kind}:{relative}:{stat.st_mtime_ns}:{stat.st_size}\0".encode() - ) - key = (remote_path, str(root), digest.hexdigest()) - else: - key = (remote_path, str(local_source), "resource") - task = self.upload_archive_tasks.get(key) - if task is None: - task = asyncio.create_task( - asyncio.to_thread(build_dir_archive, local_source, remote_path) - ) - self.upload_archive_tasks[key] = task - try: - return await asyncio.shield(task) - except Exception: - if self.upload_archive_tasks.get(key) is task: - del self.upload_archive_tasks[key] - raise + def create_runtime(self) -> Runtime: + return DaytonaRuntime(self.config) - async def cleanup_upload_archives(self) -> None: - tasks = list(self.upload_archive_tasks.values()) - results = await asyncio.gather(*tasks, return_exceptions=True) - self.upload_archive_tasks.clear() - for result in results: - if not isinstance(result, Path): - continue - try: - result.unlink(missing_ok=True) - except Exception as exc: - logger.warning( - "Failed to delete cached upload archive %s: %s", - result, - exc, - exc_info=True, - ) - async def close_sandbox_lease(self, lease: "SandboxLease") -> None: - async with self.sandbox_delete_semaphore: - await self.sandbox_delete_rate_limiter.wait() - await close_object(lease) +def make_runtime_provider(config: RuntimeConfigValue) -> RuntimeProvider: + if isinstance(config, PrimeRuntimeConfig): + return PrimeRuntimeProvider(config) + if isinstance(config, DockerRuntimeConfig): + return DockerRuntimeProvider(config) + if isinstance(config, ModalRuntimeConfig): + return ModalRuntimeProvider(config) + if isinstance(config, DaytonaRuntimeConfig): + return DaytonaRuntimeProvider(config) + return SubprocessRuntimeProvider(config) - async def clear_sandbox_creation_tasks( + +def resolve_runtime_config( + *configs: RuntimeConfigValue | None, +) -> RuntimeConfigValue: + resolved: RuntimeConfigValue | None = None + for config in configs: + if config is None: + continue + if resolved is None: + resolved = config.model_copy() + continue + if type(config) is not type(resolved): + raise ValueError( + "Runtime config resolution requires a single provider type; got " + f"{resolved.type!r} and {config.type!r}." + ) + updates = { + field: getattr(config, field) + for field in config.model_fields_set + if field != "type" + } + if updates: + resolved = resolved.model_copy(update=updates) + return resolved or SubprocessRuntimeConfig() + + +class SubprocessRuntime(Runtime): + def __init__(self, config: SubprocessRuntimeConfig) -> None: + self.config = config + self.workdir: Path | None = None + self.processes: list[asyncio.subprocess.Process] = [] + + async def start(self) -> None: + self.workdir = Path(tempfile.mkdtemp(prefix="vf-v1-", dir="/tmp")) + + async def stop(self) -> None: + for process in self.processes: + if process.returncode is None: + with contextlib.suppress(ProcessLookupError): + process.send_signal(signal.SIGINT) + for process in self.processes: + if process.returncode is None: + with contextlib.suppress(Exception): + await asyncio.wait_for(process.wait(), timeout=5) + for process in self.processes: + if process.returncode is None: + with contextlib.suppress(ProcessLookupError): + process.terminate() + for process in self.processes: + if process.returncode is None: + with contextlib.suppress(Exception): + await asyncio.wait_for(process.wait(), timeout=5) + if process.returncode is None: + with contextlib.suppress(ProcessLookupError): + process.kill() + self.processes = [] + if self.workdir is not None: + await asyncio.to_thread(shutil.rmtree, self.workdir, True) + self.workdir = None + + async def expose(self, port: int) -> str: + return f"http://127.0.0.1:{port}" + + async def public_url(self, port: int) -> str | None: + return f"http://127.0.0.1:{port}" + + async def run( self, - creations: Sequence[tuple[tuple[str, str], asyncio.Task["SandboxLease"]]], + command: list[str], *, - state: State | None = None, - scope: str | None = None, - ) -> None: - from .utils.sandbox_utils import SandboxLease as SandboxLeaseClass - - claimed_creations = [] - async with self.sandbox_lock: - for key, task in creations: - if self.sandbox_creation_tasks.get(key) is task: - del self.sandbox_creation_tasks[key] - claimed_creations.append((key, task)) - - for _, task in claimed_creations: - if not task.done(): - task.cancel() - results = await asyncio.gather( - *(task for _, task in claimed_creations), return_exceptions=True + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> CommandResult: + if not command: + raise ValueError("Runtime.run requires a command.") + full_env = { + name: os.environ[name] for name in _SUBPROCESS_ENV if name in os.environ + } + full_env.update(env or {}) + process = await asyncio.create_subprocess_exec( + *command, + cwd=self._path(cwd) if cwd is not None else self.workdir, + env=full_env, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, ) - - for (key, _), result in zip(claimed_creations, results, strict=True): - if not isinstance(result, SandboxLeaseClass): - continue - result.scope_key = key[0] - try: - await self.close_sandbox_lease(result) - except Exception as exc: - async with self.sandbox_lock: - if not result.deleted: - self.sandbox_leases[key] = result - logger.warning( - "Failed to delete sandbox %s from cancelled creation: %s", - result.id, - exc, - exc_info=True, - ) - if state is not None and scope is not None: - cleanup_errors = cast( - list[ConfigData], state.setdefault("cleanup_errors", []) - ) - cleanup_errors.append( - { - "type": type(exc).__name__, - "message": str(exc), - "scope": scope, - } - ) - - async def resolve_sandbox_lease( - self, key: tuple[str, str], factory: Callable[[], Awaitable["SandboxLease"]] - ) -> "SandboxLease": - async with self.sandbox_lock: - lease = self.sandbox_leases.get(key) - if lease is not None: - if lease.deleted: - raise RuntimeError("Sandbox lease is being deleted.") - return lease - task = self.sandbox_creation_tasks.get(key) - if task is None: - - async def create_sandbox_lease() -> "SandboxLease": - async with self.sandbox_create_semaphore: - await self.sandbox_create_rate_limiter.wait() - return await factory() - - task = asyncio.create_task(create_sandbox_lease()) - self.sandbox_creation_tasks[key] = task - try: - lease = await asyncio.shield(task) - except asyncio.CancelledError: - if task.cancelled(): - async with self.sandbox_lock: - if self.sandbox_creation_tasks.get(key) is task: - del self.sandbox_creation_tasks[key] - raise - except BaseException: - async with self.sandbox_lock: - if self.sandbox_creation_tasks.get(key) is task: - del self.sandbox_creation_tasks[key] - raise - async with self.sandbox_lock: - existing = self.sandbox_leases.get(key) - if existing is not None: - if existing.deleted: - raise RuntimeError("Sandbox lease is being deleted.") - return existing - if self.sandbox_creation_tasks.get(key) is not task: - raise RuntimeError( - "Sandbox creation was cancelled before the lease was resolved." - ) - del self.sandbox_creation_tasks[key] - if lease.deleted: - raise RuntimeError( - "Sandbox lease was deleted before it could be resolved." - ) - lease.scope_key = key[0] - self.sandbox_leases[key] = lease - return lease - - async def run_rollout_handlers( + stdout, stderr = await asyncio.wait_for(process.communicate(), timeout) + except asyncio.TimeoutError as exc: + with contextlib.suppress(ProcessLookupError): + process.kill() + await process.wait() + raise TimeoutError( + f"Runtime command timed out after {timeout} seconds." + ) from exc + return CommandResult( + returncode=int(process.returncode or 0), + stdout=stdout.decode(errors="replace"), + stderr=stderr.decode(errors="replace"), + ) + + async def read(self, path: str) -> bytes: + return await asyncio.to_thread(self._path(path).read_bytes) + + async def write(self, path: str, data: bytes) -> None: + target = self._path(path) + target.parent.mkdir(parents=True, exist_ok=True) + await asyncio.to_thread(target.write_bytes, data) + + async def run_background( self, - handlers: Iterable[Handler], - task: Task, - state: State, - **kwargs: object, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + log: str | None = None, ) -> None: - for handler in handlers: - framework_kwargs = rollout_framework_kwargs(task, state) - protected_args = set(framework_kwargs) | set(kwargs) - extra_kwargs = await self.binding_kwargs( - handler, task, state, protected_args - ) - await maybe_call_with_named_args( - handler, **extra_kwargs, **kwargs, **framework_kwargs + if not command: + raise ValueError("Runtime.run_background requires a command.") + full_env = { + name: os.environ[name] for name in _SUBPROCESS_ENV if name in os.environ + } + full_env.update(env or {}) + stdout = stderr = asyncio.subprocess.DEVNULL + log_file = None + if log is not None: + log_path = self._path(log) + log_path.parent.mkdir(parents=True, exist_ok=True) + log_file = log_path.open("ab") + stdout = stderr = log_file + try: + process = await asyncio.create_subprocess_exec( + *command, + cwd=self._path(cwd) if cwd is not None else self.workdir, + env=full_env, + stdout=stdout, + stderr=stderr, ) + finally: + if log_file is not None: + log_file.close() + self.processes.append(process) - async def run_group_handlers( - self, - handlers: Iterable[Handler], - tasks: list[Task], - states: list[State], - **kwargs: object, - ) -> None: - for handler in handlers: - framework_kwargs = group_framework_kwargs(tasks, states) - protected_args = set(framework_kwargs) | set(kwargs) - extra_kwargs = await self.group_binding_kwargs( - handler, - tasks, - states, - protected_args, - ) - await maybe_call_with_named_args( - handler, **extra_kwargs, **kwargs, **framework_kwargs - ) + def _path(self, path: str | None) -> Path: + if self.workdir is None: + raise RuntimeError("Subprocess runtime has not started.") + if path is None: + return self.workdir + value = Path(path) + return value if value.is_absolute() else self.workdir / value - async def binding_kwargs( - self, - fn: Handler, - task: Task, - state: State, - protected_args: set[str] | None = None, - ) -> RuntimeData: - name = function_name(fn) - kwargs: RuntimeData = {} - protected = protected_args or set() - for binding_key, source, owner in self._binding_entries_for_callable(fn, state): - prefix, arg_name = binding_key_parts(binding_key) - if prefix != name: - continue - if arg_name in protected: - continue - validate_bound_arg(fn, arg_name, f"Binding {binding_key!r}", protected) - if arg_name in kwargs: - raise ValueError(f"Binding arg {arg_name!r} is defined twice.") - kwargs[arg_name] = await self.resolve_owner_binding( - owner, source, task, state - ) - return kwargs - async def group_binding_kwargs( - self, - fn: Handler, - tasks: list[Task], - states: list[State], - protected_args: set[str] | None = None, - ) -> RuntimeData: - if not states: - return {} - state = states[0] - name = function_name(fn) - kwargs: RuntimeData = {} - protected = protected_args or set() - for binding_key, source, owner in self._binding_entries_for_callable(fn, state): - prefix, arg_name = binding_key_parts(binding_key) - if prefix != name: - continue - if arg_name in protected: - continue - validate_bound_arg(fn, arg_name, f"Binding {binding_key!r}", protected) - if arg_name in kwargs: - raise ValueError(f"Binding arg {arg_name!r} is defined twice.") - kwargs[arg_name] = await self.resolve_group_binding( - owner, - source, - tasks, - states, - state, - ) - return kwargs - - async def resolve_owner_binding( - self, owner: BindingOwner, source: BindingSource, task: Task, state: State - ) -> object: - if isinstance(owner, Toolset): - return await self.resolve_tool_binding(owner, source, task, state) - if self.taskset is not None and owner is self.taskset: - return await self.resolve_attached_owner_binding( - self.taskset, source, task, state - ) - if self.harness is not None and owner is self.harness: - return await self.resolve_attached_owner_binding( - self.harness, source, task, state - ) - return await self.resolve_binding(source, task, state) +class DockerRuntime(Runtime): + def __init__(self, config: DockerRuntimeConfig) -> None: + self.config = config + self.container: str | None = None - async def resolve_owner_object( - self, - owner: RuntimeOwnerMixin | Toolset | User, - name: str, - task: Task, - state: State, - ) -> object: - if owner is self.taskset: - return await self.resolve_taskset_object(name, task, state) - if owner is self.harness: - return await self.resolve_harness_object(name, task, state) - if isinstance(owner, Toolset): - return await self.resolve_toolset_object(owner, name, task, state) - if isinstance(owner, User): - return await self.resolve_user_object(owner, name, task, state) - raise RuntimeError("Runtime object owner is not attached to this runtime.") - - async def resolve_attached_owner_binding( - self, - owner: RuntimeOwnerMixin, - source: BindingSource, - task: Task, - state: State, - ) -> object: - if isinstance(source, str): - root, separator, tail = source.partition(".") - if root == "objects": - if not separator: - raise ValueError("objects binding sources must name an object.") - name, _, rest = tail.partition(".") - value = await self.resolve_owner_object(owner, name, task, state) - return read_path(value, rest) if rest else value - return await self.resolve_binding(source, task, state) - - async def resolve_binding( - self, source: BindingSource, task: Task, state: State - ) -> object: - if isinstance(source, str): - if binding_source_root(source) == "objects": - raise ValueError( - "objects.* bindings are private to the owning Taskset, Harness, " - "Toolset, or User callable. Use taskset.objects.* or " - "harness.objects.* for explicit cross-owner object bindings." - ) - root, separator, tail = source.partition(".") - if root in {"taskset", "harness"}: - if not separator: - raise ValueError( - f"{root} binding sources must use {root}.objects.name." - ) - return await self.resolve_runtime_owner_path(root, tail, task, state) - return self.resolve_path(source, task, state) - if isinstance(source, dict) and "fn" in source: - spec = cast(ConfigData, source) - validate_callable_source(spec, "Callable binding source") - fn = resolve_config_object(spec["fn"]) - if not callable(fn): - raise TypeError("Callable binding source requires callable fn.") - return await maybe_call_with_named_args(fn, task=task, state=state) - if callable(source): - return await maybe_call_with_named_args(source, task=task, state=state) - raise TypeError("Binding sources must be framework paths or callables.") - - async def resolve_group_binding( - self, - owner: BindingOwner, - source: BindingSource, - tasks: list[Task], - states: list[State], - state: State, - ) -> object: - if isinstance(source, str): - root, separator, tail = source.partition(".") - if root == "tasks": - return read_path(tasks, tail) if separator else tasks - if root == "states": - return read_path(states, tail) if separator else states - if root in {"task", "state", "tools"}: - raise ValueError("Group handler bindings must use tasks or states.") - if root == "runtime": - runtime = state.runtime_state() - return read_path(runtime, tail) if separator else runtime - if root == "objects": - if not separator: - raise ValueError("objects binding sources must name an object.") - name, _, rest = tail.partition(".") - if owner is self.taskset: - value = await self.resolve_taskset_object(name, tasks[0], state) - elif owner is self.harness: - value = await self.resolve_harness_object(name, tasks[0], state) - elif isinstance(owner, Toolset): - if toolset_object_scope(owner) == "rollout": - raise ValueError( - "objects.* group bindings require a group or global Toolset scope." - ) - value = await self.resolve_toolset_object( - owner, name, tasks[0], state - ) - else: - raise ValueError( - "objects.* group bindings require an object owner." - ) - return read_path(value, rest) if rest else value - if root in {"taskset", "harness"}: - if not separator: - raise ValueError( - f"{root} binding sources must use {root}.objects.name." - ) - return await self.resolve_runtime_owner_path( - root, tail, tasks[0], state - ) - if isinstance(source, dict) and "fn" in source: - spec = cast(ConfigData, source) - validate_callable_source(spec, "Callable binding source") - fn = resolve_config_object(spec["fn"]) - if not callable(fn): - raise TypeError("Callable binding source requires callable fn.") - return await maybe_call_with_named_args(fn, tasks=tasks, states=states) - if callable(source): - return await maybe_call_with_named_args(source, tasks=tasks, states=states) - raise TypeError("Binding sources must be framework paths or callables.") - - async def resolve_tool_binding( - self, toolset: Toolset | None, source: BindingSource, task: Task, state: State - ) -> object: - if isinstance(source, str): - root, separator, tail = source.partition(".") - if root == "objects": - if toolset is None: - raise ValueError("objects.* tool bindings require a Toolset owner.") - if not separator: - raise ValueError("objects binding sources must name an object.") - name, _, rest = tail.partition(".") - value = await self.resolve_toolset_object(toolset, name, task, state) - if rest: - return read_path(value, rest) - return value - return await self.resolve_binding(source, task, state) - - async def resolve_user_binding( - self, - user: User, - source: BindingSource, - task: Task, - state: State, - transcript: Sequence[PromptMessage] | None = None, - ) -> object: - if isinstance(source, str): - root, separator, tail = source.partition(".") - if root == "objects" and separator: - name, _, rest = tail.partition(".") - if name in user.objects: - value = await self.resolve_user_object(user, name, task, state) - else: - raise KeyError(f"Unknown user object {name!r}.") - if rest: - return read_path(value, rest) - return value - if callable(source): - return await maybe_call_with_named_args( - source, task=task, state=state, transcript=transcript - ) - return await self.resolve_binding(source, task, state) - - async def resolve_user_object( - self, user: User, name: str, task: Task, state: State - ) -> object: - if user is not self.active_user(): - raise RuntimeError("User object owner is not attached to this runtime.") - if name not in user.objects: - raise KeyError(f"Unknown user object {name!r}.") - key = (id(user), self.scope_key(user.scope, state), name) - store = self.runtime_objects["user"] - if key in store: - return store[key] - spec = user.objects[name] - obj = await resolve_object_factory(spec, f"User object {name!r}") - store[key] = obj - return obj - - async def resolve_taskset_object( - self, name: str, task: Task, state: State - ) -> object: - taskset = self.taskset - if taskset is None: - raise RuntimeError("Taskset objects require a Taskset.") - return await self.resolve_attached_owner_object( - "Taskset", taskset, "taskset", name, task, state - ) + async def start(self) -> None: + try: + version = await docker("version", "--format", "{{.Server.Version}}") + except FileNotFoundError as exc: + raise RuntimeError("docker runtime requires the docker CLI.") from exc + if version.returncode != 0: + detail = (version.stderr or version.stdout).strip() + raise RuntimeError(f"Docker daemon is not reachable: {detail}") + self.container = f"vf-v1-{uuid.uuid4().hex[:12]}" + limits: list[str] = [] + if self.config.cpu_cores is not None: + limits += ["--cpus", str(self.config.cpu_cores)] + if self.config.memory_gb is not None: + limits += ["--memory", f"{self.config.memory_gb}g"] + if self.config.gpu_count: + limits += ["--gpus", str(self.config.gpu_count)] + result = await docker( + "run", + "--detach", + "--network", + "host", + *limits, + "--workdir", + self.config.workdir, + "--name", + self.container, + self.config.image, + "sleep", + "infinity", + ) + if result.returncode != 0: + raise RuntimeError(f"docker run failed: {result.stderr.strip()}") + + async def stop(self) -> None: + if self.container is None: + return + container, self.container = self.container, None + with contextlib.suppress(Exception): + await docker("rm", "--force", container) - async def resolve_harness_object( - self, name: str, task: Task, state: State - ) -> object: - harness = self.harness - if harness is None: - raise RuntimeError("Harness objects require a Harness.") - return await self.resolve_attached_owner_object( - "Harness", harness, "harness", name, task, state - ) + async def expose(self, port: int) -> str: + return f"http://127.0.0.1:{port}" - async def resolve_attached_owner_object( + async def run( self, - label: Literal["Taskset", "Harness"], - owner: RuntimeOwnerMixin, - store_name: Literal["taskset", "harness"], - name: str, - task: Task, - state: State, - ) -> object: - objects = owner.objects - if name not in objects: - raise KeyError(f"Unknown {label} object {name!r}.") - object_bindings: dict[str, BindingSource] = {} - for binding_key, source in owner.bindings.items(): - target_name, arg_name = binding_key_parts(binding_key) - if target_name == name: - object_bindings[arg_name] = source - # Bound object factories can depend on task/state paths, so they are scoped - # to the rollout. Unbound factories are process-global runtime objects. - scope_key = self.scope_key("rollout", state) if object_bindings else "global" - key = (id(owner), scope_key, name) - store = self.runtime_objects[store_name] - if key in store: - return store[key] - kwargs: RuntimeData = {} - for arg_name, source in object_bindings.items(): - kwargs[arg_name] = await self.resolve_owner_binding( - owner, source, task, state + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> CommandResult: + container = self._container() + env_args = [ + arg + for key, value in (env or {}).items() + for arg in ("--env", f"{key}={value}") + ] + return await docker( + "exec", + *env_args, + "--workdir", + cwd or self.config.workdir, + container, + *command, + timeout=timeout, + ) + + async def read(self, path: str) -> bytes: + result = await self.run(["sh", "-c", f"base64 < {shlex.quote(path)}"]) + if result.returncode != 0: + raise RuntimeError(f"read {path!r}: {result.stderr.strip()}") + return base64.b64decode(result.stdout) + + async def write(self, path: str, data: bytes) -> None: + container = self._container() + parent = shlex.quote(str(PurePosixPath(path).parent)) + process = await asyncio.create_subprocess_exec( + "docker", + "exec", + "-i", + "--workdir", + self.config.workdir, + container, + "sh", + "-c", + f"mkdir -p {parent} && cat > {shlex.quote(path)}", + stdin=asyncio.subprocess.PIPE, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + _, stderr = await process.communicate(input=data) + if process.returncode != 0: + raise RuntimeError( + f"write {path!r}: {stderr.decode(errors='replace').strip()}" ) - obj = await resolve_object_factory( - objects[name], f"{label} object {name!r}", kwargs - ) - store[key] = obj - return obj - async def release_runtime_objects( + async def run_background( self, - scope: str | None = None, - state: State | None = None, - owner: RuntimeObjectOwner | None = None, - ) -> None: - scope_key = self.scope_key(scope, state) if scope is not None else None - owners = RUNTIME_OBJECT_OWNERS if owner is None else (owner,) - for owner_name in owners: - store = self.runtime_objects[owner_name] - for key, obj in list(store.items()): - _, object_scope_key, _ = key - if scope_key is not None and object_scope_key != scope_key: - continue - await close_object(obj) - del store[key] - - def validate_bindings( - self, state: State, *, allow_unresolved_tool_bindings: bool = False + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + log: str | None = None, ) -> None: - for owner in (self.taskset, self.harness): - self._validate_owner_bindings(owner) - for toolset in iter_toolsets(self.active_toolsets(state)): - self._validate_toolset_bindings( - toolset, - state, - allow_unresolved=allow_unresolved_tool_bindings, + if not command: + raise ValueError("Runtime.run_background requires a command.") + container = self._container() + env_args = [ + arg + for key, value in (env or {}).items() + for arg in ("--env", f"{key}={value}") + ] + log_path = shlex.quote(log or "/tmp/vf-background.log") + result = await docker( + "exec", + "--detach", + *env_args, + "--workdir", + cwd or self.config.workdir, + container, + "sh", + "-c", + f"exec {shlex.join(command)} >> {log_path} 2>&1", + ) + if result.returncode != 0: + raise RuntimeError(f"docker background command failed: {result.stderr}") + + def _container(self) -> str: + if self.container is None: + raise RuntimeError("Docker runtime has not started.") + return self.container + + +class PrimeRuntime(Runtime): + def __init__(self, config: PrimeRuntimeConfig) -> None: + self.config = config + self.client = None + self.sandbox_id: str | None = None + self.tunnels: list[RuntimeTunnel] = [] + + async def start(self) -> None: + from prime_sandboxes import ( + AdvancedConfigs, + AsyncSandboxClient, + CreateSandboxRequest, + ) + + self.client = AsyncSandboxClient() + timeout = ( + 24 * 60 + if self.config.timeout_minutes == "auto" + else self.config.timeout_minutes + ) + advanced_configs = ( + None + if self.config.idle_timeout_minutes is None + else AdvancedConfigs.model_validate( + {"idle_timeout_minutes": self.config.idle_timeout_minutes} + ) + ) + labels = ["vf-v1-runtime", *self.config.labels] + if evaluation_id := os.environ.get("EVALUATION_ID"): + labels.append(f"eval-{evaluation_id}") + if job_id := os.environ.get("PRIME_JOB_ID"): + labels.append(f"prime-job-{job_id}") + labels = list(dict.fromkeys(labels)) + try: + sandbox = await self.client.create( + CreateSandboxRequest( + name="vf-v1-runtime", + docker_image=self.config.image, + cpu_cores=self.config.cpu_cores, + memory_gb=self.config.memory_gb, + disk_size_gb=self.config.disk_gb, + gpu_count=self.config.gpu_count, + timeout_minutes=timeout, + network_access=self.config.network_access, + vm=self.config.vm, + guaranteed=self.config.guaranteed, + gpu_type=self.config.gpu_type, + region=self.config.region, + advanced_configs=advanced_configs, + labels=labels, + ) + ) + self.sandbox_id = sandbox.id + await self.client.wait_for_creation(self.sandbox_id) + await self.client.run_background_job( + self.sandbox_id, + f"mkdir -p {shlex.quote(self.config.workdir)}", ) - user = self.active_user() - if user is not None: - for name, source in user.bindings.items(): - validate_bound_arg(user.get_response, name, f"User binding {name!r}") - source_root = binding_source_root(source) - validate_binding_source(source, f"User binding {name!r}") - if source_root == "objects": - object_name = binding_object_name(cast(str, source)) - if object_name not in user.objects: - raise KeyError( - f"User binding {name!r} references unknown User object " - f"{object_name!r}." - ) - self.validate_runtime_owner_object_source( - source, f"User binding {name!r}" - ) + except Exception: + await self.stop() + raise - def _validate_owner_bindings(self, owner: RuntimeOwnerMixin | None) -> None: - if owner is None: + async def stop(self) -> None: + for tunnel in self.tunnels: + with contextlib.suppress(Exception): + tunnel.sync_stop() + self.tunnels = [] + client, sandbox_id = self.client, self.sandbox_id + self.client, self.sandbox_id = None, None + if client is None: return - targets = self._owner_binding_targets(owner) - allow_objects = owner in (self.taskset, self.harness) - for binding_key, source in owner.bindings.items(): - target_name, arg_name = binding_key_parts(binding_key) - target = targets.get(target_name) - if target is None: - raise ValueError( - f"Binding {binding_key!r} does not match a Taskset/Harness " - "callable or object factory." - ) - target_kind, fn = target - if target_kind == "object": - protected_args = frozenset() - else: - stage = handler_stage(fn, cast(CallableKind, target_kind)) - protected_args = ( - GROUP_FRAMEWORK_ARGS if stage == "group" else ROLLOUT_FRAMEWORK_ARGS - ) - if arg_name in protected_args: - continue - validate_bound_arg( - fn, - arg_name, - f"Binding {binding_key!r}", - protected_args, - allow_reserved=target_kind == "object", - ) - validate_binding_source( - source, f"Binding {binding_key!r}", allow_objects=allow_objects - ) - self.validate_runtime_owner_object_source( - source, f"Binding {binding_key!r}" - ) - source_root = binding_source_root(source) - if source_root == "objects": - object_name = binding_object_name(cast(str, source)) - if object_name not in owner.objects: - raise KeyError( - f"Binding {binding_key!r} references unknown object " - f"{object_name!r}." - ) - - def _validate_toolset_bindings( - self, toolset: Toolset, state: State, *, allow_unresolved: bool - ) -> None: - targets = self._toolset_binding_targets(toolset, state) - for binding_key, source in toolset.bindings.items(): - target_name, arg_name = binding_key_parts(binding_key) - target = targets.get(target_name) - if target is None: - if allow_unresolved and toolset_object_scope(toolset) == "rollout": - validate_binding_source(source, f"Binding {binding_key!r}") - continue - raise ValueError( - f"Binding {binding_key!r} does not match a callable or object " - "factory owned by the same Toolset." - ) - target_kind, fn = target - validate_bound_arg( - fn, - arg_name, - f"Binding {binding_key!r}", - allow_reserved=target_kind == "object", - ) - source_root = binding_source_root(source) - validate_binding_source(source, f"Binding {binding_key!r}") - self.validate_runtime_owner_object_source( - source, f"Binding {binding_key!r}" - ) - if source_root == "objects" and target_kind != "tool": - raise ValueError( - f"Binding {binding_key!r} uses objects.*, which is only valid " - "for callable tools owned by the same Toolset." - ) - if source_root == "objects": - object_name = binding_object_name(cast(str, source)) - if object_name not in toolset.objects: - raise KeyError( - f"Binding {binding_key!r} references unknown Toolset object " - f"{object_name!r}." - ) - - def _owner_binding_targets( - self, owner: RuntimeOwnerMixin - ) -> dict[str, tuple[str, Handler]]: - targets: dict[str, tuple[str, Handler]] = {} - - def add_target(kind: str, fn: Handler, name: str | None = None) -> None: - name = name or function_name(fn) - existing = targets.get(name) - if existing is not None and not same_callable(existing[1], fn): - raise ValueError( - f"Taskset/Harness binding target {name!r} is defined twice." - ) - targets[name] = (kind, fn) - - for kind in ( - "stop", - "setup", - "update", - "metric", - "reward", - "advantage", - "cleanup", - ): - for fn in lifecycle_handlers(owner, kind): - if callable(fn): - add_target(kind, cast(Handler, fn)) - for _, method in inspect.getmembers(owner, predicate=callable): - for kind in ( - "stop", - "setup", - "update", - "metric", - "reward", - "advantage", - "cleanup", - ): - if handler_is_marked(method, kind): - add_target(kind, cast(Handler, method)) - for name, spec in owner.objects.items(): - if not isinstance(name, str): - raise TypeError("Object names must be strings.") - if callable(spec): - add_target("object", cast(Handler, spec), name) - return targets - - def _binding_entries_for_callable( - self, fn: Handler, state: State - ) -> list[BindingEntry]: - target_name = function_name(fn) - entries: list[BindingEntry] = [] - for owner in (self.taskset, self.harness): - if owner is None: - continue - target = self._owner_binding_targets(owner).get(target_name) - if target is None or not same_callable(target[1], fn): - continue - self._extend_binding_entries(entries, owner.bindings, target_name, owner) - for toolset in iter_toolsets(self.active_toolsets(state)): - target = self._toolset_binding_targets(toolset, state).get(target_name) - if target is None or not same_callable(target[1], fn): - continue - self._extend_binding_entries( - entries, toolset.bindings, target_name, toolset - ) - return entries + if sandbox_id is not None: + with contextlib.suppress(Exception): + await client.delete(sandbox_id) + with contextlib.suppress(Exception): + await client.aclose() + + async def expose(self, port: int) -> str: + from prime_tunnel import Tunnel + + tunnel = Tunnel(local_port=port) + url = str(await tunnel.start()).rstrip("/") + self.tunnels.append(tunnel) + return url + + async def public_url(self, port: int) -> str | None: + if self.client is None or self.sandbox_id is None: + raise RuntimeError("Prime runtime has not started.") + try: + exposed = await self.client.expose(self.sandbox_id, port) + except Exception as exc: + raise RuntimeError( + "Prime port exposure failed. Runtime-placed servers on Prime " + "require sandbox port exposure; use a supported region/port or " + "place the server in a host/subprocess runtime." + ) from exc + return str(exposed.url).rstrip("/") - def validate_runtime_owner_object_source( - self, source: object, context: str - ) -> None: - if not isinstance(source, str): - return - root = binding_source_root(source) - if root == "taskset": - if self.taskset is None: - raise RuntimeError( - f"{context} references taskset, but no Taskset exists." - ) - object_name = owner_object_name(source) - if object_name not in self.taskset.objects: - raise KeyError( - f"{context} references unknown Taskset object {object_name!r}." - ) - if root == "harness": - if self.harness is None: - raise RuntimeError( - f"{context} references harness, but no Harness exists." - ) - object_name = owner_object_name(source) - if object_name not in self.harness.objects: - raise KeyError( - f"{context} references unknown Harness object {object_name!r}." - ) + async def run( + self, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> CommandResult: + if self.client is None or self.sandbox_id is None: + raise RuntimeError("Prime runtime has not started.") + run = self.client.run_background_job( + self.sandbox_id, + shlex.join(command), + working_dir=cwd or self.config.workdir, + env=env or {}, + ) + try: + result = await ( + asyncio.wait_for(run, timeout) if timeout is not None else run + ) + except asyncio.TimeoutError as exc: + raise TimeoutError( + f"Runtime command timed out after {timeout} seconds." + ) from exc + return CommandResult( + returncode=int(result.exit_code or 0), + stdout=result.stdout or "", + stderr=result.stderr or "", + ) + + async def read(self, path: str) -> bytes: + result = await self.run(["sh", "-c", f"base64 {shlex.quote(path)}"]) + if result.returncode != 0: + raise RuntimeError(f"read {path!r}: {result.stderr.strip()}") + return base64.b64decode(result.stdout) + + async def write(self, path: str, data: bytes) -> None: + if self.client is None or self.sandbox_id is None: + raise RuntimeError("Prime runtime has not started.") + target = ( + path + if path.startswith("/") + else f"{self.config.workdir.rstrip('/')}/{path}" + ) + parent = shlex.quote(str(PurePosixPath(target).parent)) + result = await self.run(["sh", "-c", f"mkdir -p {parent}"]) + if result.returncode != 0: + raise RuntimeError(f"write {path!r}: {result.stderr.strip()}") + try: + await self.client.upload_bytes( + self.sandbox_id, + target, + data, + filename=PurePosixPath(target).name, + ) + except Exception as exc: + raise RuntimeError(f"write {path!r}: {exc}") from exc - def _extend_binding_entries( + async def run_background( self, - entries: list[BindingEntry], - bindings: dict[str, BindingSource], - target_name: str, - owner: BindingOwner = None, + command: list[str], + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + log: str | None = None, ) -> None: - existing = {key for key, _, _ in entries} - for binding_key, source in bindings.items(): - prefix, _ = binding_key_parts(binding_key) - if prefix != target_name: - continue - if binding_key in existing: - raise ValueError(f"Binding {binding_key!r} is defined twice.") - existing.add(binding_key) - entries.append((binding_key, source, owner)) - - def _toolset_binding_targets( - self, toolset: Toolset, state: State | None = None - ) -> dict[str, tuple[str, Handler]]: - targets: dict[str, tuple[str, Handler]] = {} - - def add_target(name: str, kind: str, fn: Handler) -> None: - if name in targets: - raise ValueError(f"Toolset binding target {name!r} is defined twice.") - targets[name] = (kind, fn) - - for item in self._toolset_entries(toolset, state): - if isinstance(item, Toolset | MCPTool): - continue - if isinstance(item, Tool): - if toolset.handler is None: - raise TypeError( - f"Schema-backed tool {item.name!r} requires a Toolset handler." - ) - add_target(item.name, "tool", toolset.handler) - continue - if callable(item): - add_target(tool_name(item), "tool", cast(Handler, item)) - for name, spec in toolset.objects.items(): - if callable(spec): - add_target(name, "object", cast(Handler, spec)) - for kind in ("stop", "setup", "update", "cleanup"): - for fn in lifecycle_handlers(toolset, kind): - if callable(fn): - add_target( - function_name(fn), - kind, - cast(Handler, fn), - ) - for _, method in inspect.getmembers(toolset, predicate=callable): - if any( - handler_is_marked(method, attr) - for attr in ("stop", "setup", "update", "cleanup") - ): - handler = cast(Handler, method) - add_target( - function_name(handler), - "handler", - handler, - ) - return targets + log_path = shlex.quote(log or "/tmp/vf-background.log") + result = await self.run( + [ + "sh", + "-c", + f"nohup {shlex.join(command)} >> {log_path} 2>&1 &", + ], + cwd=cwd, + env=env, + ) + if result.returncode != 0: + raise RuntimeError(f"prime background command failed: {result.stderr}") - def active_toolsets(self, state: State) -> list[Toolset]: - task = self.task_for_state(state) - selected = self._selected_toolset_names(task) - ids_to_names = { - id(toolset): name for name, toolset in self.named_toolsets.items() - } - active: list[Toolset] = [] - for toolset in self.toolsets: - name = ids_to_names.get(id(toolset)) - if name is not None and name not in selected: - continue - active.append(toolset) - return active - - def _selected_toolset_names(self, task: Task) -> set[str]: - names = set(self.named_toolsets) - config = task.toolsets_config() - show = config.show - hide = config.hide - if show is not None: - selected = set(show) - unknown = sorted(selected - names) - if unknown: - raise KeyError(f"Unknown shown toolsets: {unknown}.") - return selected - if hide is not None: - hidden = set(hide) - unknown = sorted(hidden - names) - if unknown: - raise KeyError(f"Unknown hidden toolsets: {unknown}.") - return names - hidden - return names - - def _collect_named_toolsets(self) -> dict[str, Toolset]: - named: dict[str, Toolset] = {} - for owner in (self.taskset, self.harness): - if owner is None: - continue - for name, toolset in owner.named_toolsets.items(): - if not isinstance(name, str): - raise TypeError("Toolset names must be strings.") - if name in named: - raise ValueError(f"Toolset {name!r} is defined twice.") - if not isinstance(toolset, Toolset): - raise TypeError("named_toolsets values must be Toolsets.") - named[name] = toolset - return named - - def tools_for_toolsets( - self, - toolsets: Iterable[ToolEntry], - apply_visibility: bool, - state: State | None = None, - tool_filters: dict[str, VisibilityConfig] | None = None, - ) -> RuntimeTools: - tools: RuntimeTools = {} - - def visit(item: ToolEntry, parents: list[Toolset]) -> None: - if isinstance(item, Toolset): - for child in self._toolset_entries(item, state): - visit(child, [*parents, item]) - return - if isinstance(item, MCPTool): - return - name = tool_name(item) - if apply_visibility and not all( - self._tool_visible_for_task(toolset, name, tool_filters) - for toolset in parents - ): - return - if name in tools: - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = cast(RuntimeTool, item) - - for toolset in toolsets: - visit(toolset, []) - return tools - - def _toolset_entries( - self, toolset: Toolset, state: State | None - ) -> list[ToolEntry]: - return [*toolset.tools, *self.scoped_tool_entries(toolset, state)] - - def _tool_visible_for_task( - self, - toolset: Toolset, - name: str, - tool_filters: dict[str, VisibilityConfig] | None, - ) -> bool: - if not tool_visible(toolset, name): - return False - toolset_name = self._toolset_name(toolset) - if toolset_name is None or tool_filters is None: - return True - selected = tool_filters.get(toolset_name) - if selected is None: - return True - show = selected.show - hide = selected.hide - if show is not None and name not in show: - return False - if hide is not None and name in hide: - return False - return True - - def _toolset_name(self, toolset: Toolset) -> str | None: - for name, named_toolset in self.named_toolsets.items(): - if named_toolset is toolset: - return name - return None - def _tool_owners_for( - self, toolsets: Sequence[Toolset], state: State | None = None - ) -> dict[str, Toolset]: - owners: dict[str, Toolset] = {} - - def visit(toolset: Toolset) -> None: - for item in self._toolset_entries(toolset, state): - if isinstance(item, Toolset): - visit(item) - continue - if isinstance(item, MCPTool): - continue - name = tool_name(item) - if name in owners: - raise ValueError(f"Tool {name!r} is defined twice.") - owners[name] = toolset - - for toolset in toolsets: - visit(toolset) - return owners - - def tool_owner(self, name: str, state: State) -> Toolset | None: - return self._tool_owners_for(self.active_toolsets(state), state).get(name) - - def _owner_signals(self, owner: RuntimeOwnerMixin | None) -> list[SignalRecord]: - if owner is None: - return [] - return build_signals( - owner=owner, - scoring=owner.config.scoring, - metrics=owner.metrics, - rewards=owner.rewards, - advantages=owner.advantages, - ) +class ModalRuntime(Runtime): + def __init__(self, config: ModalRuntimeConfig) -> None: + self.config = config - def _handler_owners(self) -> tuple["Taskset | Harness | None", ...]: - return (self.taskset, self.harness) + async def start(self) -> None: + raise NotImplementedError("Modal runtime is not implemented yet.") - def _extra_handlers( - self, - attr: CallableKind, - builtins: Sequence[Handler] = (), - owners: Sequence["Taskset | Harness | Toolset | None"] | None = None, - ) -> list[Handler]: - handlers: list[Handler] = list(builtins) - for owner in owners or self._handler_owners(): - if owner is None: - continue - for handler in lifecycle_handlers(owner, attr): - if not callable(handler): - raise TypeError(f"{attr} entries must be callable.") - handlers.append(cast(Handler, handler)) - return handlers - - def _rollout_handlers( - self, - attr: CallableKind, - state: State, - stage: str | None = None, - ) -> list[Handler]: - handlers: list[Handler] = [] - for toolset in iter_toolsets(self.active_toolsets(state)): - for handler in lifecycle_handlers(toolset, attr): - if not callable(handler): - raise TypeError(f"{attr} entries must be callable.") - if stage is not None and handler_stage(handler, attr) != stage: - continue - handlers.append(cast(Handler, handler)) - for _, method in inspect.getmembers(toolset, predicate=callable): - if not handler_is_marked(method, attr): - continue - if stage is not None and handler_stage(method, attr) != stage: - continue - handlers.append(cast(Handler, method)) - return sort_handlers(unique_handlers(handlers), attr) - - def _group_handlers( - self, - attr: CallableKind, - states: Sequence[State], - stage: str | None = None, - ) -> list[Handler]: - handlers: list[Handler] = [] - for state in states: - handlers.extend(self._rollout_handlers(attr, state, stage=stage)) - return sort_handlers(unique_handlers(handlers), attr) - - async def collect_artifact( + async def stop(self) -> None: + return None + + async def expose(self, port: int) -> str: + _ = port + raise NotImplementedError("Modal runtime is not implemented yet.") + + async def run( self, - spec: ArtifactConfig, - task: Task, - state: State, + command: list[str], *, - sandbox_lease: "SandboxLease | None", - ) -> object: - path = spec.path.format(**{**dict(task), **state}) - if sandbox_lease is not None: - from .utils.sandbox_utils import read_sandbox_artifact - - try: - content = await read_sandbox_artifact( - sandbox_lease.client, sandbox_lease.id, path - ) - except FileNotFoundError: - if spec.optional: - return None - raise - return spec.parse(content) - - matches = sorted(glob.glob(path)) - if not matches: - if spec.optional: - return None - raise FileNotFoundError(f"Artifact path matched no files: {path!r}") - with open(matches[0], encoding="utf-8") as f: - return spec.parse(f.read()) - - async def collect_runtime_artifact( - self, artifact: RuntimeArtifact, task: Task, state: State - ) -> object: - sandbox_lease = self.active_artifact_sandbox_lease(artifact.owner, state) - owner_requires_sandbox = ( - isinstance(artifact.owner, Toolset | User) - and artifact.owner.sandbox is not None - ) - if sandbox_lease is None and owner_requires_sandbox: - if artifact.config.optional: - return None - raise RuntimeError( - f"Artifact {artifact.name!r} requires an active owner sandbox." - ) - return await self.collect_artifact( - artifact.config, - task, - state, - sandbox_lease=sandbox_lease, - ) + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> CommandResult: + _ = command, cwd, env, timeout + raise NotImplementedError("Modal runtime is not implemented yet.") - def active_artifact_sandbox_lease( - self, owner: ArtifactOwner, state: State - ) -> "SandboxLease | None": - if not isinstance(owner, Toolset | User): - return self.active_program_sandbox_lease(state) - from .utils.sandbox_utils import sandbox_owner_key, tool_sandbox_key - - sandbox = owner.sandbox - if sandbox is None: - return self.active_program_sandbox_lease(state) - if sandbox == "program": - return self.active_program_sandbox_lease(state) - if not isinstance(sandbox, SandboxConfig): - raise TypeError("Owner sandbox must be SandboxConfig or 'program'.") - if sandbox.prefer == "program": - lease = self.active_program_sandbox_lease(state) - if lease is not None: - return lease - scope = sandbox.scope - sandbox_key = ( - tool_sandbox_key(owner) - if isinstance(owner, Toolset) - else sandbox_owner_key(owner) - ) - return self.sandbox_leases.get((self.scope_key(scope, state), sandbox_key)) - - def resolve_path(self, path: str, task: Task, state: State) -> object: - root, separator, tail = path.partition(".") - if root == "task": - value: object = task - elif root == "state": - value = state - elif root == "runtime": - value = state.runtime_state() - elif root == "objects": - raise ValueError( - "objects.* bindings are private to the owning Taskset, Harness, " - "Toolset, or User callable. Use taskset.objects.* or " - "harness.objects.* for explicit cross-owner object bindings." - ) - elif root == "tools": - if not separator: - return self.all_tools(state) - name, _, rest = tail.partition(".") - - value = self._tool_call(name, task, state, exposed=False) - tail = rest - else: - raise ValueError(f"Unknown binding root {root!r}.") - if separator and root not in {"objects", "tools"}: - return read_path(value, tail) - if tail: - return read_path(value, tail) - return value - - async def resolve_runtime_owner_path( - self, owner_name: str, path: str, task: Task, state: State - ) -> object: - root, separator, tail = path.partition(".") - if root != "objects" or not separator: - raise ValueError( - f"{owner_name} binding sources must use {owner_name}.objects.name." - ) - name, _, rest = tail.partition(".") - if not name: - raise ValueError( - f"{owner_name} binding sources must use {owner_name}.objects.name." - ) - if owner_name == "taskset": - value = await self.resolve_taskset_object(name, task, state) - elif owner_name == "harness": - value = await self.resolve_harness_object(name, task, state) - else: - raise ValueError("Runtime owner must be 'taskset' or 'harness'.") - return read_path(value, rest) if rest else value - - async def resolve_toolset_object( - self, toolset: Toolset, name: str, task: Task, state: State - ) -> object: - if not any( - toolset is active_toolset - for active_toolset in iter_toolsets(self.active_toolsets(state)) - ): - raise RuntimeError("Toolset object owner is not active in this runtime.") - if name not in toolset.objects: - raise KeyError(f"Unknown Toolset object {name!r}.") - spec = toolset.objects[name] - scope = toolset_object_scope(toolset) - key = (id(toolset), self.scope_key(scope, state), name) - store = self.runtime_objects["toolset"] - if key in store: - return store[key] - kwargs: RuntimeData = {} - for binding_key, source in toolset.bindings.items(): - target_name, arg_name = binding_key_parts(binding_key) - if target_name == name: - kwargs[arg_name] = await self.resolve_tool_binding( - toolset, source, task, state - ) - obj = await resolve_object_factory(spec, f"Toolset object {name!r}", kwargs) - store[key] = obj - return obj - - def scope_key(self, scope: str, state: State | None = None) -> str: - if scope == "global": - return "global" - if state is None: - raise ValueError(f"{scope} object cleanup requires state.") - if scope == "group": - return str( - state.runtime_state().get("group_key") or state.get("trajectory_id") - ) - if scope == "rollout": - return str(state.get("trajectory_id")) - raise ValueError("Object scope must be 'rollout', 'group', or 'global'.") + async def read(self, path: str) -> bytes: + _ = path + raise NotImplementedError("Modal runtime is not implemented yet.") - async def release_model_client(self, state: State, *, group: bool = False) -> None: - if self.model_handle(state) is not None: - return - if not group and "group_key" in state.runtime_state(): - return - key = state.runtime_state().get("client_key") - if not isinstance(key, str): - return - client = self.model_clients.pop(key, None) - if key not in self.owned_model_clients: - return - self.owned_model_clients.remove(key) - if client is not None: - await close_object(client) - - async def release_all_model_clients(self) -> None: - for key, client in list(self.model_clients.items()): - del self.model_clients[key] - if key in self.owned_model_clients: - self.owned_model_clients.remove(key) - await close_object(client) - - async def resolve_tool_sandbox( - self, toolset: Toolset, task: Task, state: State - ) -> object: - from .utils.sandbox_utils import ( - SandboxHandle, - create_scoped_sandbox_lease, - tool_sandbox_key, - ) + async def write(self, path: str, data: bytes) -> None: + _ = path, data + raise NotImplementedError("Modal runtime is not implemented yet.") - sandbox = toolset.sandbox - if sandbox is None: - raise TypeError("Toolset sandbox must be configured.") - if isinstance(sandbox, str): - if sandbox != "program": - raise ValueError("Toolset sandbox string must be 'program'.") - lease = self.active_program_sandbox_lease(state) - if lease is None: - raise RuntimeError( - "Toolset sandbox='program' requires an active program sandbox." - ) - return SandboxHandle(lease, state) - if not isinstance(sandbox, SandboxConfig): - raise TypeError("Toolset sandbox must be SandboxConfig or 'program'.") - if sandbox.prefer is not None: - lease = self.active_program_sandbox_lease(state) - if lease is not None: - return SandboxHandle(lease, state) - scope = sandbox.scope - key = (self.scope_key(scope, state), tool_sandbox_key(toolset)) - lease = await self.resolve_sandbox_lease( - key, - lambda: create_scoped_sandbox_lease( - toolset, - key[1], - client=self.sandbox_client(), - ), - ) - return SandboxHandle(lease, state) - - def active_program_sandbox_lease(self, state: State) -> "SandboxLease | None": - sandbox_handle = self.sandbox_handle(state) - if sandbox_handle is not None: - return self.sandbox_lease_from_handle(sandbox_handle, "sandbox") - sandbox_state = state.runtime_state().get("sandbox") - if sandbox_state is None: - return None - sandbox_record = SandboxRuntimeStateConfig.model_validate(sandbox_state) - resolved_lease_key = sandbox_record.lease_key - lease = self.sandbox_leases.get(resolved_lease_key) - if lease is None: - raise RuntimeError("Program sandbox lease is no longer active.") - lease.scope_key = resolved_lease_key[0] - return lease - - async def resolve_program_sandbox( - self, sandbox_config: SandboxConfig, task: Task, state: State - ) -> "SandboxLease": - from .utils.sandbox_utils import ( - create_sandbox_lease, - program_sandbox_key, - ) - sandbox_handle = self.sandbox_handle(state) - if sandbox_handle is not None: - return self.sandbox_lease_from_handle(sandbox_handle, "sandbox") - scope = sandbox_config.scope - key = (self.scope_key(scope, state), program_sandbox_key(sandbox_config)) - lease = await self.resolve_sandbox_lease( - key, - lambda: create_sandbox_lease( - sandbox_config, - key[1], - client=self.sandbox_client(), - ), - ) - return lease - - def sandbox_lease_from_handle( - self, handle: SandboxRuntimeHandleConfig, name: str - ) -> "SandboxLease": - runtime = self.resolved_runtime(handle) - resolved_lease_key = handle.lease_key - lease = runtime.sandbox_leases.get(resolved_lease_key) - if lease is None: - raise RuntimeError(f"Resolved {name} sandbox lease is no longer active.") - lease.scope_key = resolved_lease_key[0] - return lease - - async def resolve_user_sandbox( - self, user: User, task: Task, state: State - ) -> object: - from .utils.sandbox_utils import ( - SandboxHandle, - create_scoped_sandbox_lease, - sandbox_owner_key, - ) +class DaytonaRuntime(Runtime): + def __init__(self, config: DaytonaRuntimeConfig) -> None: + self.config = config - sandbox = user.sandbox - if sandbox is None: - raise TypeError("User sandbox must be configured.") - scope = sandbox.scope - key = (self.scope_key(scope, state), sandbox_owner_key(user)) - lease = await self.resolve_sandbox_lease( - key, - lambda: create_scoped_sandbox_lease( - user, - key[1], - client=self.sandbox_client(), - ), - ) - return SandboxHandle(lease, state) - - async def release_sandboxes(self, scope: str, state: State) -> None: - scope_key = self.scope_key(scope, state) - async with self.sandbox_lock: - pending_creations = [ - (key, task) - for key, task in self.sandbox_creation_tasks.items() - if key[0] == scope_key - ] - scoped_leases = [ - (key, handle) - for key, handle in self.sandbox_leases.items() - if key[0] == scope_key and handle.scope == scope - ] - if pending_creations: - await self.clear_sandbox_creation_tasks( - pending_creations, state=state, scope=scope - ) - deletion_failures = 0 - for key, handle in scoped_leases: - try: - await self.close_sandbox_lease(handle) - except Exception as exc: - deletion_failures += 1 - logger.warning( - "Failed to delete %s sandbox %s for scope key %s: %s", - scope, - handle.id, - key[0], - exc, - exc_info=True, - ) - cleanup_errors = cast( - list[ConfigData], state.setdefault("cleanup_errors", []) - ) - cleanup_errors.append( - { - "type": type(exc).__name__, - "message": str(exc), - "scope": scope, - } - ) - else: - async with self.sandbox_lock: - if self.sandbox_leases.get(key) is handle: - del self.sandbox_leases[key] - if deletion_failures: - logger.error( - "%s/%s %s sandbox deletions failed during cleanup", - deletion_failures, - len(scoped_leases), - scope, - ) + async def start(self) -> None: + raise NotImplementedError("Daytona runtime is not implemented yet.") - async def ensure_global_sandboxes(self, state: State | None = None) -> None: - from .utils.sandbox_utils import ( - create_scoped_sandbox_lease, - sandbox_owner_key, - tool_sandbox_key, - ) + async def stop(self) -> None: + return None - toolsets = self.active_toolsets(state) if state is not None else self.toolsets - owners: list[Toolset | User] = [*iter_toolsets(toolsets)] - user = self.active_user() - if user is not None: - owners.append(user) - for owner in owners: - sandbox = owner.sandbox - if sandbox is None or sandbox == "program": - continue - if not isinstance(sandbox, SandboxConfig): - raise TypeError("Owner sandbox must be SandboxConfig or 'program'.") - if sandbox.scope != "global": - continue - sandbox_key = ( - tool_sandbox_key(owner) - if isinstance(owner, Toolset) - else sandbox_owner_key(owner) - ) - await self.resolve_sandbox_lease( - ("global", sandbox_key), - lambda owner=owner, sandbox_key=sandbox_key: ( - create_scoped_sandbox_lease( - owner, - sandbox_key, - client=self.sandbox_client(), - ) - ), - ) + async def expose(self, port: int) -> str: + _ = port + raise NotImplementedError("Daytona runtime is not implemented yet.") - def bind_global_sandboxes(self, state: State) -> None: - from .utils.sandbox_utils import attach_sandbox_ref - - for key, lease in self.sandbox_leases.items(): - scope_key, _ = key - if scope_key != "global": - continue - attach_sandbox_ref(state, lease) - - async def ensure_mcp_tools(self, state: State) -> None: - from .utils.mcp_utils import connect_mcp_tool - - for key in self.mcp_scope_keys(state): - if key in self.mcp_exit_stacks: - continue - exit_stack = AsyncExitStack() - tools: RuntimeTools = {} - tool_parents: dict[str, tuple[Toolset, ...]] = {} - try: - for toolset in self.active_toolsets(state): - await self._register_mcp_tools( - toolset, - [toolset], - connect_mcp_tool, - exit_stack, - tools, - tool_parents, - state, - key, - ) - except BaseException: - await exit_stack.aclose() - raise - self.mcp_exit_stacks[key] = exit_stack - self.mcp_tools[key] = tools - self.mcp_tool_parents[key] = tool_parents - - async def _register_mcp_tools( + async def run( self, - toolset: Toolset, - parents: list[Toolset], - connect_mcp_tool: Callable[ - [MCPTool, AsyncExitStack[bool | None]], - Awaitable[Sequence["MCPToolHandle"]], - ], - exit_stack: AsyncExitStack, - tools: RuntimeTools, - tool_parents: dict[str, tuple[Toolset, ...]], - state: State, - target_key: str, - ) -> None: - for item in self._toolset_entries(toolset, state): - if isinstance(item, Toolset): - await self._register_mcp_tools( - item, - [*parents, item], - connect_mcp_tool, - exit_stack, - tools, - tool_parents, - state, - target_key, - ) - continue - if not isinstance(item, MCPTool): - continue - if self.mcp_scope_key(toolset, state) != target_key: - continue - handles = await connect_mcp_tool(item, exit_stack) - for handle in handles: - name = tool_name(handle) - if ( - name - in self.tools_for_toolsets( - self.active_toolsets(state), - apply_visibility=False, - state=state, - ) - or name in tools - ): - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = handle - tool_parents[name] = tuple(parents) - - async def close_mcp_tools(self, state: State, scope: str = "rollout") -> None: - for key in self.mcp_scope_keys(state, scope=scope): - exit_stack = self.mcp_exit_stacks.pop(key, None) - self.mcp_tools.pop(key, None) - self.mcp_tool_parents.pop(key, None) - if exit_stack is not None: - await exit_stack.aclose() - - async def close_all_mcp_tools(self) -> None: - for key, exit_stack in list(self.mcp_exit_stacks.items()): - self.mcp_tools.pop(key, None) - self.mcp_tool_parents.pop(key, None) - del self.mcp_exit_stacks[key] - await exit_stack.aclose() - - def all_tools(self, state: State) -> RuntimeTools: - tools = self.tools_for_toolsets( - self.active_toolsets(state), apply_visibility=False, state=state - ) - for name, tool in self.mcp_tools_for_state(state, exposed=False).items(): - if name in tools: - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = tool - for name, tool in self.borrowed_tools_for_state(state).items(): - if name in tools: - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = tool - return tools - - def all_exposed_tools(self, state: State, *, validate: bool = True) -> RuntimeTools: - active_toolsets = self.active_toolsets(state) - tool_filters = self._task_tools_config( - state, active_toolsets, validate=validate - ) - tools = self.tools_for_toolsets( - active_toolsets, - apply_visibility=True, - state=state, - tool_filters=tool_filters, - ) - for name, tool in self.mcp_tools_for_state(state, exposed=True).items(): - if name in tools: - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = tool - for name, tool in self.borrowed_tools_for_state(state).items(): - if name in tools: - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = tool - return tools - - def borrowed_tools_for_state(self, state: State) -> RuntimeTools: - handle = self.resolved_handles(state).tools - if handle is None: - return {} - source_runtime = self.resolved_runtime(handle) - return { - name: cast( - RuntimeTool, BorrowedTool(source_runtime, handle.handle_id, name) - ) - for name in handle.names - } - - def _task_tools_config( - self, - state: State, - active_toolsets: Sequence[Toolset], + command: list[str], *, - validate: bool, - ) -> dict[str, VisibilityConfig]: - task = self.task_for_state(state) - task_tools = task.tools_config() - active_ids = {id(toolset) for toolset in iter_toolsets(active_toolsets)} - active_named_toolsets = { - name: toolset - for name, toolset in self.named_toolsets.items() - if id(toolset) in active_ids - } - filters: dict[str, VisibilityConfig] = {} - for name, filter_config in task_tools.items(): - if name not in active_named_toolsets: - raise KeyError(f"Unknown toolset tools filter: {name!r}.") - show = filter_config.show - hide = filter_config.hide - if validate: - available = set( - self._tool_names_for_toolset(active_named_toolsets[name], state) - ) - selected = set(show or hide or []) - unknown = sorted(selected - available) - if unknown: - raise KeyError(f"Unknown tools for toolset {name!r}: {unknown}.") - filters[name] = filter_config - return filters - - def _tool_names_for_toolset(self, toolset: Toolset, state: State) -> list[str]: - names: list[str] = [] - - def visit_toolset(item: Toolset) -> None: - names.extend(self.mcp_tools.get(self.mcp_scope_key(item, state), ())) - for child in self._toolset_entries(item, state): - visit_entry(child) - - def visit_entry(item: ToolEntry) -> None: - if isinstance(item, Toolset): - visit_toolset(item) - return - if isinstance(item, MCPTool): - return - names.append(tool_name(item)) - - visit_toolset(toolset) - return names - - def mcp_tools_for_state(self, state: State, exposed: bool) -> RuntimeTools: - tool_filters = ( - self._task_tools_config(state, self.active_toolsets(state), validate=False) - if exposed - else {} - ) - tools: RuntimeTools = {} - for key in self.mcp_scope_keys(state): - for name, tool in self.mcp_tools.get(key, {}).items(): - parents = self.mcp_tool_parents.get(key, {}).get(name, ()) - if exposed and not all( - self._tool_visible_for_task(parent, name, tool_filters) - for parent in parents - ): - continue - if name in tools: - raise ValueError(f"Tool {name!r} is defined twice.") - tools[name] = tool - return tools - - def mcp_scope_keys(self, state: State, scope: str | None = None) -> list[str]: - keys: list[str] = [] - - def visit(toolset: Toolset) -> None: - for item in self._toolset_entries(toolset, state): - if isinstance(item, Toolset): - visit(item) - continue - if not isinstance(item, MCPTool): - continue - item_scope = toolset_object_scope(toolset) - if scope is not None and item_scope != scope: - continue - key = self.mcp_scope_key(toolset, state) - if key not in keys: - keys.append(key) - - for toolset in self.active_toolsets(state): - visit(toolset) - return keys - - def mcp_scope_key(self, toolset: Toolset, state: State) -> str: - scope = toolset_object_scope(toolset) - return f"{scope}:{self.scope_key(scope, state)}:{id(toolset)}" + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float | None = None, + ) -> CommandResult: + _ = command, cwd, env, timeout + raise NotImplementedError("Daytona runtime is not implemented yet.") + + async def read(self, path: str) -> bytes: + _ = path + raise NotImplementedError("Daytona runtime is not implemented yet.") + + async def write(self, path: str, data: bytes) -> None: + _ = path, data + raise NotImplementedError("Daytona runtime is not implemented yet.") + + +async def docker(*args: str, timeout: float | None = None) -> CommandResult: + process = await asyncio.create_subprocess_exec( + "docker", + *args, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + try: + stdout, stderr = await asyncio.wait_for(process.communicate(), timeout) + except asyncio.TimeoutError as exc: + with contextlib.suppress(ProcessLookupError): + process.kill() + await process.wait() + raise TimeoutError( + f"Docker command timed out after {timeout} seconds." + ) from exc + return CommandResult( + returncode=int(process.returncode or 0), + stdout=stdout.decode(errors="replace"), + stderr=stderr.decode(errors="replace"), + ) diff --git a/verifiers/v1/runtime_handles.py b/verifiers/v1/runtime_handles.py deleted file mode 100644 index 954972d1b7..0000000000 --- a/verifiers/v1/runtime_handles.py +++ /dev/null @@ -1,51 +0,0 @@ -from typing import Literal - -from pydantic import Field -from verifiers.types import ClientType - -from .config import Config -from .types import ConfigData - - -class RuntimeHandleConfig(Config): - runtime_id: str = Field(min_length=1) - - -class ModelRuntimeHandleConfig(RuntimeHandleConfig): - client_key: str = Field(min_length=1) - model: str | None = None - client_type: ClientType | None = None - sampling_args: ConfigData | None = None - - -class TrajectoryRuntimeHandleConfig(RuntimeHandleConfig): - trajectory_id: str = Field(min_length=1) - mode: Literal["append"] = "append" - start: int = 0 - - -class SandboxRuntimeStateConfig(Config): - id: str = Field(min_length=1) - scope: str = Field(min_length=1) - key: str = Field(min_length=1) - lease_key: tuple[str, str] - - -class SandboxRuntimeHandleConfig(RuntimeHandleConfig): - id: str = Field(min_length=1) - scope: str = Field(min_length=1) - key: str = Field(min_length=1) - lease_key: tuple[str, str] - - -class ToolsRuntimeHandleConfig(RuntimeHandleConfig): - handle_id: str = Field(min_length=1) - names: list[str] = Field(default_factory=list) - - -class ResolvedRuntimeHandlesConfig(Config): - model: ModelRuntimeHandleConfig | None = None - endpoint: RuntimeHandleConfig | None = None - trajectory: TrajectoryRuntimeHandleConfig | None = None - sandbox: SandboxRuntimeHandleConfig | None = None - tools: ToolsRuntimeHandleConfig | None = None diff --git a/verifiers/v1/sandbox.py b/verifiers/v1/sandbox.py deleted file mode 100644 index cefe606034..0000000000 --- a/verifiers/v1/sandbox.py +++ /dev/null @@ -1,66 +0,0 @@ -from typing import Literal - -from pydantic import Field, field_validator - -from .config import Config -from .types import ConfigData -from .utils.config_utils import ( - explicit_config_data, - resolved_config_data, -) - - -class SandboxConfig(Config): - image: str = "python:3.11-slim" - start_command: str = "tail -f /dev/null" - cpu_cores: float = 1.0 - memory_gb: float = 2.0 - disk_size_gb: float = 5.0 - gpu_count: int = 0 - gpu_type: str | None = None - vm: bool | None = None - network_access: bool = True - timeout_minutes: int = 60 - create_timeout: int | None = None - wait_timeout: int | None = None - environment_vars: dict[str, str] = {} - secrets: dict[str, str] = {} - team_id: str | None = None - region: str | None = None - registry_credentials_id: str | None = None - guaranteed: bool = False - workdir: str | None = None - command_timeout: int | None = None - poll_interval: int = 3 - packages: list[str] = [] - install_timeout: int = 300 - setup_commands: list[str] = [] - setup_timeout: int = 300 - labels: list[str] = [] - scope: Literal["rollout", "group", "global"] = "rollout" - prefer: Literal["program"] | None = None - create_concurrency: int = Field(default=128, ge=1) - create_rate_per_second: float | None = Field(default=None, gt=0) - delete_concurrency: int = Field(default=128, ge=1) - delete_rate_per_second: float | None = Field(default=None, gt=0) - - @field_validator("packages", "setup_commands", "labels", mode="before") - @classmethod - def validate_string_list(cls, value: object) -> object: - if isinstance(value, str): - return [value] - return value - - @field_validator("environment_vars", "secrets", mode="before") - @classmethod - def validate_string_mapping(cls, value: object) -> object: - if value is None: - return {} - if isinstance(value, dict): - return {str(key): str(item) for key, item in value.items()} - return value - - def data(self, *, fill_defaults: bool = True) -> ConfigData: - if fill_defaults: - return resolved_config_data(self) - return explicit_config_data(self, SandboxConfig) diff --git a/verifiers/v1/state.py b/verifiers/v1/state.py index d8502caae9..851ab2e320 100644 --- a/verifiers/v1/state.py +++ b/verifiers/v1/state.py @@ -1,10 +1,390 @@ -"""V1 state contract exports. +from __future__ import annotations -V1 uses the shared top-level ``verifiers.State`` type. V1 tasks opt state -instances into strict runtime/lifecycle-field handling when they are passed to -``State.for_task(...)``. -""" +import json +import time +import uuid +from copy import deepcopy +from typing import TYPE_CHECKING, cast -from verifiers.types import State +from pydantic import ( + BaseModel, + computed_field, + Field, + StrictBool, + StrictFloat, + StrictInt, + model_validator, +) -__all__ = ["State"] +from verifiers.types import ( + ErrorData, + FinishReason, + Messages, + ResponseTokens, + ToolCall, + ToolMessage, + Usage, +) +from verifiers.utils.error_utils import error_data, validate_error_data +from verifiers.utils.save_utils import serialize_messages_for_output + +from .types import JsonData, ModelConfig +from .utils.json_utils import json_data +from .utils.task_freeze_utils import assert_serializable + +if TYPE_CHECKING: + from .task import Task + + +class TimeSpan(BaseModel, extra="forbid"): + start: float = 0.0 + end: float = 0.0 + + @property + def duration(self) -> float: + return max(0.0, self.end - self.start) if self.end else 0.0 + + def begin(self) -> None: + self.start = time.time() + + def finish(self) -> None: + self.end = time.time() + + +class Timing(BaseModel, extra="forbid"): + start_time: float = Field(default_factory=time.time) + setup: TimeSpan = Field(default_factory=TimeSpan) + generation: TimeSpan = Field(default_factory=TimeSpan) + scoring: TimeSpan = Field(default_factory=TimeSpan) + cleanup: TimeSpan = Field(default_factory=TimeSpan) + model: list[TimeSpan] = Field(default_factory=list) + runtime: list[TimeSpan] = Field(default_factory=list) + + @property + def total(self) -> float: + end = max( + self.setup.end, + self.generation.end, + self.scoring.end, + self.cleanup.end, + self.start_time, + ) + return max(0.0, end - self.start_time) + + +class TurnUsage(BaseModel, extra="forbid"): + prompt_tokens: StrictInt = 0 + reasoning_tokens: StrictInt = 0 + completion_tokens: StrictInt = 0 + total_tokens: StrictInt = 0 + + @classmethod + def from_usage(cls, usage: Usage | None) -> "TurnUsage | None": + if usage is None: + return None + return cls( + prompt_tokens=usage.prompt_tokens, + reasoning_tokens=usage.reasoning_tokens, + completion_tokens=usage.completion_tokens, + total_tokens=usage.total_tokens, + ) + + +class TurnTokens(BaseModel, extra="forbid"): + prompt_ids: list[StrictInt] = Field(default_factory=list) + prompt_mask: list[StrictInt] = Field(default_factory=list) + prompt_advantages: list[StrictFloat] | None = None + completion_ids: list[StrictInt] = Field(default_factory=list) + completion_mask: list[StrictInt] = Field(default_factory=list) + completion_logprobs: list[StrictFloat] = Field(default_factory=list) + completion_advantages: list[StrictFloat] | None = None + overlong_prompt: StrictBool = False + is_truncated: StrictBool = False + + @classmethod + def from_response( + cls, tokens: ResponseTokens | None, *, is_truncated: bool = False + ) -> "TurnTokens | None": + if tokens is None: + return None + return cls( + prompt_ids=list(tokens.prompt_ids), + prompt_mask=list(tokens.prompt_mask), + completion_ids=list(tokens.completion_ids), + completion_mask=list(tokens.completion_mask), + completion_logprobs=list(tokens.completion_logprobs), + is_truncated=is_truncated, + ) + + @model_validator(mode="after") + def validate_lengths(self) -> "TurnTokens": + if len(self.prompt_ids) != len(self.prompt_mask): + raise ValueError("TurnTokens prompt_ids and prompt_mask lengths differ.") + if self.prompt_advantages is not None and len(self.prompt_ids) != len( + self.prompt_advantages + ): + raise ValueError( + "TurnTokens prompt_ids and prompt_advantages lengths differ." + ) + if len(self.completion_ids) != len(self.completion_mask): + raise ValueError( + "TurnTokens completion_ids and completion_mask lengths differ." + ) + if len(self.completion_ids) != len(self.completion_logprobs): + raise ValueError( + "TurnTokens completion_ids and completion_logprobs lengths differ." + ) + if self.completion_advantages is not None and len(self.completion_ids) != len( + self.completion_advantages + ): + raise ValueError( + "TurnTokens completion_ids and completion_advantages lengths differ." + ) + return self + + +class Turn(BaseModel, extra="forbid"): + """One model request/response boundary in a rollout transcript.""" + + id: str = Field(default_factory=lambda: uuid.uuid4().hex) + prompt: Messages + completion: Messages = Field(default_factory=list) + tool_calls: list[ToolCall] = Field(default_factory=list) + tool_results: list[ToolMessage] = Field(default_factory=list) + response_id: str | None = None + model: str | None = None + created: StrictInt | None = None + finish_reason: FinishReason = None + usage: TurnUsage | None = None + tokens: TurnTokens | None = None + reward: float | None = None + is_truncated: bool = False + timing: TimeSpan = Field(default_factory=TimeSpan) + + +class Extras(BaseModel, extra="forbid"): + @staticmethod + def schema_for(extras: "Extras | None") -> type["Extras"] | None: + if extras is None: + return None + schema = type(extras) + if not issubclass(schema, Extras): + raise TypeError("extras config must be a vf.Extras object.") + return schema + + @staticmethod + def defaults_for(extras: "Extras | None") -> JsonData: + if extras is None: + return {} + return json_data(extras.model_dump(mode="json", exclude_none=True)) + + @staticmethod + def merge_defaults( + taskset_defaults: JsonData, harness_defaults: JsonData + ) -> JsonData: + conflicts = sorted(set(taskset_defaults) & set(harness_defaults)) + if conflicts: + raise ValueError( + f"Extras config keys are defined twice: {', '.join(conflicts)}." + ) + return {**deepcopy(taskset_defaults), **deepcopy(harness_defaults)} + + @staticmethod + def realize_schema( + taskset_schema: type["Extras"] | None, + harness_schema: type["Extras"] | None, + ) -> type["Extras"] | None: + schemas = [schema for schema in (taskset_schema, harness_schema) if schema] + if not schemas: + return None + if len(schemas) == 1: + return schemas[0] + seen: dict[str, type[BaseModel]] = {} + for schema in schemas: + for field_name in schema.model_fields: + if field_name in seen: + raise ValueError( + f"Extras field {field_name!r} is defined by both " + f"{seen[field_name].__name__} and {schema.__name__}." + ) + seen[field_name] = schema + return cast( + type[Extras], + type( + "RealizedExtras", + tuple(schemas), + {"__module__": __name__}, + ), + ) + + +class State(BaseModel, extra="forbid"): + """Strict serializable v1 rollout state.""" + + id: str = Field(default_factory=lambda: uuid.uuid4().hex) + task_id: str | None = None + transcript: list[Turn] = Field(default_factory=list) + extras: JsonData = Field(default_factory=dict) + metrics: dict[str, float] = Field(default_factory=dict) + reward: float = 0.0 + artifacts: JsonData = Field(default_factory=dict) + usage: dict[str, float] = Field(default_factory=dict) + timing: Timing = Field(default_factory=Timing) + is_completed: bool = False + is_truncated: bool = False + stop_condition: str | None = None + error: ErrorData | None = None + group_id: str | None = None + metadata: JsonData = Field(default_factory=dict) + model: ModelConfig | None = None + teacher: ModelConfig | None = None + + @computed_field + @property + def prompt(self) -> Messages: + if not self.transcript: + return [] + return self.transcript[-1].prompt + + @computed_field + @property + def completion(self) -> Messages: + if not self.transcript: + return [] + return self.transcript[-1].completion + + @computed_field + @property + def messages(self) -> Messages: + if not self.transcript: + return [] + latest = self.transcript[-1] + return [*latest.prompt, *latest.completion] + + def stop(self, condition: str = "state_done") -> None: + if not condition: + raise ValueError("State.stop condition must be non-empty.") + self.is_completed = True + self.stop_condition = condition + + def capture_error(self, error: BaseException) -> None: + self.error = error_data(error) + self.stop("has_error") + + def assert_serializable(self) -> None: + assert_serializable(self.model_dump(mode="json", exclude_none=True)) + + @staticmethod + def serialized_messages(messages: object) -> list[JsonData]: + serialized: list[JsonData] = [] + for index, message in enumerate(serialize_messages_for_output(messages)): + serialized.append(json_data(message, context=f"message[{index}]")) + return serialized + + @staticmethod + def turn_record(turn: Turn) -> dict[str, object]: + return { + "id": turn.id, + "prompt": State.serialized_messages(turn.prompt), + "completion": State.serialized_messages(turn.completion), + "tool_calls": [ + json_data(tool_call, context="tool_call") + for tool_call in turn.tool_calls + ], + "tool_results": State.serialized_messages(turn.tool_results), + "response_id": turn.response_id, + "model": turn.model, + "created": turn.created, + "finish_reason": turn.finish_reason, + "usage": turn.usage.model_dump(mode="json", exclude_none=True) + if turn.usage is not None + else None, + "tokens": turn.tokens.model_dump(mode="json", exclude_none=True) + if turn.tokens is not None + else None, + "reward": turn.reward, + "is_truncated": turn.is_truncated, + "timing": turn.timing.model_dump(mode="json", exclude_none=True), + } + + def to_output( + self, task: "Task", state_columns: list[str] | None = None + ) -> dict[str, object]: + prompt = self.prompt if self.transcript else task.prompt + serialize_messages = type(self).serialized_messages + turn_record = type(self).turn_record + output: dict[str, object] = { + "example_id": task.row_id, + "prompt": serialize_messages(prompt), + "completion": serialize_messages(self.completion), + "reward": float(self.reward), + "timing": self.timing.model_dump(mode="json"), + "is_completed": self.is_completed, + "is_truncated": self.is_truncated, + "metrics": dict(self.metrics), + "extras": deepcopy(self.extras), + "stop_condition": self.stop_condition, + "transcript": [turn_record(turn) for turn in self.transcript], + } + answer = getattr(task, "answer", None) + if answer is not None: + output["answer"] = str(answer) + info = getattr(task, "info", None) + if isinstance(info, dict) and info: + output["info"] = deepcopy(info) + if self.error is not None: + output["error"] = validate_error_data(self.error) + output["error_chain"] = self.error["error_chain_repr"] + output["long_error_chain"] = self.error["error_chain_str"] + usage = dict(self.usage) + for turn in self.transcript: + if turn.usage is None: + continue + usage["input_tokens"] = usage.get("input_tokens", 0.0) + float( + turn.usage.prompt_tokens + ) + usage["output_tokens"] = usage.get("output_tokens", 0.0) + float( + turn.usage.completion_tokens + ) + if usage: + output["token_usage"] = { + "input_tokens": float(usage.get("input_tokens", 0.0)), + "output_tokens": float(usage.get("output_tokens", 0.0)), + } + reserved_output_fields = set(output) + for key, value in self.metrics.items(): + if key in output: + raise ValueError( + f"Metric name {key!r} conflicts with a reserved output field." + ) + output[key] = value + for column in state_columns or []: + if column in output: + if column in reserved_output_fields: + continue + raise ValueError( + f"State column {column!r} conflicts with an existing output field." + ) + if column == "prompt": + output[column] = serialize_messages(prompt) + elif column == "completion": + output[column] = serialize_messages(self.completion) + elif column == "messages": + output[column] = serialize_messages(self.messages) + elif column == "extras": + output[column] = deepcopy(self.extras) + elif column == "transcript": + output[column] = [turn_record(turn) for turn in self.transcript] + elif column in type(self).model_fields: + value = getattr(self, column) + model_dump = getattr(value, "model_dump", None) + if callable(model_dump): + output[column] = model_dump(mode="json", exclude_none=True) + else: + output[column] = deepcopy(value) + elif column in self.extras: + output[column] = deepcopy(self.extras[column]) + else: + output[column] = None + json.dumps(output) + return output diff --git a/verifiers/v1/task.py b/verifiers/v1/task.py index 68c503a871..eee83d6843 100644 --- a/verifiers/v1/task.py +++ b/verifiers/v1/task.py @@ -1,150 +1,179 @@ +from __future__ import annotations + +import hashlib +import json +from collections.abc import Mapping from copy import deepcopy -from .artifact import ArtifactsConfig -from .model import model_config_data -from .sandbox import SandboxConfig -from .toolset import VisibilityConfig -from .utils.task_freeze_utils import assert_serializable, freeze_value -from .utils.prompt_utils import normalize_prompt, normalize_system_prompt -from .types import JsonData, JsonValue +from pydantic import ( + BaseModel, + Field, + SerializationInfo, + SerializerFunctionWrapHandler, + TypeAdapter, + model_serializer, + model_validator, +) +from typing_extensions import Self + +from verifiers.types import Messages, UserMessage +from .types import JsonData +from .utils.prompt_utils import SystemPrompt, dump_messages, normalize_system_prompt +from .utils.task_freeze_utils import assert_serializable -class Task(dict): - _vf_state_contract = "v1" +MESSAGES_ADAPTER = TypeAdapter(Messages) - def __init__(self, task: JsonData | None = None): - super().__init__(deepcopy(dict(task or {}))) - self._frozen = False - def freeze(self) -> "Task": - if "runtime" in self: +class TaskVisibility(BaseModel, extra="forbid", frozen=True): + show: list[str] | None = None + hide: list[str] | None = None + + @model_validator(mode="after") + def validate_visibility(self) -> Self: + if self.show is not None and self.hide is not None: + raise ValueError("Task visibility accepts show or hide, not both.") + return self + + +class Resources(BaseModel, extra="forbid", frozen=True): + cpu_cores: float | None = None + memory_gb: float | None = None + gpu_count: int | None = None + disk_gb: float | None = None + + +class Task(BaseModel, extra="forbid", frozen=True): + """Immutable serializable task specification. Subclass for task-specific data.""" + + task_id: str = "" + row_id: int = 0 + prompt: Messages = Field(default_factory=list) + system_prompt: SystemPrompt = None + toolsets: TaskVisibility | None = None + tools: TaskVisibility | None = None + user: bool | None = None + name: str | None = None + description: str | None = None + image: str | None = None + resources: Resources = Field(default_factory=Resources) + max_turns: int | None = None + + def __init__( + self, task: Mapping[str, object] | None = None, **data: object + ) -> None: + if task is not None and data: raise TypeError( - "task.runtime is not supported; use top-level task fields or state.runtime." - ) - if "prompt" in self: - super().__setitem__( - "prompt", normalize_prompt(self["prompt"], field_name="task.prompt") + "Task accepts either a mapping or keyword fields, not both." ) - if "system_prompt" in self: - super().__setitem__( - "system_prompt", - normalize_system_prompt( - self["system_prompt"], field_name="task.system_prompt" - ), - ) - if "tools" in self: - super().__setitem__( - "tools", - { - name: config.model_dump(mode="json", exclude_none=True) - for name, config in self.tools_config().items() - }, - ) - if "toolsets" in self: - super().__setitem__( - "toolsets", - self.toolsets_config().model_dump(mode="json", exclude_none=True), - ) - sandbox_config = self.sandbox_config() - if sandbox_config is not None: - super().__setitem__( - "sandbox", - sandbox_config.data(fill_defaults=False), - ) - if "program" in self and not isinstance(self["program"], dict): - raise TypeError("task.program must be a mapping.") - if "artifacts" in self: - super().__setitem__( - "artifacts", - { - name: artifact.data() - for name, artifact in self.artifacts_config() - .artifacts("task.artifacts") - .items() - }, - ) - if "model" in self: - super().__setitem__("model", model_config_data(self["model"])) - if "max_turns" in self and ( - not isinstance(self["max_turns"], int) - or isinstance(self["max_turns"], bool) + super().__init__(**deepcopy(dict(task or data))) + + @model_validator(mode="before") + @classmethod + def normalize_input(cls, value: object) -> object: + if isinstance(value, Task): + return deepcopy(value.model_dump(mode="python")) + if isinstance(value, Mapping): + raw = deepcopy(dict(value)) + if "id" in raw and "task_id" not in raw: + raw["task_id"] = raw.pop("id") + if "example_id" in raw and "row_id" not in raw: + raw["row_id"] = raw.pop("example_id") + raw.pop("example_id", None) + if "prompt" in raw: + raw["prompt"] = cls.normalize_prompt(raw["prompt"]) + if "system_prompt" in raw: + raw["system_prompt"] = cls.normalize_system_prompt(raw["system_prompt"]) + return raw + return value + + @model_validator(mode="after") + def validate_task(self) -> Self: + if isinstance(self.row_id, bool): + raise TypeError("task.row_id must be an integer.") + if self.max_turns is not None and ( + isinstance(self.max_turns, bool) or not isinstance(self.max_turns, int) ): raise TypeError("task.max_turns must be an integer.") - for key, value in list(self.items()): - super().__setitem__(key, freeze_value(value)) - assert_serializable(self) - self._frozen = True + object.__setattr__(self, "prompt", type(self).normalize_prompt(self.prompt)) + object.__setattr__( + self, + "system_prompt", + type(self).normalize_system_prompt(self.system_prompt), + ) + if not self.task_id: + object.__setattr__(self, "task_id", self.default_task_id()) + assert_serializable(self.model_dump(mode="json", exclude_none=True)) return self - @property - def frozen(self) -> bool: - return self._frozen - - def toolsets_config(self) -> VisibilityConfig: - raw_toolsets = self.get("toolsets") or {} - if not isinstance(raw_toolsets, dict): - raise TypeError("task.toolsets must be a mapping.") - return VisibilityConfig.model_validate(raw_toolsets) - - def tools_config(self) -> dict[str, VisibilityConfig]: - raw_tools = self.get("tools") or {} - if not isinstance(raw_tools, dict): - raise TypeError("task.tools must be a toolset-keyed mapping.") - if "show" in raw_tools or "hide" in raw_tools: - raise ValueError("task.tools must be keyed by toolset name.") - configs: dict[str, VisibilityConfig] = {} - for name, raw_filter in raw_tools.items(): - if not isinstance(name, str): - raise TypeError("task.tools keys must be toolset names.") - configs[name] = VisibilityConfig.model_validate(raw_filter) - return configs - - def sandbox_config(self) -> SandboxConfig | None: - raw_sandbox = self.get("sandbox") - if raw_sandbox is None: + @model_serializer(mode="wrap") + def serialize_task( + self, + handler: SerializerFunctionWrapHandler, + info: SerializationInfo, + ) -> dict[str, object]: + data = handler(self) + if not isinstance(data, dict): + raise TypeError("Task serializer expected a JSON object.") + serialized = {str(key): value for key, value in data.items()} + if info.mode != "json": + return serialized + if "prompt" in serialized: + serialized["prompt"] = dump_messages(self.prompt) + if "system_prompt" in serialized: + if self.system_prompt: + serialized["system_prompt"] = list( + type(self).normalize_system_prompt(self.system_prompt) + ) + else: + serialized.pop("system_prompt", None) + return serialized + + @classmethod + def normalize_prompt(cls, value: object) -> Messages: + messages: Messages + if isinstance(value, str): + messages = [UserMessage(content=value)] + else: + messages = MESSAGES_ADAPTER.validate_python(value or []) + for message in messages: + if getattr(message, "role", None) == "system": + raise ValueError("task.prompt must not contain system messages.") + return messages + + @classmethod + def normalize_system_prompt(cls, value: object) -> list[JsonData]: + return normalize_system_prompt( + cls.system_prompt_input(value), + field_name="task.system_prompt", + ) + + @classmethod + def system_prompt_input(cls, value: object) -> SystemPrompt: + if value is None: return None - if not isinstance(raw_sandbox, dict): - raise TypeError("task.sandbox must be a mapping.") - return SandboxConfig.model_validate(raw_sandbox) - - def artifacts_config(self) -> ArtifactsConfig: - raw_artifacts = self.get("artifacts") or {} - if not isinstance(raw_artifacts, dict): - raise TypeError("task.artifacts must be a mapping.") - return ArtifactsConfig.model_validate(raw_artifacts) - - def __setitem__(self, key: str, value: object) -> None: - self._raise_if_frozen() - super().__setitem__(key, value) - - def __delitem__(self, key: str) -> None: - self._raise_if_frozen() - super().__delitem__(key) - - def update(self, *args: object, **kwargs: object) -> None: - self._raise_if_frozen() - super().update(*args, **kwargs) - - def setdefault(self, key: str, default: object = None) -> object: - self._raise_if_frozen() - return super().setdefault(key, default) - - def pop(self, key: str, default: object = None) -> object: - self._raise_if_frozen() - return super().pop(key, default) - - def popitem(self) -> tuple[str, JsonValue]: - raise TypeError("Task.popitem() is not supported.") - - def clear(self) -> None: - self._raise_if_frozen() - super().clear() - - def __ior__(self, value: object, /) -> "Task": - self._raise_if_frozen() - self.update(value) - return self - - def _raise_if_frozen(self) -> None: - if self._frozen: - raise TypeError("Task is immutable after freeze.") + if isinstance(value, str): + return value + if isinstance(value, list): + return MESSAGES_ADAPTER.validate_python(value) + from .utils.prompt_utils import SystemPromptConfig + + if isinstance(value, Mapping): + return SystemPromptConfig.model_validate(dict(value)) + if isinstance(value, SystemPromptConfig): + return value + raise TypeError("task.system_prompt must be a string, messages list, or null.") + + def default_task_id(self) -> str: + data = self.model_dump( + mode="json", + exclude={"task_id"}, + exclude_none=True, + exclude_defaults=True, + ) + payload = { + "type": f"{type(self).__module__}.{type(self).__qualname__}", + "task": data, + } + raw = json.dumps(payload, sort_keys=True, separators=(",", ":")) + return hashlib.sha256(raw.encode()).hexdigest()[:24] diff --git a/verifiers/v1/taskset.py b/verifiers/v1/taskset.py index 1b22d9d56b..49e064f035 100644 --- a/verifiers/v1/taskset.py +++ b/verifiers/v1/taskset.py @@ -1,24 +1,21 @@ +from __future__ import annotations + +from collections.abc import Mapping from importlib.abc import Traversable from pathlib import Path from typing import Generic, TypeVar, cast, final from datasets import Dataset -from pydantic import AliasChoices, Field +from pydantic import Field, field_serializer, model_validator -from .config import ( - ConfigSource, - LifecycleConfig, -) -from .artifact import ArtifactsConfig -from .state import State +from .config import Config, ConfigSource +from .decorators import discover_decorated +from .state import Extras, State from .task import Task +from .toolset import ServerConfig, ToolsetConfig, ToolsetConfigs +from .runtime import RuntimeConfig +from .types import Handler, JsonData, TaskSplit, Tasks from .user import UserConfig -from .utils.binding_utils import ( - BindingSources, - BindingsConfig, - ObjectsConfig, -) -from .utils.prompt_utils import SystemPrompt, normalize_system_prompt from .utils.config_utils import ( coerce_config, config_ref_context, @@ -26,47 +23,134 @@ registered_config_type, register_config_type, ) -from .utils.runtime_owner_utils import RuntimeOwnerMixin +from .utils.prompt_utils import SystemPrompt, normalize_system_prompt +from .utils.scoring_utils import build_signals from .utils.taskset_utils import ( - dataset_from_result, + dataset_from_result_typed, discover_sibling_dir, prepare_task, task_from_dataset_record, ) -from .types import ( - JsonData, - Objects, - TaskSplit, - Tasks, -) +LifecycleKind = str -class TasksetConfig(LifecycleConfig): - # Core fields configure taskset-owned loaders and runtime behavior. - taskset_id: str | None = Field( - default=None, - validation_alias=AliasChoices("taskset_id", "id"), - ) + +class TasksetConfig(Config): + id: str | None = None system_prompt: SystemPrompt = None user: UserConfig | None = None - bindings: BindingsConfig = BindingsConfig() - objects: ObjectsConfig = ObjectsConfig() - artifacts: ArtifactsConfig = ArtifactsConfig() + toolsets: ToolsetConfigs = Field(default_factory=dict) + runtime: RuntimeConfig | None = None + extras: Extras | None = None + + @model_validator(mode="before") + @classmethod + def resolve_server_sources(cls, value: object) -> object: + if not isinstance(value, Mapping): + return value + data = dict(value) + if "toolsets" in data: + data["toolsets"] = cls.resolve_toolsets_config( + data["toolsets"], + cls.default_toolsets_config(), + ) + if "user" in data: + data["user"] = cls.resolve_user_config( + data["user"], + cls.default_user_config(), + ) + return data + + @field_serializer("user") + def serialize_user(self, value: UserConfig | None) -> dict[str, object] | None: + if value is None: + return None + return value.model_dump(mode="json", exclude_none=True) + + @field_serializer("toolsets") + def serialize_toolsets(self, value: ToolsetConfigs) -> dict[str, dict[str, object]]: + return { + name: config.model_dump(mode="json", exclude_none=True) + for name, config in value.items() + } + + @classmethod + def default_toolsets_config(cls) -> ToolsetConfigs: + value = cls.model_fields["toolsets"].get_default(call_default_factory=True) + if value is None: + return {} + if not isinstance(value, Mapping): + raise TypeError(f"{cls.__name__}.toolsets must be a mapping.") + toolsets: ToolsetConfigs = {} + for name, item in value.items(): + if not isinstance(name, str) or not name: + raise TypeError(f"{cls.__name__}.toolsets keys must be strings.") + toolsets[name] = ServerConfig.resolve_config( + name, + item, + default=None, + base_type=ToolsetConfig, + ) + return toolsets + + @staticmethod + def resolve_toolsets_config( + value: object, defaults: ToolsetConfigs + ) -> ToolsetConfigs: + if not isinstance(value, Mapping): + raise TypeError("TasksetConfig.toolsets must be a mapping.") + toolsets: ToolsetConfigs = dict(defaults) + for name, item in value.items(): + if not isinstance(name, str) or not name: + raise TypeError( + "TasksetConfig.toolsets keys must be non-empty strings." + ) + toolsets[name] = ServerConfig.resolve_config( + name, + item, + default=defaults.get(name), + base_type=ToolsetConfig, + ) + return toolsets + + @staticmethod + def enabled_toolsets(toolsets: ToolsetConfigs) -> ToolsetConfigs: + return { + name: toolset for name, toolset in toolsets.items() if bool(toolset.enabled) + } @classmethod - def __pydantic_init_subclass__(cls, **kwargs: object) -> None: - super().__pydantic_init_subclass__(**kwargs) - field = cls.model_fields.get("taskset_id") - if field is not None: - field.validation_alias = AliasChoices("taskset_id", "id") - cls.model_rebuild(force=True) + def default_user_config(cls) -> UserConfig | None: + value = cls.model_fields["user"].get_default(call_default_factory=True) + if value is None: + return None + return ServerConfig.resolve_config( + "user", + value, + default=None, + base_type=UserConfig, + ) + + @staticmethod + def resolve_user_config( + value: object, default: UserConfig | None + ) -> UserConfig | None: + if value is None: + return None + return ServerConfig.resolve_config( + "user", + value, + default=default, + base_type=UserConfig, + ) ConfigT = TypeVar("ConfigT", bound=TasksetConfig) -class Taskset(RuntimeOwnerMixin[ConfigT], Generic[ConfigT]): +class Taskset(Generic[ConfigT]): config: ConfigT + task_type: type[Task] = Task def __init_subclass__(cls, **kwargs: object) -> None: super().__init_subclass__(**kwargs) @@ -84,26 +168,24 @@ def __init__(self, config: ConfigSource = None): config_type = registered_config_type(type(self), TasksetConfig) self.config = cast(ConfigT, coerce_config(config_type, config)) with config_ref_context(self.config): - self.initialize_runtime_refresh() - resolved_taskset_id = self.config.taskset_id - if resolved_taskset_id is not None and not isinstance( - resolved_taskset_id, str - ): - raise TypeError("taskset_id must be a string.") - self.taskset_id = resolved_taskset_id or type(self).__name__ - system_prompt_value = self.load_system_prompt(self.config) + resolved_id = self.config.id + if resolved_id is not None and not isinstance(resolved_id, str): + raise TypeError("taskset id must be a string.") + self.id = resolved_id or type(self).__name__ self.system_prompt = normalize_system_prompt( - system_prompt_value, + self.load_system_prompt(self.config), field_name="taskset.system_prompt", ) - self.initialize_runtime_user(self.config.user) - self.bindings: BindingSources = self.config.bindings.entries( - "taskset.bindings" - ) - self.objects: Objects = self.load_objects(self.config.objects) - self.artifacts = self.load_artifacts(self.config.artifacts) - self.initialize_runtime_toolsets(self.config, self.config.toolsets) - self.initialize_runtime_handlers() + self.user = self.config.user + self.toolsets = TasksetConfig.enabled_toolsets(self.config.toolsets) + self.handlers = self.load_handlers() + self.signals = build_signals(self) + for signal in self.signals: + if signal["kind"] == "advantage": + raise ValueError( + "Taskset signals must be metrics or rewards; configure " + "env advantages with Env(advantage=...)." + ) self._dataset: Dataset | None = None self._eval_dataset: Dataset | None = None @@ -114,10 +196,29 @@ def get_upload_dirs(self) -> dict[str, Traversable | Path]: skills = self.get_skills_dir() return {} if skills is None else {"skills": skills} + def load_system_prompt(self, config: ConfigT) -> SystemPrompt: + return config.system_prompt + + def load_handlers(self) -> dict[LifecycleKind, list[Handler]]: + handlers: dict[LifecycleKind, list[Handler]] = { + "stop": [], + "setup": [], + "update": [], + "cleanup": [], + "teardown": [], + } + for kind in ("stop", "setup", "update", "cleanup", "teardown"): + handlers[kind].extend(discover_decorated(self, kind)) + return handlers + + @property + def has_group_signals(self) -> bool: + return any(signal["stage"] == "group" for signal in self.signals) + def to_task(self, task: Task | JsonData) -> Task: if isinstance(task, Task): - return prepare_task(task, self.taskset_id) - return task_from_dataset_record(task, self.taskset_id) + return prepare_task(task) + return task_from_dataset_record(task, self.task_type) def load_tasks(self, split: TaskSplit = "train") -> Tasks: if split not in ("train", "eval"): @@ -128,30 +229,27 @@ async def init_group( self, task: Task, num_rollouts: int ) -> tuple[list[Task], list[State]]: tasks = [task for _ in range(num_rollouts)] - return tasks, [State.for_task(task) for task in tasks] + return tasks, [State(task_id=task.task_id) for task in tasks] def get_dataset(self) -> Dataset: if self._dataset is None: with config_ref_context(self.config): - self._dataset = dataset_from_result( - self.load_tasks(split="train"), self.taskset_id + self._dataset = dataset_from_result_typed( + self.load_tasks(split="train"), self.task_type ) return self._dataset def get_eval_dataset(self) -> Dataset: if self._eval_dataset is None: with config_ref_context(self.config): - self._eval_dataset = dataset_from_result( - self.load_tasks(split="eval"), self.taskset_id + self._eval_dataset = dataset_from_result_typed( + self.load_tasks(split="eval"), self.task_type ) return self._eval_dataset def __iter__(self): for record in self.get_dataset(): - yield task_from_dataset_record(dict(record), self.taskset_id) + yield self.to_task(dict(record)) def __len__(self) -> int: return len(self.get_dataset()) - - def load_system_prompt(self, config: ConfigT) -> SystemPrompt: - return config.system_prompt diff --git a/verifiers/v1/toolset.py b/verifiers/v1/toolset.py index 1ced144fe1..430016020b 100644 --- a/verifiers/v1/toolset.py +++ b/verifiers/v1/toolset.py @@ -1,55 +1,30 @@ -from collections.abc import Iterable -from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Generic, Literal, TypeAlias, TypeVar, cast, final - -from pydantic import StrictBool, model_validator -from verifiers.types import Tool - -from .artifact import Artifacts, ArtifactsConfig -from .config import ( - CallableEntry, - Config, - ConfigSource, - resolve_config_object, -) -from .sandbox import SandboxConfig -from .utils.binding_utils import BindingSources, BindingsConfig, binding_sources -from .utils.binding_utils import ObjectsConfig +from __future__ import annotations + +import inspect +import importlib.util +import asyncio +from collections.abc import Callable, Coroutine, Mapping +from dataclasses import dataclass +from typing import Generic, Literal, TypeAlias, TypeVar, cast + +from pydantic import BaseModel, Field, model_validator + +from .config import Config, ConfigSource +from .runtime import RuntimeConfig, SubprocessRuntimeConfig +from .types import JsonData from .utils.config_utils import ( coerce_config, config_type_from_class, + explicit_config_data, + import_config_ref, registered_config_type, register_config_type, ) -from .utils.config_callable_utils import config_callables -from .types import Handler, Objects -from .utils.toolset_utils import ( - collect_toolsets as collect_toolsets, - flatten_toolsets as flatten_toolsets, - iter_toolsets as iter_toolsets, - normalize_toolset as normalize_toolset, - normalize_toolset_collection as normalize_toolset_collection, - normalize_toolset_result as normalize_toolset_result, - tool_item as tool_item, - tool_items as tool_items, - tool_name as tool_name, -) - -if TYPE_CHECKING: - from .state import State - from .task import Task - -ToolsetCallableEntry: TypeAlias = CallableEntry | Handler - -class MCPToolConfig(Config): - command: str - args: list[str] = [] - env: dict[str, str] | None = None - cwd: str | None = None - -ToolEntryConfig: TypeAlias = str | MCPToolConfig +Scope: TypeAlias = Literal["rollout", "env"] +ServerPlacement: TypeAlias = Literal["dedicated", "colocated", "remote"] +ConfigT = TypeVar("ConfigT", bound="ServerConfig") class VisibilityConfig(Config): @@ -60,53 +35,170 @@ class VisibilityConfig(Config): def validate_visibility(self) -> "VisibilityConfig": if self.show is not None and self.hide is not None: raise ValueError("Visibility accepts show or hide, not both.") - for field_name, names in (("show", self.show), ("hide", self.hide)): - if names is not None and len(names) != len(set(names)): - raise ValueError(f"Visibility {field_name} contains duplicate names.") return self -class ToolsetConfig(VisibilityConfig): - tools: list[ToolEntryConfig] = [] - handler: str | None = None - bindings: BindingsConfig = BindingsConfig() - objects: ObjectsConfig = ObjectsConfig() - artifacts: ArtifactsConfig = ArtifactsConfig() - write: StrictBool = False - scope: Literal["rollout", "group", "global"] | None = None - sandbox: SandboxConfig | Literal["program"] | None = None - stops: list[CallableEntry] = [] - setups: list[CallableEntry] = [] - updates: list[CallableEntry] = [] - cleanups: list[CallableEntry] = [] - teardowns: list[CallableEntry] = [] +class ServerConfig(VisibilityConfig): + source: str | None = None + enabled: bool = True + scope: Scope = "rollout" + placement: ServerPlacement = "dedicated" + runtime: RuntimeConfig | None = Field(default_factory=SubprocessRuntimeConfig) + url: str | None = None + headers: dict[str, str] = Field(default_factory=dict) + env: dict[str, str] = Field(default_factory=dict) + resources: JsonData = Field(default_factory=dict) + startup_timeout_seconds: float = 18.0 + @model_validator(mode="after") + def validate_server(self) -> "ServerConfig": + if self.scope == "env" and self.placement == "colocated": + raise ValueError("env-scope servers cannot use colocated placement.") + if self.placement == "remote": + if not self.url: + raise ValueError("Remote server configs require url.") + return self + if self.url is not None: + raise ValueError("Only remote server configs may set url.") + return self -ConfigT = TypeVar("ConfigT", bound=ToolsetConfig) + def default_server_ref(self) -> str: + config_type = type(self) + module_name = config_type.__module__ + if module_name.startswith("verifiers.v1."): + raise ValueError( + f"{config_type.__name__} cannot infer a toolset implementation from " + "the framework package." + ) + package = ServerConfig.config_package(module_name) + class_name = config_type.__name__ + if module_name.endswith(".config"): + if class_name.endswith("ToolsetConfig"): + impl_name = f"{class_name.removesuffix('Config')}" + else: + basename = package.rsplit(".", 1)[-1] + impl_name = f"{ServerConfig.snake_to_pascal(basename)}Toolset" + return f"{package}.toolset:{impl_name}" + if class_name.endswith("Config"): + impl_name = f"{class_name.removesuffix('Config')}" + else: + impl_name = f"{class_name}Toolset" + return f"{module_name}:{impl_name}" + + def implementation_ref(self) -> str: + if self.placement == "remote": + raise ValueError("Remote server configs do not have an implementation ref.") + ref = type(self).resolve_ref(self.default_server_ref(), type(self)) + if not ref: + raise ValueError("Server implementation ref must be non-empty.") + return ref + + def load(self) -> "Toolset": + return Toolset.load_ref(self.implementation_ref(), self) + + @classmethod + def resolve_config( + cls, + name: str, + value: object, + *, + default: "ServerConfig | None", + base_type: type[ConfigT], + ) -> ConfigT: + if isinstance(value, base_type): + if default is not None: + cls.validate_default_source( + name, + source=value.source, + default=default, + base_type=base_type, + ) + return value + if value is not None and not isinstance(value, BaseModel | Mapping): + raise TypeError("Server config values must be mappings or config objects.") + data = explicit_config_data(cast(ConfigSource, value)) + source = data.get("source") + if default is not None: + cls.validate_default_source( + name, + source=source, + default=default, + base_type=base_type, + ) + config_type = type(default) + merged = default.model_dump(mode="json", exclude_none=True) + merged.update(data) + return cast(ConfigT, config_type.model_validate(merged)) + if data.get("enabled") is False and source is None: + raise ValueError( + f"Server {name!r} is not declared by the taskset and cannot be " + "disabled." + ) + if not isinstance(source, str) or not source: + raise ValueError( + f"Server {name!r} is not declared by the taskset; set source to a " + f"{base_type.__name__} class." + ) + config_type = cls.source_type(source, base_type) + return cast(ConfigT, config_type.model_validate(data)) + + @classmethod + def validate_default_source( + cls, + name: str, + *, + source: object, + default: "ServerConfig", + base_type: type[ConfigT], + ) -> None: + if source is None: + return + if not isinstance(source, str) or not source: + raise TypeError(f"Server {name!r} source must be a non-empty string.") + config_type = cls.source_type(source, base_type) + if config_type is not type(default): + raise TypeError( + f"Server {name!r} source must match taskset-defined " + f"{type(default).__name__}; got {config_type.__name__}." + ) + + @classmethod + def source_type(cls, source: str, base_type: type[ConfigT]) -> type[ConfigT]: + obj = import_config_ref(source) + if isinstance(obj, type) and issubclass(obj, base_type): + return obj + raise TypeError( + f"Server source {source!r} must point to a {base_type.__name__}." + ) + + @staticmethod + def config_package(module_name: str) -> str: + if module_name.endswith(".config"): + return module_name.rsplit(".", 1)[0] + return module_name.rsplit(".", 1)[0] + + @staticmethod + def resolve_ref(ref: str, config_type: type["ServerConfig"]) -> str: + module_name, separator, attr_path = ref.partition(":") + if not separator: + raise ValueError(f"Server ref {ref!r} must use 'module:object'.") + if module_name.startswith("."): + package = ServerConfig.config_package(config_type.__module__) + module_name = importlib.util.resolve_name(module_name, package) + return f"{module_name}:{attr_path}" + + @staticmethod + def snake_to_pascal(value: str) -> str: + return "".join(part[:1].upper() + part[1:] for part in value.split("_") if part) + + +class ToolsetConfig(ServerConfig): + pass -@dataclass(frozen=True) class Toolset(Generic[ConfigT]): - # Tool surface. - tools: "tuple[ToolEntry, ...]" = () - handler: Handler | None = None - show: tuple[str, ...] | None = None - hide: tuple[str, ...] | None = None - # Local dependencies and runtime policy. - bindings: BindingSources = field(default_factory=dict) - objects: Objects = field(default_factory=dict) - artifacts: Artifacts = field(default_factory=dict) - write: bool = False - scope: str | None = None - sandbox: SandboxConfig | Literal["program"] | None = None - # Lifecycle collections. - stops: tuple[Handler, ...] = () - setups: tuple[Handler, ...] = () - updates: tuple[Handler, ...] = () - cleanups: tuple[Handler, ...] = () - teardowns: tuple[Handler, ...] = () - # Config. - config: ConfigT | None = None + config: ConfigT + name: str def __init_subclass__(cls, **kwargs: object) -> None: super().__init_subclass__(**kwargs) @@ -114,179 +206,147 @@ def __init_subclass__(cls, **kwargs: object) -> None: cls, inherited=False, owner_base=Toolset, - config_base=ToolsetConfig, + config_base=ServerConfig, ) if config_type is not None: register_config_type(cls, config_type) - @final - def __init__( - self, - # Tool surface. - tools: "ToolEntries | None" = None, - handler: ToolsetCallableEntry | None = None, - show: Iterable[str] | None = None, - hide: Iterable[str] | None = None, - # Local dependencies and runtime policy. - bindings: BindingSources | BindingsConfig | None = None, - objects: ObjectsConfig | None = None, - artifacts: ArtifactsConfig | None = None, - write: bool | None = None, - scope: str | None = None, - sandbox: SandboxConfig | Literal["program"] | None = None, - # Lifecycle collections. - stops: Iterable[ToolsetCallableEntry] | None = None, - setups: Iterable[ToolsetCallableEntry] | None = None, - updates: Iterable[ToolsetCallableEntry] | None = None, - cleanups: Iterable[ToolsetCallableEntry] | None = None, - teardowns: Iterable[ToolsetCallableEntry] | None = None, - # Config. - config: ConfigSource = None, - ): - if config is not None: - if any( - value is not None - for value in ( - tools, - handler, - show, - hide, - bindings, - objects, - artifacts, - write, - scope, - sandbox, - stops, - setups, - updates, - cleanups, - teardowns, - ) - ): - raise ValueError( - "Toolset accepts either config or constructor fields, not both." - ) - config_type = registered_config_type(type(self), ToolsetConfig) - config_value = coerce_config(config_type, config) - tools = config_value.tools - handler = config_value.handler - show = config_value.show - hide = config_value.hide - bindings = config_value.bindings - objects = config_value.objects - artifacts = config_value.artifacts - write = config_value.write - scope = config_value.scope - sandbox = config_value.sandbox - stops = config_value.stops - setups = config_value.setups - updates = config_value.updates - cleanups = config_value.cleanups - teardowns = config_value.teardowns - else: - config_value = None - tool_values = tool_items(tools) - if show is not None and hide is not None: - raise ValueError("Toolset accepts show or hide, not both.") - if isinstance(show, str) or isinstance(hide, str): - raise TypeError("Toolset show/hide must be lists of names.") - show_names = tuple(show) if show is not None else None - hide_names = tuple(hide) if hide is not None else None - if show_names is not None and not all( - isinstance(name, str) for name in show_names - ): - raise TypeError("Toolset show must contain only strings.") - if hide_names is not None and not all( - isinstance(name, str) for name in hide_names - ): - raise TypeError("Toolset hide must contain only strings.") - resolved_handler: object = handler - if handler is not None: - resolved_handler = resolve_config_object(handler) - if not callable(resolved_handler): - raise TypeError("Toolset handler must resolve to a callable.") - if write is not None and not isinstance(write, bool): - raise TypeError("Toolset write must be a boolean.") - object.__setattr__(self, "tools", tuple(tool_values)) - object.__setattr__(self, "handler", cast(Handler | None, resolved_handler)) - object.__setattr__(self, "show", show_names) - object.__setattr__(self, "hide", hide_names) - object.__setattr__( - self, - "bindings", - binding_sources(bindings, "toolset.bindings"), - ) - object.__setattr__( - self, - "objects", - self.load_objects(ObjectsConfig.model_validate(objects or {})), - ) - object.__setattr__( - self, - "artifacts", - self.load_artifacts(ArtifactsConfig.model_validate(artifacts or {})), - ) - object.__setattr__(self, "write", bool(write)) - if scope is not None and scope not in {"rollout", "group", "global"}: - raise ValueError("Toolset scope must be 'rollout', 'group', or 'global'.") - object.__setattr__(self, "scope", scope) - if ( - sandbox is not None - and sandbox != "program" - and not isinstance(sandbox, SandboxConfig) - ): - raise TypeError("Toolset sandbox must be SandboxConfig or 'program'.") - object.__setattr__(self, "sandbox", sandbox) - object.__setattr__(self, "stops", tuple(config_callables(stops or (), "stop"))) - object.__setattr__( - self, "setups", tuple(config_callables(setups or (), "setup")) - ) - object.__setattr__( - self, "updates", tuple(config_callables(updates or (), "update")) - ) - object.__setattr__( - self, "cleanups", tuple(config_callables(cleanups or (), "cleanup")) - ) - object.__setattr__( - self, "teardowns", tuple(config_callables(teardowns or (), "teardown")) - ) - object.__setattr__(self, "config", config_value) + def __init__(self, config: ConfigSource = None): + config_type = registered_config_type(type(self), ServerConfig) + self.config = cast(ConfigT, coerce_config(config_type, config)) + self.name = type(self).default_name() + self.resources: dict[str, object] = {} + + def start(self) -> None: + return None + + def stop(self) -> None: + return None - def load_objects(self, config: ObjectsConfig) -> Objects: - return config.objects("toolset.objects") + def load_resources(self) -> None: + for method_name, method in inspect.getmembers(self, predicate=callable): + spec = getattr( + getattr(type(self), method_name, None), "__vf_resource__", None + ) + if not isinstance(spec, ResourceSpec): + continue + name = spec.name or method_name + value = method() + if inspect.isawaitable(value): + value = asyncio.run(cast(Coroutine[object, object, object], value)) + self.resources[name] = value - def load_artifacts(self, config: ArtifactsConfig) -> Artifacts: - return config.artifacts("toolset.artifacts") + @staticmethod + def load_ref(server: str, config: ServerConfig) -> "Toolset": + obj = import_config_ref(server) + if isinstance(obj, type) and issubclass(obj, Toolset): + return obj(config=config) + if callable(obj): + loader = cast(Callable[[ServerConfig], object], obj) + loaded = loader(config) + if isinstance(loaded, Toolset): + return loaded + raise TypeError(f"Server {server!r} must be a Toolset class or loader.") - async def get_object(self, name: str, task: "Task", state: "State") -> object: - return await state._runtime().resolve_owner_object(self, name, task, state) + @classmethod + def tool_specs(cls) -> dict[str, "ToolSpec"]: + specs: dict[str, ToolSpec] = {} + for _, member in inspect.getmembers(cls, predicate=callable): + spec = getattr(member, "__vf_tool__", None) + if not isinstance(spec, ToolSpec): + continue + tool_name = spec.name or getattr(member, "__name__", "") + if not isinstance(tool_name, str) or not tool_name: + raise TypeError("Tool names must be non-empty strings.") + if tool_name in specs: + raise ValueError(f"Tool {tool_name!r} is defined twice.") + specs[tool_name] = spec + return specs + + @classmethod + def default_name(cls) -> str: + name = cls.__name__ + if name.endswith("Toolset") and len(name) > len("Toolset"): + name = name[: -len("Toolset")] + return cls.name_from_class(name or "toolset") + + @staticmethod + def name_from_class(value: str) -> str: + result: list[str] = [] + for index, char in enumerate(value): + if char.isupper() and index > 0 and not value[index - 1].isupper(): + result.append("_") + result.append(char.lower()) + return "".join(result).replace("-", "_") + + +class ToolBinding(Config): + args: dict[str, str] = Field(default_factory=dict) + sets: dict[str, str] = Field(default_factory=dict) + extends: dict[str, str] = Field(default_factory=dict) + hidden: bool = False @dataclass(frozen=True) -class MCPTool: - command: str - args: tuple[str, ...] = () - env: dict[str, str] | None = None - cwd: str | None = None - - def __init__( - self, - command: str, - args: Iterable[str] = (), - env: dict[str, str] | None = None, - cwd: str | None = None, - ): - object.__setattr__(self, "command", command) - object.__setattr__(self, "args", tuple(args)) - object.__setattr__(self, "env", dict(env) if env is not None else None) - object.__setattr__(self, "cwd", cwd) - - -ToolEntry: TypeAlias = Handler | str | Tool | Toolset | MCPTool | MCPToolConfig -ToolEntries: TypeAlias = ToolEntry | Iterable[ToolEntry] -ToolsetItem: TypeAlias = Toolset | ToolEntry -ToolsetCollection: TypeAlias = ( - ToolsetItem | Iterable[ToolsetItem] | dict[str, ToolsetItem | ToolsetConfig] -) -Toolsets: TypeAlias = ToolsetCollection | None +class ToolSpec: + name: str | None + args: dict[str, str] + sets: dict[str, str] + extends: dict[str, str] + hidden: bool + + +@dataclass(frozen=True) +class ResourceSpec: + name: str | None + + +ToolFunc = TypeVar("ToolFunc", bound=Callable[..., object]) + + +def tool( + func: ToolFunc | None = None, + *, + args: Mapping[str, str] | None = None, + sets: Mapping[str, str] | None = None, + extends: Mapping[str, str] | None = None, + name: str | None = None, + hidden: bool = False, +) -> ToolFunc | Callable[[ToolFunc], ToolFunc]: + def decorate(item: ToolFunc) -> ToolFunc: + setattr( + item, + "__vf_tool__", + ToolSpec( + name=name, + args=dict(args or {}), + sets=dict(sets or {}), + extends=dict(extends or {}), + hidden=hidden, + ), + ) + return item + + if func is not None: + return decorate(func) + return decorate + + +ResourceFunc = TypeVar("ResourceFunc", bound=Callable[..., object]) + + +def resource( + func: ResourceFunc | None = None, + *, + name: str | None = None, +) -> ResourceFunc | Callable[[ResourceFunc], ResourceFunc]: + def decorate(item: ResourceFunc) -> ResourceFunc: + setattr(item, "__vf_resource__", ResourceSpec(name=name)) + return item + + if func is not None: + return decorate(func) + return decorate + + +ToolsetConfigs: TypeAlias = dict[str, ToolsetConfig] diff --git a/verifiers/v1/types.py b/verifiers/v1/types.py index e820c499f3..dec66a63a5 100644 --- a/verifiers/v1/types.py +++ b/verifiers/v1/types.py @@ -1,38 +1,46 @@ +from __future__ import annotations + from collections.abc import Awaitable, Callable, Iterable, Sequence +from dataclasses import dataclass from typing import Literal, TYPE_CHECKING, TypeAlias from datasets import Dataset +from pydantic import Field from verifiers.clients import Client -from verifiers.types import ClientConfig, Message, MessageContent, Messages +from verifiers.types import ( + ClientConfig, + Message, + MessageContent, + Messages, + Response, + SamplingArgs, + Tool, +) from typing_extensions import TypeAliasType +from .config import Config + if TYPE_CHECKING: + from renderers import Renderer, RendererPool + + from .mcp import MCPToolRegistry + from .runtime import Runtime + from .state import State from .task import Task + RendererHandle: TypeAlias = Renderer | RendererPool +else: + RendererHandle: TypeAlias = object + JsonScalar: TypeAlias = str | int | float | bool | None JsonValue = TypeAliasType( "JsonValue", JsonScalar | list["JsonValue"] | dict[str, "JsonValue"], ) JsonData: TypeAlias = dict[str, JsonValue] -ConfigValue = TypeAliasType( - "ConfigValue", - JsonScalar - | list["ConfigValue"] - | tuple["ConfigValue", ...] - | dict[str, "ConfigValue"], -) -ConfigData: TypeAlias = dict[str, ConfigValue] HandlerResult: TypeAlias = ( - JsonValue - | JsonData - | ConfigData - | Message - | Messages - | MessageContent - | Sequence[float] - | None + JsonValue | JsonData | Message | Messages | MessageContent | Sequence[float] | None ) Handler: TypeAlias = Callable[..., HandlerResult | Awaitable[HandlerResult]] @@ -42,12 +50,117 @@ PromptMessage: TypeAlias = Message | JsonData PromptInput: TypeAlias = str | Sequence[PromptMessage] -ModelClient: TypeAlias = Client | ClientConfig -RuntimeObject: TypeAlias = object -RuntimeData: TypeAlias = dict[str, RuntimeObject] -RuntimeCallableResult: TypeAlias = RuntimeObject | Awaitable[RuntimeObject] -RuntimeCallable: TypeAlias = Callable[..., RuntimeCallableResult] -ObjectFactoryResult: TypeAlias = RuntimeObject | Awaitable[RuntimeObject] -ObjectFactory: TypeAlias = Callable[..., ObjectFactoryResult] -Objects: TypeAlias = dict[str, ObjectFactory] -ToolParameters: TypeAlias = dict[str, RuntimeObject] + +class ModelConfig(Config): + client: ClientConfig = Field(default_factory=ClientConfig) + model: str + sampling_args: JsonData = Field(default_factory=dict) + + +@dataclass(frozen=True) +class ModelClient: + config: ModelConfig + client: Client + renderer: RendererHandle | None = None + + async def get_response( + self, + *, + prompt: Messages, + state: "State | None" = None, + model: str | None = None, + sampling_args: SamplingArgs | None = None, + tools: list[Tool] | None = None, + ) -> Response: + kwargs: dict[str, object] = {} + if state is not None: + kwargs["state"] = client_state_record(state) + return await self.client.get_response( + prompt=prompt, + model=model or self.config.model, + sampling_args=sampling_args or dict(self.config.sampling_args), + tools=tools, + **kwargs, + ) + + def get_renderer(self) -> RendererHandle | None: + if self.renderer is not None: + return self.renderer + if self.config.client.client_type != "renderer": + return None + from verifiers.clients import RendererClient + + if not isinstance(self.client, RendererClient): + return None + return self.client.get_renderer( + self.config.model, + sampling_args=dict(self.config.sampling_args), + ) + + +@dataclass +class Context: + task: Task + state: State + model_client: ModelClient + teacher: ModelClient | None = None + runtime: Runtime | None = None + toolsets: MCPToolRegistry | None = None + user: MCPToolRegistry | None = None + parent: Context | None = None + score: bool = False + scoring: bool = False + + @property + def client(self) -> Client: + return self.model_client.client + + @property + def model(self) -> str: + return self.model_client.config.model + + @property + def sampling_args(self) -> SamplingArgs: + return dict(self.model_client.config.sampling_args) + + def has_active_scoring(self) -> bool: + context: Context | None = self + while context is not None: + if context.scoring: + return True + context = context.parent + return False + + +def client_state_record(state: "State") -> JsonData: + from .utils.json_utils import json_data, json_value + + transcript: list[JsonValue] = [] + for turn in state.transcript: + tokens = ( + json_value(turn.tokens.model_dump(mode="json")) if turn.tokens else None + ) + transcript.append( + json_value( + { + "prompt": turn.prompt, + "completion": turn.completion, + "tokens": tokens, + "reward": turn.reward, + "is_truncated": turn.is_truncated, + }, + context="client state turn", + ) + ) + record = { + "id": state.id, + "task_id": state.task_id, + "group_id": state.group_id, + "extras": state.extras, + "metadata": state.metadata, + "transcript": transcript, + } + return json_data( + {key: value for key, value in record.items() if value is not None}, + context="client state", + ) diff --git a/verifiers/v1/user.py b/verifiers/v1/user.py index 8c067db1aa..e4a04eca39 100644 --- a/verifiers/v1/user.py +++ b/verifiers/v1/user.py @@ -1,132 +1,67 @@ -from collections.abc import Sequence -from typing import TYPE_CHECKING, Generic, Literal, TypeVar, cast, final +from __future__ import annotations -from verifiers.types import Message, UserMessage -from verifiers.utils.message_utils import normalize_messages +from collections.abc import Callable, Mapping +from typing import Generic, TypeVar -from .artifact import Artifacts, ArtifactsConfig -from .config import Config, ConfigSource -from .sandbox import SandboxConfig -from .utils.binding_utils import ( - BindingSources, - BindingsConfig, - ObjectsConfig, +from .toolset import ( + Scope, + ServerConfig, + Toolset, + tool, ) -from .utils.config_utils import ( - coerce_config, - config_type_from_class, - registered_config_type, - register_config_type, -) -from .utils.trajectory_utils import completion_from_trajectory -from .state import State -from .types import JsonData, Objects, PromptMessage - -if TYPE_CHECKING: - from .task import Task - -UserScope = Literal["rollout", "group", "global"] - - -class UserConfig(Config): - scope: UserScope = "rollout" - bindings: BindingsConfig = BindingsConfig() - objects: ObjectsConfig = ObjectsConfig() - artifacts: ArtifactsConfig = ArtifactsConfig() - sandbox: SandboxConfig | None = None - - -def state_messages( - state: State, transcript: Sequence[PromptMessage] | None = None -) -> list[Message]: - if transcript is not None: - return normalize_messages(transcript, field_name="user.transcript") - prompt = state.get("prompt") - completion = state.get("completion") - if isinstance(prompt, list) and isinstance(completion, list): - return normalize_messages( - [ - *cast(list[PromptMessage], prompt), - *cast(list[PromptMessage], completion), - ], - field_name="state.messages", - ) - if isinstance(completion, list): - return normalize_messages( - cast(list[PromptMessage], completion), field_name="state.completion" - ) - trajectory = state.get("trajectory") - if isinstance(trajectory, Sequence) and not isinstance(trajectory, str): - return normalize_messages( - completion_from_trajectory(cast(Sequence[JsonData], trajectory)), - field_name="state.trajectory", - ) - return [] - - -ConfigT = TypeVar("ConfigT", bound=UserConfig) -user_type_registry: dict[type[UserConfig], type["User"]] = {} - - -class User(Generic[ConfigT]): - config: ConfigT - scope: UserScope - bindings: BindingSources - objects: Objects - artifacts: Artifacts - sandbox: SandboxConfig | None - - def __init_subclass__(cls, **kwargs: object) -> None: - super().__init_subclass__(**kwargs) - config_type = config_type_from_class( - cls, - inherited=False, - owner_base=User, - config_base=UserConfig, - ) - if config_type is not None: - register_config_type(cls, config_type) - user_type_registry[cast(type[UserConfig], config_type)] = cls - - @final - def __init__( - self, - *, - config: ConfigSource = None, - ): - config_type = registered_config_type(type(self), UserConfig) - self.config = cast(ConfigT, coerce_config(config_type, config)) - if self.config.scope not in {"rollout", "group", "global"}: - raise ValueError("User scope must be 'rollout', 'group', or 'global'.") - bindings = self.config.bindings.entries("User bindings", key_style="arg") - if "messages" in bindings: - raise ValueError("User messages are provided directly to get_response.") - self.scope = self.config.scope - self.bindings = bindings - self.objects = self.load_objects(self.config.objects) - self.artifacts = self.load_artifacts(self.config.artifacts) - self.sandbox = self.config.sandbox - - def load_objects(self, config: ObjectsConfig) -> Objects: - return config.objects("user.objects") - - def load_artifacts(self, config: ArtifactsConfig) -> Artifacts: - return config.artifacts("user.artifacts") - - async def get_object(self, name: str, task: "Task", state: State) -> object: - return await state._runtime().resolve_owner_object(self, name, task, state) - - async def get_response( - self, task: "Task", state: State, messages: list[Message] - ) -> list[UserMessage]: - return [] -def user_from_config(config: UserConfig) -> User: - for config_type in type(config).__mro__: - if not issubclass(config_type, UserConfig): - continue - user_type = user_type_registry.get(config_type) - if user_type is not None: - return user_type(config=config) - raise TypeError(f"No User subclass is registered for {type(config).__name__}.") +class UserConfig(ServerConfig): + scope: Scope = "rollout" + + def default_server_ref(self) -> str: + config_type = type(self) + module_name = config_type.__module__ + if module_name.startswith("verifiers.v1."): + raise ValueError( + f"{config_type.__name__} cannot infer a user implementation from " + "the framework package." + ) + if module_name.endswith(".config"): + package = type(self).config_package(module_name) + if config_type.__name__ == "UserConfig": + impl_name = "User" + elif config_type.__name__.endswith("Config"): + impl_name = config_type.__name__.removesuffix("Config") + else: + impl_name = "User" + return f"{package}.user:{impl_name}" + if config_type.__name__.endswith("Config"): + ref = f"{module_name}:{config_type.__name__.removesuffix('Config')}" + else: + ref = f"{module_name}:{config_type.__name__}User" + return type(self).resolve_ref(ref, config_type) + + def load(self) -> "User": + server = self.implementation_ref() + user = Toolset.load_ref(server, self) + if not isinstance(user, User): + raise TypeError(f"User server {server!r} did not return a User.") + return user + + +UserConfigT = TypeVar("UserConfigT", bound=UserConfig) + + +class User(Toolset[UserConfigT], Generic[UserConfigT]): + pass + + +UserFunc = TypeVar("UserFunc", bound=Callable[..., object]) + + +def user( + func: UserFunc | None = None, + *, + args: Mapping[str, str] | None = None, + sets: Mapping[str, str] | None = None, + extends: Mapping[str, str] | None = None, +) -> UserFunc | Callable[[UserFunc], UserFunc]: + return tool( + func, args=args, sets=sets, extends=extends, name="respond", hidden=True + ) diff --git a/verifiers/v1/utils/binding_utils.py b/verifiers/v1/utils/binding_utils.py deleted file mode 100644 index 0b72f21a6d..0000000000 --- a/verifiers/v1/utils/binding_utils.py +++ /dev/null @@ -1,327 +0,0 @@ -import inspect -from collections.abc import Set -from typing import Literal, TypeAlias, cast - -from pydantic import ConfigDict, model_validator -from typing_extensions import Self - -from ..config import Config, validate_serializable_value -from ..types import ConfigData, Handler, ObjectFactory, Objects -from .config_utils import resolve_config_object -from .object_utils import validate_object_factory_spec, validate_object_loader_spec - - -BindingRoot: TypeAlias = Literal[ - "task", - "state", - "tasks", - "states", - "runtime", - "objects", - "tools", - "taskset", - "harness", -] -CallableBindingSource: TypeAlias = Handler | ConfigData -BindingSource: TypeAlias = str | CallableBindingSource -BindingSources: TypeAlias = dict[str, BindingSource] -ObjectRefs: TypeAlias = dict[str, str] - - -class BindingsConfig(Config): - model_config = ConfigDict(extra="allow") - - @model_validator(mode="before") - @classmethod - def validate_mapping_input(cls, value: object) -> object: - if isinstance(value, BindingsConfig): - return value - if value is None: - return {} - if not isinstance(value, dict): - raise TypeError("BindingsConfig must be a mapping.") - for key, source in value.items(): - if not isinstance(key, str): - raise TypeError("BindingsConfig keys must be strings.") - validate_serializable_value(source, f"bindings.{key}") - validate_binding_source(source, f"bindings source for {key!r}") - return value - - @model_validator(mode="after") - def validate_config_entries(self) -> Self: - for key, source in self.raw_entries().items(): - validate_serializable_value(source, f"bindings.{key}") - validate_binding_source(source, f"bindings source for {key!r}") - return self - - def raw_entries(self) -> BindingSources: - return cast(BindingSources, dict(self.model_extra or {})) - - def entries( - self, - field: str = "bindings", - *, - allow_objects: bool = True, - validate_sources: bool = True, - key_style: Literal["callable", "arg"] = "callable", - ) -> BindingSources: - result: BindingSources = {} - for raw_key, source in self.raw_entries().items(): - if not isinstance(raw_key, str): - raise TypeError(f"{field} keys must be strings.") - if key_style == "callable": - binding_key_parts(raw_key) - elif not raw_key or "." in raw_key: - raise ValueError(f"{field} keys must be argument names.") - if validate_sources: - validate_binding_source( - source, - f"{field} source for {raw_key!r}", - allow_objects=allow_objects, - ) - result[raw_key] = source - return result - - -def binding_sources( - value: BindingSources | BindingsConfig | None, - field: str = "bindings", -) -> BindingSources: - if value is None: - return {} - if isinstance(value, BindingsConfig): - return value.entries(field) - if not isinstance(value, dict): - raise TypeError(f"{field} must be a mapping.") - result: BindingSources = {} - for raw_key, source in value.items(): - if not isinstance(raw_key, str): - raise TypeError(f"{field} keys must be strings.") - binding_key_parts(raw_key) - validate_binding_source(source, f"{field} source for {raw_key!r}") - result[raw_key] = source - return result - - -class ObjectsConfig(Config): - model_config = ConfigDict(extra="allow") - - @model_validator(mode="before") - @classmethod - def validate_mapping_input(cls, value: object) -> object: - if isinstance(value, ObjectsConfig): - return value - if value is None: - return {} - if not isinstance(value, dict): - raise TypeError("ObjectsConfig must be a mapping.") - for key, source in value.items(): - if not isinstance(key, str): - raise TypeError("ObjectsConfig keys must be strings.") - if not isinstance(source, str): - raise TypeError(f"objects entry {key!r} must be an import ref string.") - validate_object_loader_spec(source, f"objects entry {key!r}") - return value - - def refs(self) -> ObjectRefs: - return cast(ObjectRefs, dict(self.model_extra or {})) - - def objects(self, field: str = "objects") -> Objects: - resolved: Objects = {} - for name, source in self.refs().items(): - factory = resolve_config_object(source) - validate_object_factory_spec(factory, f"{field}.{name}") - resolved[name] = cast(ObjectFactory, factory) - return resolved - - @model_validator(mode="after") - def validate_entries(self) -> Self: - for key, source in self.refs().items(): - if not isinstance(source, str): - raise TypeError(f"objects entry {key!r} must be an import ref string.") - validate_object_loader_spec(source, f"objects entry {key!r}") - return self - - -VALID_BINDING_ROOTS: frozenset[str] = frozenset( - { - "task", - "state", - "tasks", - "states", - "runtime", - "objects", - "tools", - "taskset", - "harness", - } -) -ROLLOUT_FRAMEWORK_ARGS: frozenset[str] = frozenset( - { - "answer", - "completion", - "error", - "example_id", - "info", - "metrics", - "prompt", - "question", - "reward", - "runtime", - "state", - "task", - "task_id", - "timing", - "trajectory", - } -) -GROUP_FRAMEWORK_ARGS: frozenset[str] = frozenset({"states", "tasks"}) - - -def validate_binding_source( - source: object, context: str, *, allow_objects: bool = True -) -> None: - if ( - not isinstance(source, str) - and not callable(source) - and not isinstance(source, dict) - ): - raise TypeError(f"{context} must be a framework path or callable.") - root = binding_source_root(source) - validate_binding_source_root(root, context, allow_objects=allow_objects) - if root == "objects": - if not isinstance(source, str): - raise TypeError(f"{context} must be a string source.") - binding_object_name(source) - if root in {"taskset", "harness"}: - if not isinstance(source, str): - raise TypeError(f"{context} must be a string source.") - owner_object_name(source) - if isinstance(source, dict): - validate_callable_source(cast(ConfigData, source), context) - - -def validate_callable_source(source: ConfigData, context: str) -> None: - if "fn" not in source: - raise TypeError(f"{context} mapping sources must use an 'fn' key.") - unknown = set(source) - {"fn"} - if unknown: - raise ValueError(f"{context} has unknown keys: {sorted(unknown)}.") - - -def function_name(fn: Handler) -> str: - name = getattr(fn, "__name__", None) - if not isinstance(name, str) or not name: - raise ValueError("Callable bindings require a stable __name__.") - return name - - -def binding_key_parts(key: str) -> tuple[str, str]: - if not isinstance(key, str): - raise TypeError("Binding keys must be strings.") - target, separator, arg_name = key.partition(".") - if separator != "." or not target or not arg_name or "." in arg_name: - raise ValueError(f"Binding key {key!r} must be 'callable.arg'.") - return target, arg_name - - -def binding_source_root(source: object) -> BindingRoot | None: - if not isinstance(source, str): - return None - root, _, _ = source.partition(".") - if root in VALID_BINDING_ROOTS: - return cast(BindingRoot, root) - raise ValueError( - "Binding string sources must start with task, state, tasks, states, " - f"runtime, objects, tools, taskset, or harness; got {source!r}." - ) - - -def validate_binding_source_root( - root: BindingRoot | None, context: str, *, allow_objects: bool = True -) -> None: - if root is None: - return - if root == "objects" and not allow_objects: - raise ValueError(f"{context} cannot use objects.* sources.") - - -def binding_object_name(source: str) -> str: - if not isinstance(source, str): - raise TypeError("Object binding source must be a string.") - root, separator, tail = source.partition(".") - if root != "objects" or not separator: - raise ValueError("Object binding source must be 'objects.name'.") - name, _, _ = tail.partition(".") - if not name: - raise ValueError("Object binding source must be 'objects.name'.") - return name - - -def owner_object_name(source: str) -> str: - if not isinstance(source, str): - raise TypeError("Owner object binding source must be a string.") - root, separator, tail = source.partition(".") - if root not in {"taskset", "harness"} or not separator: - raise ValueError("Owner object binding source must be 'owner.objects.name'.") - objects_root, objects_separator, object_tail = tail.partition(".") - if objects_root != "objects" or not objects_separator: - raise ValueError("Owner object binding source must be 'owner.objects.name'.") - name, _, _ = object_tail.partition(".") - if not name: - raise ValueError("Owner object binding source must be 'owner.objects.name'.") - return name - - -def validate_bound_arg( - fn: object, - arg_name: str, - context: str, - protected_args: Set[str] = frozenset(), - *, - allow_reserved: bool = False, -) -> None: - if arg_name in protected_args: - return - if not allow_reserved and arg_name in {"task", "state", "runtime"}: - raise ValueError(f"{context} cannot bind reserved arg {arg_name!r}.") - if not callable(fn): - raise TypeError(f"{context} target is not callable.") - try: - signature = inspect.signature(fn) - except (TypeError, ValueError) as exc: - raise TypeError(f"{context} target signature cannot be inspected.") from exc - if arg_name not in signature.parameters: - name = ( - getattr(fn, "__name__", None) - or getattr(fn, "name", None) - or type(fn).__name__ - ) - raise TypeError( - f"{context} targets {name!r}, but {name!r} does not declare " - f"arg {arg_name!r}." - ) - - -def same_callable(left: Handler, right: Handler) -> bool: - if left is right: - return True - left_self = getattr(left, "__self__", None) - right_self = getattr(right, "__self__", None) - left_func = getattr(left, "__func__", None) - right_func = getattr(right, "__func__", None) - return left_self is right_self and left_func is not None and left_func is right_func - - -def read_path(value: object, path: str) -> object: - current = value - for part in path.split("."): - if not part: - raise ValueError(f"Invalid empty path segment in {path!r}.") - if isinstance(current, dict): - current = cast(ConfigData, current)[part] - elif isinstance(current, list): - current = current[int(part)] - else: - current = getattr(current, part) - return current diff --git a/verifiers/v1/utils/config_callable_utils.py b/verifiers/v1/utils/config_callable_utils.py deleted file mode 100644 index a8b0603c86..0000000000 --- a/verifiers/v1/utils/config_callable_utils.py +++ /dev/null @@ -1,127 +0,0 @@ -import functools -import inspect -from collections.abc import Iterable -from typing import Literal, TypeAlias, cast - -from pydantic import BaseModel - -from .config_utils import resolve_config_object -from ..types import ConfigData, Handler - -CallableKind: TypeAlias = Literal[ - "stop", "setup", "update", "metric", "reward", "advantage", "cleanup", "teardown" -] - -CALLABLE_KIND_FIELDS: dict[CallableKind, str] = { - "stop": "stops", - "setup": "setups", - "update": "updates", - "metric": "metrics", - "reward": "rewards", - "advantage": "advantages", - "cleanup": "cleanups", - "teardown": "teardowns", -} - - -def merge_config_callables( - values: Iterable[Handler], - config: object, - kind: CallableKind, -) -> list[Handler]: - return [*config_callables(values, kind), *config_callables(config, kind)] - - -def merge_config_handler_map( - values: dict[CallableKind, Iterable[Handler]], - config: object, -) -> dict[CallableKind, list[Handler]]: - return { - kind: merge_config_callables( - constructor_values, getattr(config, CALLABLE_KIND_FIELDS[kind]), kind - ) - for kind, constructor_values in values.items() - } - - -def config_callables(value: object, kind: CallableKind) -> list[Handler]: - if value is None: - return [] - if isinstance(value, str): - return [callable_config_item(value, kind)] - if isinstance(value, dict): - return [callable_config_item(value, kind)] - if isinstance(value, Iterable): - return [callable_config_item(item, kind) for item in value] - return [callable_config_item(value, kind)] - - -def callable_config_item(value: object, kind: CallableKind) -> Handler: - value = resolve_config_object(value) - if isinstance(value, BaseModel): - value = value.model_dump(exclude_none=True) - if isinstance(value, dict): - return callable_from_mapping(cast(ConfigData, value), kind) - if not callable(value): - raise TypeError(f"{kind} config entries must resolve to callables.") - return cast(Handler, value) - - -def callable_from_mapping(spec: ConfigData, kind: CallableKind) -> Handler: - allowed = callable_config_keys(kind) - unknown = set(spec) - allowed - if unknown: - raise ValueError(f"{kind} callable config has unknown keys: {sorted(unknown)}.") - if bool(spec.get("skip", False)): - raise ValueError( - f"{kind} callable config should be removed instead of skipped." - ) - fn = resolve_config_object(spec.get("fn")) - if not callable(fn): - raise TypeError(f"{kind} callable config requires callable fn.") - metadata = {key: spec[key] for key in spec if key not in {"fn", "skip"}} - return configured_callable(cast(Handler, fn), kind, metadata) - - -def callable_config_keys(kind: CallableKind) -> set[str]: - keys = {"fn", "priority", "skip"} - if kind in {"update", "metric", "reward", "cleanup"}: - keys.add("stage") - if kind == "reward": - keys.add("weight") - return keys - - -def configured_callable( - fn: Handler, - kind: CallableKind, - metadata: ConfigData, -) -> Handler: - if not metadata: - return fn - - @functools.wraps(fn) - async def wrapper(**kwargs: object) -> object: - result = fn(**kwargs) - if inspect.isawaitable(result): - return await result - return result - - setattr(wrapper, "__signature__", inspect.signature(fn)) - setattr(wrapper, kind, True) - if "priority" in metadata: - priority = metadata["priority"] - if not isinstance(priority, int) or isinstance(priority, bool): - raise TypeError(f"{kind} priority must be an integer.") - setattr(wrapper, f"{kind}_priority", priority) - if "stage" in metadata: - stage = metadata["stage"] - if stage not in {"rollout", "group"}: - raise ValueError(f"{kind} stage must be 'rollout' or 'group'.") - setattr(wrapper, f"{kind}_stage", stage) - if "weight" in metadata: - weight = metadata["weight"] - if not isinstance(weight, int | float) or isinstance(weight, bool): - raise TypeError("reward weight must be numeric.") - setattr(wrapper, "reward_weight", float(weight)) - return cast(Handler, wrapper) diff --git a/verifiers/v1/utils/config_utils.py b/verifiers/v1/utils/config_utils.py index 4d5431e712..1b98d9eb69 100644 --- a/verifiers/v1/utils/config_utils.py +++ b/verifiers/v1/utils/config_utils.py @@ -1,6 +1,6 @@ import importlib import sys -from collections.abc import Iterator +from collections.abc import Iterator, Mapping from contextlib import contextmanager from contextvars import ContextVar from typing import TypeVar, cast, get_args, get_origin, get_type_hints @@ -8,19 +8,13 @@ from pydantic import BaseModel from pydantic_core import PydanticUndefined -from ..types import ConfigData, ConfigValue - ConfigT = TypeVar("ConfigT", bound=BaseModel) ConfigOwner = type[object] config_type_registry: dict[ConfigOwner, type[BaseModel]] = {} FRAMEWORK_CONFIG_MODULES = { "verifiers.v1.config", "verifiers.v1.env", - "verifiers.v1.artifact", "verifiers.v1.harness", - "verifiers.v1.model", - "verifiers.v1.program", - "verifiers.v1.sandbox", "verifiers.v1.taskset", "verifiers.v1.toolset", "verifiers.v1.user", @@ -31,17 +25,18 @@ def explicit_config_data( - value: BaseModel | ConfigData | None, target: type[BaseModel] | None = None -) -> ConfigData: + value: BaseModel | Mapping[str, object] | None, + target: type[BaseModel] | None = None, +) -> dict[str, object]: if value is None: - data: ConfigData = {} + data: dict[str, object] = {} elif isinstance(value, BaseModel): data = explicit_model_config_data(value) if target is not None: data = { key: item for key, item in data.items() if key in target.model_fields } - elif isinstance(value, dict): + elif isinstance(value, Mapping): data = string_mapping(value) else: raise TypeError("Config must be a mapping or config object.") @@ -49,7 +44,8 @@ def explicit_config_data( def coerce_config( - config_cls: type[ConfigT], value: BaseModel | ConfigData | None = None + config_cls: type[ConfigT], + value: BaseModel | Mapping[str, object] | None = None, ) -> ConfigT: if value is None: return config_cls() @@ -170,25 +166,26 @@ def resolve_config_annotation(owner_type: ConfigOwner, annotation: object) -> ob def resolved_config_data( - value: BaseModel | ConfigData | None, target: type[BaseModel] | None = None -) -> ConfigData: + value: BaseModel | Mapping[str, object] | None, + target: type[BaseModel] | None = None, +) -> dict[str, object]: if value is None: - data: ConfigData = {} + data: dict[str, object] = {} elif isinstance(value, BaseModel): - data = cast(ConfigData, value.model_dump(exclude_none=True)) + data = dict(value.model_dump(exclude_none=True)) if target is not None: data = { key: item for key, item in data.items() if key in target.model_fields } - elif isinstance(value, dict): + elif isinstance(value, Mapping): data = string_mapping(value) else: raise TypeError("Config must be a mapping or config object.") return data -def explicit_model_config_data(value: BaseModel) -> ConfigData: - data: ConfigData = {} +def explicit_model_config_data(value: BaseModel) -> dict[str, object]: + data: dict[str, object] = {} for key in value.model_fields_set: item = getattr(value, key) data[key] = config_dump_value(item) @@ -199,16 +196,17 @@ def explicit_model_config_data(value: BaseModel) -> ConfigData: return data -def config_dump_value(value: object) -> ConfigValue: +def config_dump_value(value: object) -> object: if isinstance(value, BaseModel): return explicit_model_config_data(value) - if isinstance(value, dict): + if isinstance(value, Mapping): return { - key: config_dump_value(item) for key, item in string_mapping(value).items() + key: config_dump_value(item) + for key, item in string_mapping(cast(Mapping[str, object], value)).items() } if isinstance(value, list | tuple): return [config_dump_value(item) for item in value] - return cast(ConfigValue, value) + return value def resolve_config_object(value: object) -> object: @@ -250,7 +248,9 @@ def config_ref_parts(ref: str) -> tuple[str, str]: @contextmanager -def config_ref_context(config: BaseModel | ConfigData | None) -> Iterator[None]: +def config_ref_context( + config: BaseModel | Mapping[str, object] | None, +) -> Iterator[None]: module_name = config_ref_module(config) if module_name is None: yield @@ -262,7 +262,7 @@ def config_ref_context(config: BaseModel | ConfigData | None) -> Iterator[None]: _CONFIG_REF_MODULE.reset(token) -def config_ref_module(config: BaseModel | ConfigData | None) -> str | None: +def config_ref_module(config: BaseModel | Mapping[str, object] | None) -> str | None: if isinstance(config, BaseModel): module_name = type(config).__module__ if module_name not in FRAMEWORK_CONFIG_MODULES: @@ -270,12 +270,12 @@ def config_ref_module(config: BaseModel | ConfigData | None) -> str | None: return None -def string_mapping(value: dict) -> ConfigData: - result: ConfigData = {} +def string_mapping(value: Mapping[str, object]) -> dict[str, object]: + result: dict[str, object] = {} for key, item in value.items(): if not isinstance(key, str): raise TypeError("Config mappings require string keys.") - result[key] = cast(ConfigValue, item) + result[key] = item return result diff --git a/verifiers/v1/utils/endpoint_utils.py b/verifiers/v1/utils/endpoint_utils.py deleted file mode 100644 index 117e54770c..0000000000 --- a/verifiers/v1/utils/endpoint_utils.py +++ /dev/null @@ -1,828 +0,0 @@ -import asyncio -import json -import logging -import os -import time -import uuid -from collections.abc import Awaitable, Callable -from typing import Literal, Protocol, TypeAlias, cast - -from anthropic import Anthropic, AsyncAnthropic -from openai import AsyncOpenAI, OpenAI - -from verifiers.errors import Error, OverlongPromptError, TunnelError -from verifiers.types import ( - AssistantMessage, - ClientType, - ContentPart, - EndpointApi, - EndpointClient, - EndpointConfig, - MessageContent, - Messages, - Response, - SystemMessage, - Tool, - ToolCall, - ToolMessage, - UserMessage, -) -from verifiers.utils.interception_utils import ( - InterceptionServer, - deliver_response, - synthesize_stream, -) -from verifiers.utils.message_utils import normalize_messages -from verifiers.utils.response_utils import parse_response_message - -from ..runtime import ModelRequestContext, Runtime, TrajectoryVisibility -from ..state import State -from ..task import Task -from ..types import JsonData, PromptMessage, RuntimeObject, ToolParameters -from .serialization_utils import serializable - -VF_TRAJECTORY_VISIBILITY_HEADER = "x-verifiers-trajectory" -VF_ENDPOINT_API_KEY_VAR = "VF_ENDPOINT_API_KEY" -NormalizedEndpointApi: TypeAlias = Literal[ - "chat_completions", - "completions", - "responses", - "messages", -] -EndpointInterceptData: TypeAlias = dict[str, RuntimeObject] - - -class TunnelHandle(Protocol): - is_running: bool - url: str | None - - async def start(self) -> str: ... - - async def check_registered(self) -> bool: ... - - def sync_stop(self) -> None: ... - - -def client_from_state( - state: State, - api: EndpointApi | ClientType = "chat_completions", - *, - sync: bool = False, -) -> EndpointClient: - endpoint = endpoint_from_state(state) - return endpoint.client(state, api=api, sync=sync) - - -def endpoint_config_from_state( - state: State, - api: EndpointApi | ClientType = "chat_completions", -) -> EndpointConfig: - endpoint = endpoint_from_state(state) - return endpoint.config(state, api=api) - - -def endpoint_from_state(state: State) -> "Endpoint": - runtime = state._runtime() - harness = runtime.harness - if harness is None: - raise RuntimeError("State does not have an active model endpoint.") - endpoint = harness.endpoint - if not isinstance(endpoint, Endpoint): - raise RuntimeError("State does not have an active model endpoint.") - return endpoint - - -def endpoint_api_client_type( - api: NormalizedEndpointApi, -) -> Literal[ - "openai_chat_completions", - "openai_completions", - "openai_responses", - "anthropic_messages", -]: - if api == "chat_completions": - return "openai_chat_completions" - if api == "completions": - return "openai_completions" - if api == "responses": - return "openai_responses" - return "anthropic_messages" - - -def normalize_endpoint_api( - api: EndpointApi | ClientType, -) -> NormalizedEndpointApi: - if api in { - "chat_completions", - "openai_chat_completions", - "chat", - }: - return "chat_completions" - if api in {"responses", "openai_responses"}: - return "responses" - if api in {"messages", "anthropic_messages"}: - return "messages" - if api in {"completions", "openai_completions"}: - return "completions" - if api == "openai_chat_completions_token": - raise ValueError( - "state.get_client(...) does not expose token-level chat completions clients." - ) - if api == "renderer": - raise ValueError("state.get_client(...) does not expose renderer clients.") - if api == "nemorl_chat_completions": - raise ValueError( - "state.get_client(...) does not expose NeMoRL chat completions clients." - ) - raise ValueError(f"Unknown endpoint API {api!r}.") - - -class Endpoint: - TUNNEL_CHECK_INTERVAL = 60.0 - - def __init__( - self, - port: int | None = None, - secret: str | None = None, - use_tunnel: bool = False, - logger: logging.Logger | None = None, - ): - self.use_tunnel = use_tunnel - self.logger = logger or logging.getLogger(__name__) - self.server = InterceptionServer( - port if port is not None else 0, - secret=secret or os.environ.get("ENDPOINT_SECRET"), - ) - self.secret = self.server.secret - self._tunnel: TunnelHandle | None = None - self._tunnel_lock = asyncio.Lock() - self._tunnel_last_checked = 0.0 - self._rollout_queues: dict[str, asyncio.Queue[str]] = {} - - async def start(self) -> None: - await self.server.start() - - async def register_rollout( - self, - state: State, - tool_handler: object | None = None, - tool_defs: list[Tool] | None = None, - user_handler: object | None = None, - stop_handler: object | None = None, - model_handler: object | None = None, - ) -> str: - await self.start() - rollout_key = f"rollout_{uuid.uuid4().hex[:8]}" - request_queue = self.server.register_rollout( - rollout_key, - state=state, - tool_handler=tool_handler, - tool_defs=tool_defs, - user_handler=user_handler, - stop_handler=stop_handler, - model_handler=model_handler, - ) - self._rollout_queues[rollout_key] = cast(asyncio.Queue[str], request_queue) - endpoint_root_url = f"{await self.url_base()}/rollout/{rollout_key}" - api_key_var = f"{VF_ENDPOINT_API_KEY_VAR}_{rollout_key.upper()}" - state["endpoint_rollout_key"] = rollout_key - state["endpoint_root_url"] = endpoint_root_url - state["endpoint_base_url"] = f"{endpoint_root_url}/v1" - state["endpoint_api_key_var"] = api_key_var - return state["endpoint_base_url"] - - def client( - self, - state: State, - api: EndpointApi | ClientType = "chat_completions", - *, - sync: bool = False, - ) -> EndpointClient: - api = normalize_endpoint_api(api) - api_key = self.secret or "intercepted" - if api == "messages": - base_url = str(state["endpoint_root_url"]) - if sync: - return Anthropic(api_key=api_key, base_url=base_url) - return AsyncAnthropic(api_key=api_key, base_url=base_url) - base_url = str(state["endpoint_base_url"]) - if sync: - return OpenAI(api_key=api_key, base_url=base_url) - return AsyncOpenAI(api_key=api_key, base_url=base_url) - - def config( - self, - state: State, - api: EndpointApi | ClientType = "chat_completions", - ) -> EndpointConfig: - api = normalize_endpoint_api(api) - base_url = ( - str(state["endpoint_root_url"]) - if api == "messages" - else str(state["endpoint_base_url"]) - ) - return EndpointConfig( - model=state.get_model(), - base_url=base_url, - api_key_var=str(state["endpoint_api_key_var"]), - api_client_type=endpoint_api_client_type(api), - ) - - def unregister_rollout(self, rollout_key: str) -> None: - self._rollout_queues.pop(rollout_key, None) - self.server.unregister_rollout(rollout_key) - - def rollout_queue(self, rollout_key: str) -> asyncio.Queue[str]: - return self._rollout_queues[rollout_key] - - def get_request(self, request_id: str) -> EndpointInterceptData: - return cast(EndpointInterceptData, self.server.intercepts[request_id]) - - def request_context( - self, request_id: str, request: EndpointInterceptData - ) -> ModelRequestContext: - headers = request.get("headers") or {} - if not isinstance(headers, dict): - raise TypeError("Endpoint request headers must be a mapping.") - header_data: dict[str, str] = {} - for key, value in headers.items(): - if not isinstance(key, str) or not isinstance(value, str): - raise TypeError("Endpoint request headers must be strings.") - header_data[key.lower()] = value - return ModelRequestContext( - source="endpoint", - endpoint_request_id=request_id, - headers=header_data, - trajectory_visibility=self.trajectory_visibility(header_data), - ) - - def trajectory_visibility(self, headers: dict[str, str]) -> TrajectoryVisibility: - value = headers.get(VF_TRAJECTORY_VISIBILITY_HEADER) - if value is None: - return "append" - if not isinstance(value, str): - raise TypeError( - f"{VF_TRAJECTORY_VISIBILITY_HEADER} must be 'append' or 'hidden'." - ) - visibility = value.strip().lower() - if visibility not in {"append", "hidden"}: - raise ValueError( - f"{VF_TRAJECTORY_VISIBILITY_HEADER} must be 'append' or 'hidden'." - ) - return cast(TrajectoryVisibility, visibility) - - async def url_base(self) -> str: - if self.use_tunnel: - return await self.get_tunnel_url() - return f"http://127.0.0.1:{self.server.port}" - - async def get_tunnel_url(self) -> str: - from prime_tunnel import Tunnel - - async with self._tunnel_lock: - tunnel = self._tunnel - if tunnel is not None and not tunnel.is_running: - tunnel.sync_stop() - self._tunnel = None - - tunnel = self._tunnel - if tunnel is not None: - now = time.time() - if now - self._tunnel_last_checked > self.TUNNEL_CHECK_INTERVAL: - self._tunnel_last_checked = now - if not await tunnel.check_registered(): - tunnel.sync_stop() - self._tunnel = None - - if self._tunnel is None: - tunnel = cast(TunnelHandle, Tunnel(local_port=self.server.port)) - url = await tunnel.start() - self._tunnel = tunnel - self._tunnel_last_checked = time.time() - return str(url) - - tunnel = self._tunnel - if tunnel.url is None: - raise TunnelError("Tunnel started but URL is unavailable.") - return str(tunnel.url) - - async def check_tunnel(self) -> None: - tunnel = self._tunnel - if tunnel is not None and not tunnel.is_running: - raise TunnelError("Tunnel process died during rollout.") - - async def teardown(self) -> None: - async with self._tunnel_lock: - tunnel = self._tunnel - if tunnel is not None: - tunnel.sync_stop() - self._tunnel = None - await self.server.stop() - - -async def run_intercepted_program( - program: Callable[[Task, State], Awaitable[State | JsonData | None]], - endpoint: Endpoint, - runtime: Runtime, - task: Task, - state: State, -) -> State | JsonData | None: - async def call_tool(name: str, arguments: ToolParameters) -> object: - return await runtime.call_tool(name, task, state, **dict(arguments)) - - async def call_user(transcript: list[PromptMessage]) -> list[JsonData]: - return await runtime.user_messages(task, state, transcript=transcript) - - async def check_stop() -> JsonData: - done = await runtime.is_completed(task, state) - stop_condition = state.get("stop_condition") - return { - "done": done, - "stop_condition": stop_condition - if isinstance(stop_condition, str) or stop_condition is None - else str(stop_condition), - } - - model_tasks: set[asyncio.Task[Response]] = set() - - async def call_model(messages: list[PromptMessage], tools: object) -> JsonData: - # Sandbox sends canonical Messages; host resolves the client, tokenizes, - # and records the step. Tool defs come from the runtime. - del tools - prompt = normalize_messages( - cast(Messages, messages), field_name="vf.model.messages" - ) - request = asyncio.ensure_future( - runtime.submit_model_request( - prompt, - task, - state, - tool_defs=runtime.tool_defs(state), - context=ModelRequestContext(source="endpoint"), - ) - ) - model_tasks.add(request) - try: - response = await request - except Error as exc: - if isinstance(exc, OverlongPromptError): - state["prompt_too_long"] = True - state._set_truncated(True) - state._set_stop_condition("prompt_too_long", overwrite=True) - else: - state._set_error(exc) - state._set_stop_condition("has_error", overwrite=True) - raise - finally: - if request.done(): - model_tasks.discard(request) - completion = await parse_response_message(response) - return cast(JsonData, serializable(completion[0])) - - await endpoint.register_rollout( - state, - tool_handler=call_tool, - tool_defs=runtime.tool_defs(state), - user_handler=call_user, - stop_handler=check_stop, - model_handler=call_model, - ) - - async def execute_program() -> State | JsonData | None: - return await program(task, state) - - execution = asyncio.create_task(execute_program()) - rollout_key = str(state["endpoint_rollout_key"]) - queue = endpoint.rollout_queue(rollout_key) - pending: set[asyncio.Task[None]] = set() - try: - while True: - await raise_finished_forward_errors(pending) - if execution.done(): - await raise_execution_error(execution) - if not queue.empty(): - request_id = queue.get_nowait() - pending.add( - asyncio.create_task( - forward_request(endpoint, runtime, task, state, request_id) - ) - ) - continue - if not pending: - break - await asyncio.wait( - pending, - timeout=1.0, - return_when=asyncio.FIRST_COMPLETED, - ) - await endpoint.check_tunnel() - continue - queue_task = asyncio.create_task(queue.get()) - wait_set = {queue_task, execution, *pending} - try: - done, _ = await asyncio.wait( - wait_set, - timeout=1.0, - return_when=asyncio.FIRST_COMPLETED, - ) - if queue_task in done: - request_id = queue_task.result() - pending.add( - asyncio.create_task( - forward_request(endpoint, runtime, task, state, request_id) - ) - ) - continue - if execution in done: - continue - if pending.intersection(done): - continue - await endpoint.check_tunnel() - finally: - if not queue_task.done(): - queue_task.cancel() - await asyncio.gather(queue_task, return_exceptions=True) - if execution.done() and queue.empty() and not pending: - break - await raise_finished_forward_errors(pending) - return await execution - finally: - if not execution.done(): - execution.cancel() - await asyncio.gather(execution, return_exceptions=True) - await cancel_forwarders(pending) - await cancel_forwarders(cast("set[asyncio.Task[None]]", model_tasks)) - endpoint.unregister_rollout(rollout_key) - - -async def raise_finished_forward_errors(pending: set[asyncio.Task[None]]) -> None: - finished = {task for task in pending if task.done()} - for task in finished: - pending.remove(task) - await task - - -async def cancel_forwarders(pending: set[asyncio.Task[None]]) -> None: - for task in pending: - if not task.done(): - task.cancel() - if pending: - await asyncio.gather(*pending, return_exceptions=True) - - -async def raise_execution_error( - execution: asyncio.Task[State | JsonData | None], -) -> None: - if execution.cancelled(): - await execution - error = execution.exception() - if error is not None: - raise error - - -async def forward_request( - endpoint: Endpoint, - runtime: Runtime, - task: Task, - state: State, - request_id: str, -) -> None: - request = endpoint.get_request(request_id) - prompt = normalize_endpoint_prompt(request) - tool_defs = normalize_endpoint_tools( - request.get("tools"), str(request.get("protocol")) - ) - response = None - error: BaseException | None = None - try: - response = await runtime.submit_model_request( - prompt, - task, - state, - tool_defs=tool_defs, - context=endpoint.request_context(request_id, request), - ) - except BaseException as e: - error = e - if isinstance(e, Error): - state._set_error(e) - raise - finally: - if bool(request.get("stream")): - if request.get("protocol") != "openai_chat_completions": - raise NotImplementedError( - "Streaming interception is currently supported for OpenAI Chat Completions." - ) - await synthesize_stream(request, response, error) - else: - deliver_response(request, response, error) - - -def normalize_endpoint_prompt(request: EndpointInterceptData) -> Messages: - protocol = request.get("protocol") - if protocol == "anthropic_messages": - return normalize_anthropic_messages(request) - if protocol == "openai_responses": - return normalize_openai_responses_input(request.get("input")) - if protocol == "openai_completions": - return normalize_endpoint_messages(request.get("prompt")) - return normalize_endpoint_messages(request.get("messages")) - - -def normalize_endpoint_messages(messages: object) -> Messages: - if isinstance(messages, str): - return normalize_messages(messages, field_name="endpoint.messages") - if isinstance(messages, list): - return normalize_messages( - cast(Messages, messages), field_name="endpoint.messages" - ) - raise TypeError("Endpoint messages must be vf.Messages or str.") - - -def normalize_anthropic_messages(request: EndpointInterceptData) -> Messages: - messages: Messages = [] - system = request.get("system") - if isinstance(system, str) and system: - messages.append(SystemMessage(content=system)) - raw_messages = request.get("messages") - if not isinstance(raw_messages, list): - raise TypeError("Anthropic endpoint messages must be a list.") - for raw_message in raw_messages: - if not isinstance(raw_message, dict): - raise TypeError("Anthropic endpoint message entries must be dicts.") - raw_message = cast(JsonData, raw_message) - role = raw_message.get("role") - content = raw_message.get("content") - if role == "user": - messages.extend(normalize_anthropic_user_message(content)) - elif role == "assistant": - messages.append(normalize_anthropic_assistant_message(content)) - else: - raise ValueError(f"Unsupported Anthropic message role: {role!r}") - return messages - - -def normalize_anthropic_user_message(content: object) -> Messages: - if isinstance(content, str): - return [UserMessage(content=content)] - if not isinstance(content, list): - return [UserMessage(content=str(content))] - messages: Messages = [] - text_parts: list[str] = [] - for block in content: - if not isinstance(block, dict): - continue - block = cast(JsonData, block) - block_type = block.get("type") - if block_type == "text" and isinstance(block.get("text"), str): - text_parts.append(str(block["text"])) - elif block_type == "tool_result": - tool_use_id = block.get("tool_use_id") - if not isinstance(tool_use_id, str): - continue - messages.append( - ToolMessage( - tool_call_id=tool_use_id, - content=anthropic_tool_result_content(block.get("content")), - ) - ) - if text_parts: - messages.insert(0, UserMessage(content="\n".join(text_parts))) - return messages - - -def normalize_anthropic_assistant_message(content: object) -> AssistantMessage: - if isinstance(content, str): - return AssistantMessage(content=content) - if not isinstance(content, list): - return AssistantMessage(content=str(content)) - text_parts: list[str] = [] - tool_calls: list[ToolCall] = [] - for block in content: - if not isinstance(block, dict): - continue - block = cast(JsonData, block) - block_type = block.get("type") - if block_type == "text" and isinstance(block.get("text"), str): - text_parts.append(str(block["text"])) - elif block_type == "tool_use": - tool_id = block.get("id") - name = block.get("name") - if isinstance(tool_id, str) and isinstance(name, str): - tool_calls.append( - ToolCall( - id=tool_id, - name=name, - arguments=json.dumps(block.get("input") or {}), - ) - ) - return AssistantMessage( - content="\n".join(text_parts) if text_parts else None, - tool_calls=tool_calls or None, - ) - - -def anthropic_block_content_text(content: object) -> str: - if isinstance(content, str): - return content - if isinstance(content, list): - text_parts: list[str] = [] - for block in content: - if not isinstance(block, dict): - continue - block = cast(JsonData, block) - text = block.get("text") - if isinstance(text, str): - text_parts.append(text) - return "\n".join(text_parts) - return str(content) - - -def normalize_openai_responses_input(raw_input: object) -> Messages: - if isinstance(raw_input, str): - return [UserMessage(content=raw_input)] - if not isinstance(raw_input, list): - raise TypeError("OpenAI Responses input must be a string or list.") - messages: Messages = [] - for item in raw_input: - if not isinstance(item, dict): - raise TypeError("OpenAI Responses input entries must be dicts.") - item = cast(JsonData, item) - item_type = item.get("type") - if item_type == "function_call": - call_id = item.get("call_id") or item.get("id") - name = item.get("name") - arguments = item.get("arguments") - if ( - isinstance(call_id, str) - and isinstance(name, str) - and isinstance(arguments, str) - ): - messages.append( - AssistantMessage( - tool_calls=[ - ToolCall(id=call_id, name=name, arguments=arguments) - ] - ) - ) - continue - if item_type == "function_call_output": - call_id = item.get("call_id") - if isinstance(call_id, str): - messages.append( - ToolMessage( - tool_call_id=call_id, - content=responses_tool_output_content(item.get("output")), - ) - ) - continue - role = item.get("role") - content = responses_content_text(item.get("content")) - if role in {"system", "developer"}: - messages.append(SystemMessage(content=content)) - elif role == "assistant": - messages.append(AssistantMessage(content=content)) - else: - messages.append(UserMessage(content=content)) - return messages - - -def responses_content_text(content: object) -> str: - if isinstance(content, str): - return content - if isinstance(content, list): - text_parts: list[str] = [] - for part in content: - if isinstance(part, dict): - part = cast(JsonData, part) - text = part.get("text") - if isinstance(text, str): - text_parts.append(text) - return "\n".join(text_parts) - return "" if content is None else str(content) - - -def responses_tool_output_content(output: object) -> MessageContent: - """Responses function_call_output -> internal tool content, keeping images - (input_image -> image_url); text-only falls back to a string.""" - - if not isinstance(output, list): - return responses_content_text(output) - parts: list[ContentPart] = [] - has_image = False - for item in output: - if not isinstance(item, dict): - continue - item = cast(JsonData, item) - if item.get("type") == "input_image": - url = item.get("image_url") - if isinstance(url, str) and url: - parts.append({"type": "image_url", "image_url": {"url": url}}) - has_image = True - else: - text = item.get("text") - if isinstance(text, str): - parts.append({"type": "text", "text": text}) - return parts if has_image else responses_content_text(output) - - -def anthropic_tool_result_content(content: object) -> MessageContent: - """Anthropic tool_result content -> internal tool content, keeping images - (image block -> image_url); text-only falls back to a string.""" - - if not isinstance(content, list): - return anthropic_block_content_text(content) - parts: list[ContentPart] = [] - has_image = False - for block in content: - if not isinstance(block, dict): - continue - block = cast(JsonData, block) - if block.get("type") == "image": - source = block.get("source") - url = "" - if isinstance(source, dict): - if source.get("type") == "base64": - media_type = str(source.get("media_type") or "image/png") - data = str(source.get("data") or "") - url = f"data:{media_type};base64,{data}" - elif source.get("type") == "url": - url = str(source.get("url") or "") - if url: - parts.append({"type": "image_url", "image_url": {"url": url}}) - has_image = True - else: - text = block.get("text") - if isinstance(text, str): - parts.append({"type": "text", "text": text}) - return parts if has_image else anthropic_block_content_text(content) - - -def normalize_endpoint_tools(tools: object, protocol: str) -> list[Tool] | None: - if tools is None: - return None - if not isinstance(tools, list): - raise TypeError("Endpoint tools must be a list.") - normalized: list[Tool] = [] - for raw_tool in tools: - if isinstance(raw_tool, Tool): - normalized.append(raw_tool) - continue - if not isinstance(raw_tool, dict): - raise TypeError("Endpoint tool definitions must be dicts.") - raw_tool_data = cast(ToolParameters, raw_tool) - if protocol == "anthropic_messages": - normalized.append( - Tool( - name=str(raw_tool_data.get("name", "")), - description=str(raw_tool_data.get("description", "")), - parameters=endpoint_tool_parameters( - raw_tool_data.get("input_schema") - ), - ) - ) - continue - if protocol == "openai_responses": - normalized.append( - Tool( - name=str(raw_tool_data.get("name", "")), - description=str(raw_tool_data.get("description", "")), - parameters=endpoint_tool_parameters( - raw_tool_data.get("parameters") - ), - strict=cast(bool | None, raw_tool_data.get("strict")), - ) - ) - continue - function_payload = raw_tool_data.get("function") - if raw_tool_data.get("type") == "function" and isinstance( - function_payload, dict - ): - function_payload = cast(ToolParameters, function_payload) - normalized.append( - Tool( - name=str(function_payload.get("name", "")), - description=str(function_payload.get("description", "")), - parameters=endpoint_tool_parameters( - function_payload.get("parameters") - ), - strict=cast(bool | None, function_payload.get("strict")), - ) - ) - else: - normalized.append(Tool.model_validate(raw_tool_data)) - return normalized - - -def endpoint_tool_parameters(value: object) -> ToolParameters: - if value is None: - return {} - if not isinstance(value, dict): - raise TypeError("Endpoint tool parameters must be a mapping.") - return {str(key): item for key, item in value.items()} - - -def assistant_completion_from_messages( - prompt: list[JsonData], messages: list[JsonData] -) -> list[JsonData]: - return messages[len(prompt) :] diff --git a/verifiers/v1/utils/json_utils.py b/verifiers/v1/utils/json_utils.py index cb01cbb4dc..b9fb6d7ccf 100644 --- a/verifiers/v1/utils/json_utils.py +++ b/verifiers/v1/utils/json_utils.py @@ -1,10 +1,44 @@ import json -from typing import cast -from ..types import ConfigData +from pydantic import BaseModel -def json_args(value: str) -> ConfigData: +from ..types import JsonData, JsonValue + + +def json_args(value: str) -> JsonData: parsed = json.loads(value or "{}") - if not isinstance(parsed, dict): - raise ValueError("Tool call arguments must decode to a JSON object.") - return cast(ConfigData, parsed) + return json_data(parsed, context="Tool call arguments") + + +def jsonable(value: object) -> object: + if isinstance(value, BaseModel): + return jsonable(value.model_dump(mode="json", exclude_none=True)) + model_dump = getattr(value, "model_dump", None) + if callable(model_dump): + return jsonable(model_dump(mode="json", exclude_none=True)) + if isinstance(value, dict): + return {str(key): jsonable(item) for key, item in value.items()} + if isinstance(value, list | tuple): + return [jsonable(item) for item in value] + return value + + +def json_value(value: object, *, context: str = "Value") -> JsonValue: + resolved = jsonable(value) + if resolved is None or isinstance(resolved, str | int | float | bool): + return resolved + if isinstance(resolved, list): + return [json_value(item, context=context) for item in resolved] + if isinstance(resolved, dict): + return { + str(key): json_value(item, context=f"{context}.{key}") + for key, item in resolved.items() + } + raise TypeError(f"{context} must be JSON serializable.") + + +def json_data(value: object, *, context: str = "Value") -> JsonData: + resolved = json_value(value, context=context) + if not isinstance(resolved, dict): + raise TypeError(f"{context} must be a JSON object.") + return resolved diff --git a/verifiers/v1/utils/judge_utils.py b/verifiers/v1/utils/judge_utils.py deleted file mode 100644 index 55cea13a8b..0000000000 --- a/verifiers/v1/utils/judge_utils.py +++ /dev/null @@ -1,54 +0,0 @@ -import json -from typing import cast -from ..types import ConfigData - - -def parse_judge_json(text: str) -> ConfigData: - value = parsed_json_object(text) - if value is not None: - return value - start = text.find("{") - end = text.rfind("}") - if start >= 0 and end > start: - value = parsed_json_object(text[start : end + 1]) - if value is not None: - return value - return {"score": 0.0, "reason": "judge did not return JSON", "raw": text} - - -def parsed_json_object(text: str) -> ConfigData | None: - try: - value = json.loads(text) - except json.JSONDecodeError: - return None - if isinstance(value, dict): - return cast(ConfigData, value) - return None - - -def clamp_float(value: object) -> float: - if not isinstance(value, int | float | str) or isinstance(value, bool): - return 0.0 - try: - number = float(value) - except (TypeError, ValueError): - return 0.0 - return max(0.0, min(1.0, number)) - - -def truncate_command_record(record: object) -> object: - if not isinstance(record, dict): - return record - record = cast(ConfigData, record) - return { - **dict(record), - "command": truncate_text(str(record.get("command") or ""), limit=2_000), - "stdout": truncate_text(str(record.get("stdout") or "")), - "stderr": truncate_text(str(record.get("stderr") or "")), - } - - -def truncate_text(text: str, limit: int = 6_000) -> str: - if len(text) <= limit: - return text - return text[:limit] + "\n..." diff --git a/verifiers/v1/utils/lifecycle_utils.py b/verifiers/v1/utils/lifecycle_utils.py deleted file mode 100644 index 86103732bf..0000000000 --- a/verifiers/v1/utils/lifecycle_utils.py +++ /dev/null @@ -1,105 +0,0 @@ -import inspect -from collections.abc import Iterable -from typing import TYPE_CHECKING, Literal, cast - -from verifiers.utils.async_utils import maybe_call_with_named_args - -from .config_callable_utils import CallableKind -from ..state import State -from ..types import Handler - -if TYPE_CHECKING: - from ..harness import Harness - from ..taskset import Taskset - from ..toolset import Toolset - from ..user import User - -LifecycleStage = Literal["rollout", "group"] - - -def collect_handlers( - owners: Iterable["Taskset | Harness | Toolset | User | None"], - attr: str, - extra: Iterable[Handler] = (), - stage: LifecycleStage | None = None, -) -> list[Handler]: - handlers: list[Handler] = [] - for owner in owners: - if owner is None: - continue - for _, method in inspect.getmembers(owner, predicate=callable): - if handler_is_marked(method, cast(CallableKind, attr)): - handlers.append(cast(Handler, method)) - handlers.extend(extra) - if stage is not None: - handlers = [ - handler - for handler in handlers - if handler_stage(handler, cast(CallableKind, attr)) == stage - ] - return sort_handlers(unique_handlers(handlers), attr) - - -def validate_handler_args( - handlers: Iterable[Handler], - expected: set[str], - attr: str, - stage: LifecycleStage, -) -> None: - context = f"{stage} {attr}" - for handler in handlers: - signature = inspect.signature(handler) - for parameter in signature.parameters.values(): - if parameter.kind == parameter.POSITIONAL_ONLY: - raise TypeError( - f"{context} handler {handler!r} must use named parameters." - ) - if ( - parameter.kind == parameter.VAR_POSITIONAL - and parameter.name not in expected - ): - raise TypeError(f"{context} handler {handler!r} must not use *args.") - - -async def run_handlers(handlers: Iterable[Handler], **kwargs: object) -> None: - for handler in handlers: - await maybe_call_with_named_args(handler, **kwargs) - - -def unique_handlers( - handlers: Iterable[Handler], -) -> list[Handler]: - unique: list[Handler] = [] - seen: set[tuple[int, int]] = set() - for handler in handlers: - key = ( - id(getattr(handler, "__self__", None)), - id(getattr(handler, "__func__", handler)), - ) - if key in seen: - continue - seen.add(key) - unique.append(handler) - return unique - - -def sort_handlers(handlers: Iterable[Handler], attr: str) -> list[Handler]: - return sorted( - handlers, - key=lambda handler: ( - -int(getattr(handler, f"{attr}_priority", 0)), - str(getattr(handler, "__name__", "")), - ), - ) - - -async def state_done(state: State) -> bool: - return bool(state.get("done")) - - -def handler_is_marked(handler: object, kind: CallableKind) -> bool: - return getattr(handler, kind, False) is True - - -def handler_stage(handler: object, kind: CallableKind) -> LifecycleStage: - return cast(LifecycleStage, getattr(handler, f"{kind}_stage", "rollout")) diff --git a/verifiers/v1/utils/logging_utils.py b/verifiers/v1/utils/logging_utils.py deleted file mode 100644 index 1c92debaec..0000000000 --- a/verifiers/v1/utils/logging_utils.py +++ /dev/null @@ -1,61 +0,0 @@ -import logging -from collections import Counter - -from verifiers.utils.display_utils import format_numeric, format_timing_plain -from verifiers.utils.error_utils import ErrorChain, is_error_data -from verifiers.utils.logging_utils import truncate - -from ..state import State - -logger = logging.getLogger("verifiers.v1.rollout") - - -def log_rollout_start(state: State) -> None: - logger.info( - f"Started example_id={state.get('example_id')} " - f"| trajectory_id={state.get('trajectory_id')}" - ) - - -def log_rollout_finish(state: State) -> None: - tools = Counter( - call["name"] - for step in state.get("trajectory") or [] - for message in step.get("completion") or [] - for call in message.get("tool_calls") or [] - ) - timing = state.get("timing") or {} - metrics = state.get("metrics") or {} - - def duration(phase: str) -> float: - return (timing.get(phase) or {}).get("duration", 0.0) - - tool_summary = ", ".join(f"{name}: {n}" for name, n in tools.most_common()) - metric_summary = ", ".join( - f"{name}: {format_numeric(value)}" for name, value in metrics.items() - ) - parts = [ - f"Finished example_id={state.get('example_id')}", - f"trajectory_id={state.get('trajectory_id')}", - f"tools=[{tool_summary}]", - "timing=" - + format_timing_plain( - setup=duration("setup"), - generation=duration("generation"), - scoring=duration("scoring"), - overhead=timing.get("overhead", 0.0), - model=duration("model"), - env=duration("env"), - ), - f"stop={state.get('stop_condition')}", - f"reward={format_numeric(state.get('reward') or 0.0)}", - "metrics={" + metric_summary + "}", - ] - error = state.get("error") - if isinstance(error, BaseException): - parts.append(f"error={truncate(str(ErrorChain(error)), 200)}") - elif is_error_data(error): - parts.append(f"error={truncate(error['error_chain_str'], 200)}") - if state.get("is_truncated"): - parts.append("truncated=True") - logger.info(" | ".join(parts)) diff --git a/verifiers/v1/utils/mcp_proxy_utils.py b/verifiers/v1/utils/mcp_proxy_utils.py deleted file mode 100644 index 58c71e3937..0000000000 --- a/verifiers/v1/utils/mcp_proxy_utils.py +++ /dev/null @@ -1,223 +0,0 @@ -import json -from typing import Literal, cast - -from ..types import ConfigData -from .sandbox_python_utils import SANDBOX_PYTHON, python_package_list - -ProgramChannel = Literal["callable", "mcp"] - -MCP_PROXY_PATH = "/tmp/vf_mcp_tools.py" -MCP_PROXY_CONFIG_PATH = "/tmp/vf_mcp_tools.json" -MCP_PACKAGE = "mcp>=1.14.1" -REQUESTS_PACKAGE = "requests" - - -PROGRAM_CHANNELS = {"callable", "mcp"} -PROGRAM_CHANNEL_METADATA = {"priority"} - - -def validate_program_channels(value: object) -> tuple[ProgramChannel, ...]: - if value is None: - return () - if isinstance(value, str): - if value not in PROGRAM_CHANNELS: - raise ValueError("program.channels must be 'callable' or 'mcp'.") - return (cast(ProgramChannel, value),) - if isinstance(value, list): - result: list[ProgramChannel] = [] - for item in value: - for channel in validate_program_channels(item): - if channel in result: - raise ValueError( - f"program.channels defines {channel!r} more than once." - ) - result.append(channel) - return tuple(result) - if isinstance(value, dict): - if not all(isinstance(key, str) for key in value): - raise TypeError("program.channels mapping keys must be strings.") - spec = cast(ConfigData, value) - unknown = sorted(set(spec) - PROGRAM_CHANNELS - PROGRAM_CHANNEL_METADATA) - if unknown: - raise ValueError(f"program.channels has unknown channel: {unknown}.") - if "priority" in spec: - priority = spec["priority"] - if not isinstance(priority, int) or isinstance(priority, bool): - raise TypeError("program.channels priority must be an integer.") - result = [cast(ProgramChannel, key) for key in spec if key in PROGRAM_CHANNELS] - if not result: - raise ValueError("program.channels mapping must define a channel.") - return tuple(result) - raise TypeError("program.channels must be a string, mapping, or list.") - - -def proxy_program( - program: ConfigData, tool_base_url: str, tool_auth_var: str -) -> ConfigData: - files = dict(cast(ConfigData, program.get("files") or {})) - if MCP_PROXY_PATH in files and files[MCP_PROXY_PATH] != proxy_source(): - raise ValueError(f"program.files cannot override {MCP_PROXY_PATH}.") - config = { - "tool_base_url": tool_base_url.rstrip("/"), - "tool_auth_var": tool_auth_var, - } - config_json = json.dumps(config) - if MCP_PROXY_CONFIG_PATH in files and files[MCP_PROXY_CONFIG_PATH] != config_json: - raise ValueError(f"program.files cannot override {MCP_PROXY_CONFIG_PATH}.") - files[MCP_PROXY_PATH] = proxy_source() - files[MCP_PROXY_CONFIG_PATH] = config_json - return {**dict(program), "files": files} - - -def proxy_command() -> list[str]: - return [SANDBOX_PYTHON, MCP_PROXY_PATH, MCP_PROXY_CONFIG_PATH] - - -def proxy_sandbox(sandbox_config: ConfigData) -> ConfigData: - config = dict(sandbox_config) - packages = python_package_list(config.get("packages")) - if not any(str(package).startswith("mcp") for package in packages): - packages.append(MCP_PACKAGE) - if not any(str(package).startswith("requests") for package in packages): - packages.append(REQUESTS_PACKAGE) - config["packages"] = packages - return config - - -def proxy_source() -> str: - return r""" -import asyncio -import json -import os -import sys - -import requests - -from mcp.server import Server -from mcp.server.stdio import stdio_server -from mcp.types import CallToolResult, TextContent, Tool - -CONFIG = None - - -def config() -> dict: - global CONFIG - if CONFIG is None: - if len(sys.argv) != 2: - raise RuntimeError("MCP proxy requires a config path argument.") - with open(sys.argv[1]) as f: - CONFIG = json.load(f) - return CONFIG - - -def tool_base_url() -> str: - value = config().get("tool_base_url") - if not value: - raise RuntimeError("tool_base_url is required.") - return str(value).rstrip("/") - - -def auth_headers() -> dict[str, str]: - token_var = config().get("tool_auth_var") - if not token_var: - raise RuntimeError("tool_auth_var is required.") - token = os.environ.get(str(token_var), "") - headers = { - "Content-Type": "application/json", - "Accept": "application/json", - "User-Agent": "python-requests/2.32.3", - } - if token: - headers["Authorization"] = f"Bearer {token}" - return headers - - -def get_json(url: str, params: dict | None = None) -> dict: - try: - response = requests.get( - url, - params=params, - headers=auth_headers(), - timeout=300, - ) - response.raise_for_status() - except requests.HTTPError as exc: - raise RuntimeError(response.text) from exc - except requests.RequestException as exc: - raise RuntimeError(str(exc)) from exc - return response.json() - - -def post_json(url: str, payload: dict) -> dict: - try: - response = requests.post( - url, - json=payload, - headers=auth_headers(), - timeout=300, - ) - response.raise_for_status() - except requests.HTTPError as exc: - raise RuntimeError(response.text) from exc - except requests.RequestException as exc: - raise RuntimeError(str(exc)) from exc - return response.json() - - -def tool_text(value: object) -> str: - if isinstance(value, str): - return value - return json.dumps(value, ensure_ascii=False) - - -server = Server("verifiers-tools") - - -@server.list_tools() -async def list_tools() -> list[Tool]: - payload = await asyncio.to_thread(get_json, tool_base_url(), {"protocol": "vf"}) - tools = [] - for item in payload.get("tools") or []: - tools.append( - Tool( - name=str(item["name"]), - description=str(item.get("description") or ""), - inputSchema=item.get("parameters") or {"type": "object", "properties": {}}, - ) - ) - return tools - - -@server.call_tool(validate_input=False) -async def call_tool(name: str, arguments: dict) -> CallToolResult: - payload = await asyncio.to_thread( - post_json, - f"{tool_base_url()}/{name}", - {"arguments": arguments or {}}, - ) - if "error" in payload: - return CallToolResult( - content=[TextContent(type="text", text=str(payload["error"]))], - isError=True, - ) - result = payload.get("result") - structured = result if isinstance(result, dict) else None - return CallToolResult( - content=[TextContent(type="text", text=tool_text(result))], - structuredContent=structured, - isError=False, - ) - - -async def main() -> None: - async with stdio_server() as (read_stream, write_stream): - await server.run( - read_stream, - write_stream, - server.create_initialization_options(), - ) - - -if __name__ == "__main__": - asyncio.run(main()) -""" diff --git a/verifiers/v1/utils/mcp_utils.py b/verifiers/v1/utils/mcp_utils.py deleted file mode 100644 index b9db306dec..0000000000 --- a/verifiers/v1/utils/mcp_utils.py +++ /dev/null @@ -1,150 +0,0 @@ -import asyncio -from contextlib import AsyncExitStack -from typing import cast - -from verifiers.errors import ToolError -from verifiers.types import Tool - -from ..toolset import MCPTool -from ..types import RuntimeData - - -class MCPToolHandle: - def __init__(self, session: "MCPToolSession", tool_def: Tool): - self.session = session - self.name = tool_def.name - self.tool_def = tool_def - - async def __call__(self, **kwargs: object) -> object: - result = await self.session.call_tool(self.name, dict(kwargs)) - return mcp_result_value(result) - - -class MCPToolSession: - def __init__(self, spec: MCPTool): - self.spec = spec - self.handles: list[MCPToolHandle] = [] - self._queue: asyncio.Queue[tuple[str, str, RuntimeData, asyncio.Future]] = ( - asyncio.Queue() - ) - self._ready: asyncio.Future[list[MCPToolHandle]] | None = None - self._task: asyncio.Task[None] | None = None - - async def __aenter__(self) -> "MCPToolSession": - loop = asyncio.get_running_loop() - self._ready = loop.create_future() - self._task = loop.create_task(self._run()) - self.handles = await self._ready - return self - - async def __aexit__(self, exc_type, exc, tb) -> None: - await self.close() - - async def call_tool(self, name: str, arguments: RuntimeData) -> object: - loop = asyncio.get_running_loop() - future = loop.create_future() - await self._queue.put(("call", name, arguments, future)) - return await future - - async def close(self) -> None: - task = self._task - if task is None or task.done(): - return - loop = asyncio.get_running_loop() - future = loop.create_future() - await self._queue.put(("close", "", {}, future)) - await future - await task - - async def _run(self) -> None: - from mcp import ClientSession - from mcp.client.stdio import StdioServerParameters, stdio_client - - server = StdioServerParameters( - command=self.spec.command, - args=list(self.spec.args), - env=dict(self.spec.env) if self.spec.env is not None else None, - cwd=self.spec.cwd, - ) - ready = self._ready - if ready is None: - raise RuntimeError("MCPToolSession started without a ready future.") - try: - async with stdio_client(server) as (read_stream, write_stream): - async with ClientSession(read_stream, write_stream) as session: - await session.initialize() - tools_result = await session.list_tools() - self.handles = [ - MCPToolHandle(self, mcp_tool_def(tool)) - for tool in tools_result.tools - ] - ready.set_result(self.handles) - while True: - action, name, arguments, future = await self._queue.get() - try: - if action == "close": - cast(asyncio.Future[None], future).set_result(None) - return - if action != "call": - raise RuntimeError(f"Unknown MCP action: {action}") - result = await session.call_tool(name, arguments) - cast(asyncio.Future, future).set_result(result) - except BaseException as exc: - cast(asyncio.Future, future).set_exception(exc) - except BaseException as exc: - if not ready.done(): - ready.set_exception(exc) - raise - - -async def connect_mcp_tool( - spec: MCPTool, exit_stack: AsyncExitStack -) -> list[MCPToolHandle]: - session = await exit_stack.enter_async_context(MCPToolSession(spec)) - return session.handles - - -def mcp_tool_def(tool: object) -> Tool: - schema = getattr(tool, "inputSchema", None) or getattr(tool, "input_schema", None) - model_dump = getattr(tool, "model_dump", None) - if schema is None and callable(model_dump): - dumped = model_dump() - schema = dumped.get("inputSchema") or dumped.get("input_schema") - if not isinstance(schema, dict): - schema = {"type": "object", "properties": {}} - name = getattr(tool, "name", None) - if not isinstance(name, str) or not name: - raise TypeError("MCP tools require a name.") - return Tool( - name=name, - description=str(getattr(tool, "description", "") or ""), - parameters={str(key): item for key, item in schema.items()}, - strict=None, - ) - - -def mcp_result_value(result: object) -> object: - content = getattr(result, "content", []) - if bool(getattr(result, "isError", False)): - raise ToolError(str(mcp_content_value(content))) - return mcp_content_value(content) - - -def mcp_content_value(content: object) -> object: - if not isinstance(content, list): - return serializable_content(content) - values = [serializable_content(item) for item in content] - if len(values) == 1: - return values[0] - return values - - -def serializable_content(item: object) -> object: - item_type = getattr(item, "type", None) - text = getattr(item, "text", None) - if item_type == "text" and isinstance(text, str): - return text - model_dump = getattr(item, "model_dump", None) - if callable(model_dump): - return model_dump(exclude_none=True) - return item diff --git a/verifiers/v1/utils/object_utils.py b/verifiers/v1/utils/object_utils.py deleted file mode 100644 index e42bc980e5..0000000000 --- a/verifiers/v1/utils/object_utils.py +++ /dev/null @@ -1,61 +0,0 @@ -import inspect -from collections.abc import Awaitable -from typing import cast - -from verifiers.utils.async_utils import maybe_call_with_named_args - -from ..types import ObjectFactory, RuntimeObject - - -async def close_object(obj: RuntimeObject) -> None: - for name in ("aclose", "close", "delete", "teardown"): - fn = getattr(obj, name, None) - if callable(fn): - await maybe_call_with_named_args(fn) - return - - -async def resolve_object_factory( - spec: ObjectFactory, context: str, kwargs: dict[str, RuntimeObject] | None = None -) -> RuntimeObject: - if not callable(spec): - raise TypeError(f"{context} must be an import ref or factory function.") - if not (inspect.isfunction(spec) or inspect.isclass(spec)): - raise TypeError(f"{context} must be a factory function or class.") - validate_object_factory(spec, context, kwargs or {}) - value = cast(ObjectFactory, spec)(**(kwargs or {})) - if inspect.isawaitable(value): - return await cast(Awaitable[RuntimeObject], value) - return value - - -def validate_object_loader_spec(spec: RuntimeObject, context: str) -> None: - if isinstance(spec, str): - return - if not callable(spec): - raise TypeError(f"{context} must be an import ref or factory function.") - if not (inspect.isfunction(spec) or inspect.isclass(spec)): - raise TypeError(f"{context} must be a factory function or class.") - validate_object_factory_spec(spec, context) - - -def validate_object_factory_spec(spec: RuntimeObject, context: str) -> None: - if not inspect.isclass(spec): - name = getattr(spec, "__name__", "") - if name == "": - raise TypeError(f"{context} must be a named factory function.") - try: - inspect.signature(cast(ObjectFactory, spec)) - except (TypeError, ValueError) as exc: - raise TypeError(f"{context} factory signature cannot be inspected.") from exc - - -def validate_object_factory( - spec: RuntimeObject, context: str, kwargs: dict[str, RuntimeObject] -) -> None: - validate_object_factory_spec(spec, context) - signature = inspect.signature(cast(ObjectFactory, spec)) - try: - signature.bind(**kwargs) - except TypeError as exc: - raise TypeError(f"{context} has unbound factory arguments.") from exc diff --git a/verifiers/v1/utils/program_utils.py b/verifiers/v1/utils/program_utils.py deleted file mode 100644 index 9c08a26683..0000000000 --- a/verifiers/v1/utils/program_utils.py +++ /dev/null @@ -1,555 +0,0 @@ -import asyncio -import os -import shlex -from typing import TypeAlias, cast - -from verifiers.errors import InfraError -from verifiers.utils.async_utils import maybe_call_with_named_args - -from .config_utils import resolve_config_object, string_mapping -from .binding_utils import ( - BindingsConfig, - binding_key_parts, - function_name, - read_path, - validate_binding_source, - validate_bound_arg, -) -from ..runtime import Runtime -from ..state import State -from ..task import Task -from .mcp_proxy_utils import validate_program_channels -from ..artifact import ArtifactsConfig -from ..program import ProgramChannel, ProgramValue -from ..sandbox import SandboxConfig -from ..types import ConfigData, Handler, RuntimeData - -ProgramMappingInput: TypeAlias = dict[str, ProgramValue] | None -ProgramListInput: TypeAlias = ProgramValue | None - -PROGRAM_KIND_KEYS = {"base", "fn", "command"} -PROGRAM_OPTION_KEYS = { - "sandbox", - "files", - "dirs", - "setup", - "setup_timeout", - "bindings", - "env", - "artifacts", - "channels", -} -PROGRAM_KEYS = PROGRAM_KIND_KEYS | PROGRAM_OPTION_KEYS | {"args"} -SANDBOX_ONLY_PROGRAM_KEYS = {"files", "dirs", "setup", "setup_timeout", "artifacts"} -TASK_PROGRAM_KEYS = { - "files", - "dirs", - "setup", - "bindings", - "env", - "artifacts", - "args", -} - - -async def run_local_command( - program: ConfigData, task: Task, state: State, runtime: Runtime -) -> State: - if "mcp" in program_channels(program): - raise ValueError("program.channels='mcp' requires sandbox command placement.") - validate_program_bindings(program) - argv = await command_argv(program, task, state, runtime) - env = await command_env(program, task, state, runtime, include_base=True) - proc = await asyncio.create_subprocess_exec( - *argv, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - env=env, - ) - stdout, stderr = await proc.communicate() - state["command"] = { - "argv": argv, - "returncode": proc.returncode, - "stdout": stdout.decode(errors="replace"), - "stderr": stderr.decode(errors="replace"), - } - state["completion"] = [ - {"role": "assistant", "content": state["command"]["stdout"].strip()} - ] - if proc.returncode: - raise InfraError( - f"Command exited with {proc.returncode}: {state['command']['stderr']}" - ) - state._set_stop_condition("command_completed") - return state - - -async def command_argv( - program: ConfigData, task: Task, state: State, runtime: Runtime -) -> list[str]: - command = program.get("command") - if isinstance(command, str): - argv = shlex.split(command) - elif isinstance(command, list): - argv = [ - str( - await resolve_program_value( - cast(ProgramValue, part), task, state, runtime, program - ) - ) - for part in command - ] - else: - raise TypeError("program.command must be a string or list.") - args = program.get("args", []) - if not isinstance(args, list): - raise TypeError("program.args must be a list.") - for arg in args: - argv.append( - str( - await resolve_program_value( - cast(ProgramValue, arg), task, state, runtime, program - ) - ) - ) - if not argv: - raise ValueError("program.command cannot be empty.") - return argv - - -async def command_env( - program: ConfigData, - task: Task, - state: State, - runtime: Runtime, - include_base: bool, -) -> dict[str, str]: - env = dict(os.environ) if include_base else {} - endpoint_base_url = state.get("endpoint_base_url") - if isinstance(endpoint_base_url, str): - harness = runtime.harness - if harness is None: - raise RuntimeError("Runtime has no active model endpoint.") - api_key = str(harness.endpoint.secret or "intercepted") - endpoint_root_url = state.get("endpoint_root_url") - env["OPENAI_BASE_URL"] = endpoint_base_url - env["OPENAI_API_KEY"] = api_key - api_key_var = state.get("endpoint_api_key_var") - if isinstance(api_key_var, str): - env[api_key_var] = api_key - if isinstance(endpoint_root_url, str): - env["ANTHROPIC_BASE_URL"] = endpoint_root_url - env["ANTHROPIC_API_KEY"] = api_key - raw_env = program.get("env", {}) - if not isinstance(raw_env, dict): - raise TypeError("program.env must be a mapping.") - for key, value in raw_env.items(): - if not isinstance(key, str): - raise TypeError("program.env keys must be strings.") - env[key] = str( - await resolve_program_value( - cast(ProgramValue, value), task, state, runtime, program - ) - ) - return env - - -async def resolve_program_value( - value: ProgramValue, - task: Task, - state: State, - runtime: Runtime, - program: ConfigData | None = None, -) -> object: - callable_spec = program_value_callable(value) - if callable_spec is not None: - fn, configured_kwargs = callable_spec - kwargs = await program_binding_kwargs(fn, program, task, state, runtime) - for key, item in configured_kwargs.items(): - kwargs[key] = await resolve_program_value( - cast(ProgramValue, item), task, state, runtime, program - ) - return await maybe_call_with_named_args( - fn, task=task, state=state, runtime=runtime, **kwargs - ) - if isinstance(value, str): - root, separator, tail = value.partition(".") - if separator and root == "task": - return read_path(task, tail) - if separator and root == "state": - return read_path(state, tail) - if separator and root == "runtime": - return read_path(state.runtime_state(), tail) - if isinstance(value, dict): - if len(value) != 1: - raise ValueError("Program value mappings must have exactly one root.") - root, path = next(iter(value.items())) - if root == "task": - return read_path(task, str(path)) - if root == "state": - return read_path(state, str(path)) - if root == "runtime": - return read_path(state.runtime_state(), str(path)) - raise ValueError(f"Unknown program value root {root!r}.") - return value - - -def program_value_callable(value: ProgramValue) -> tuple[Handler, ConfigData] | None: - if isinstance(value, dict) and "fn" in value: - spec = cast(ConfigData, value) - validate_program_callable_source(spec) - fn = resolve_config_object(spec["fn"]) - if not callable(fn): - raise TypeError("Program callable value requires callable fn.") - kwargs = {key: item for key, item in spec.items() if key != "fn"} - return cast(Handler, fn), kwargs - return None - - -def validate_program_callable_source(source: ConfigData) -> None: - fn = source.get("fn") - if not isinstance(fn, str): - raise TypeError("Program callable value fn must be an import ref string.") - for key in source: - if not isinstance(key, str) or not key: - raise TypeError("Program callable value keys must be non-empty strings.") - - -async def program_binding_kwargs( - fn: Handler, - program: ConfigData | None, - task: Task, - state: State, - runtime: Runtime, -) -> RuntimeData: - if program is None: - return {} - raw_bindings = BindingsConfig.model_validate(program.get("bindings") or {}).entries( - "program.bindings", allow_objects=False - ) - if not raw_bindings: - return {} - name = function_name(fn) - kwargs: RuntimeData = {} - for binding_key, source in raw_bindings.items(): - target_name, arg_name = binding_key_parts(binding_key) - if target_name != name: - continue - validate_bound_arg(fn, arg_name, f"Program binding {binding_key!r}") - validate_binding_source( - source, f"Program binding {binding_key!r}", allow_objects=False - ) - if arg_name in kwargs: - raise ValueError(f"Program binding arg {arg_name!r} is defined twice.") - kwargs[arg_name] = await runtime.resolve_binding(source, task, state) - return kwargs - - -def validate_program_bindings(program: ConfigData) -> None: - raw_bindings = BindingsConfig.model_validate(program.get("bindings") or {}).entries( - "program.bindings", allow_objects=False - ) - if not raw_bindings: - return - targets = program_binding_targets(program) - for binding_key, source in raw_bindings.items(): - target_name, arg_name = binding_key_parts(binding_key) - fn = targets.get(target_name) - if fn is None: - if target_name in program_setup_callable_names(program): - raise ValueError( - "program.setup callables cannot use program.bindings; move " - "bound runtime setup under program.channels.." - ) - raise ValueError( - f"Program binding {binding_key!r} does not match a callable " - "owned by the same program." - ) - validate_bound_arg(fn, arg_name, f"Program binding {binding_key!r}") - validate_binding_source( - source, f"Program binding {binding_key!r}", allow_objects=False - ) - - -def program_binding_targets( - program: ConfigData, -) -> dict[str, Handler]: - targets: dict[str, Handler] = {} - - def add(value: object) -> None: - callable_spec = program_value_callable(cast(ProgramValue, value)) - if callable_spec is None: - return - fn, _ = callable_spec - name = function_name(fn) - existing = targets.get(name) - if existing is not None and existing is not fn: - raise ValueError(f"Program binding target {name!r} is defined twice.") - targets[name] = fn - - def add_items(value: object) -> None: - if isinstance(value, list): - for item in value: - add(item) - elif value is not None and not isinstance(value, str): - add(value) - - command = program.get("command") - if isinstance(command, list): - for item in command: - add(item) - add_items(program.get("args")) - for _, item, _ in program_channel_setup(program): - add(item) - for key in ("files", "dirs", "env"): - value = program.get(key) - if isinstance(value, dict): - for item in value.values(): - add(item) - return targets - - -def program_setup_callable_names(program: ConfigData) -> set[str]: - names: set[str] = set() - setup = program.get("setup") - items = setup if isinstance(setup, list) else [setup] - for item in items: - callable_spec = program_value_callable(cast(ProgramValue, item)) - if callable_spec is not None: - fn, _ = callable_spec - names.add(function_name(fn)) - return names - - -def float_config(config: ConfigData, key: str, default: float) -> float: - value = config.get(key) - if value is None: - return default - if isinstance(value, bool) or not isinstance(value, int | float | str): - raise TypeError(f"{key} must be numeric.") - return float(value) - - -def int_config(config: ConfigData, key: str, default: int) -> int: - value = config.get(key) - if value is None: - return default - if isinstance(value, bool) or not isinstance(value, int | float | str): - raise TypeError(f"{key} must be numeric.") - return int(value) - - -def program_channels(program: ConfigData) -> tuple[ProgramChannel, ...]: - return validate_program_channels(program.get("channels")) - - -def program_kind(program: ConfigData) -> str: - base = program.get("base", False) - if not isinstance(base, bool): - raise TypeError("program.base must be a boolean.") - kinds = [] - if base: - kinds.append("base") - if "fn" in program: - kinds.append("fn") - if "command" in program: - kinds.append("command") - if not kinds and any(key in program for key in PROGRAM_OPTION_KEYS): - if "sandbox" not in program or program.get("sandbox") is False: - raise ValueError("option-only program mappings require sandbox placement.") - kinds.append("base") - if len(kinds) != 1: - raise ValueError( - "program mapping must specify exactly one of base=true, fn, or command." - ) - return kinds[0] - - -def validate_program_options( - program: ConfigData, - kind: str, - sandbox_config: SandboxConfig | None, -) -> None: - unknown = sorted(set(program) - PROGRAM_KEYS) - if unknown: - raise ValueError(f"Unknown program keys: {unknown}.") - validate_program_bindings(program) - if sandbox_config is None: - sandbox_only = sorted(set(program) & SANDBOX_ONLY_PROGRAM_KEYS) - if sandbox_only: - raise ValueError(f"Program keys {sandbox_only} require sandbox placement.") - channels = set(program_channels(program)) - if "mcp" in channels: - if kind != "command": - raise ValueError( - "program.channels='mcp' is only supported for command programs." - ) - if sandbox_config is None: - raise ValueError("program.channels='mcp' requires program.sandbox.") - if "callable" in channels and kind == "command": - raise ValueError( - "program.channels='callable' is only supported for base and fn programs." - ) - if kind == "base" and sandbox_config is None: - inert = sorted(set(program) & (PROGRAM_OPTION_KEYS - {"sandbox"})) - if inert: - raise ValueError(f"Base program keys {inert} require sandbox placement.") - - -def validate_program_sandbox_scope(sandbox_config: SandboxConfig) -> None: - if sandbox_config.scope not in {"rollout", "group", "global"}: - raise ValueError("program sandbox scope must be rollout, group, or global.") - - -def merge_task_program(program: ConfigData, task: Task, *, kind: str) -> ConfigData: - task_program = task.get("program") - if task_program is None: - return program - if not isinstance(task_program, dict): - raise TypeError("task.program must be a mapping.") - task_program = cast(ConfigData, task_program) - unknown = sorted(set(task_program) - TASK_PROGRAM_KEYS) - if unknown: - raise ValueError( - "task.program can only define files, dirs, setup, bindings, env, " - f"artifacts, and args; got {unknown}." - ) - if kind != "command" and "args" in task_program: - raise ValueError("task.program.args is only supported for command programs.") - merged = dict(program) - for key in ("files", "dirs", "env", "artifacts"): - merged[key] = merge_program_mapping_option( - cast(ProgramMappingInput, program.get(key)), - cast(ProgramMappingInput, task_program.get(key)), - key, - ) - merged["bindings"] = merge_program_bindings( - cast(ProgramMappingInput, program.get("bindings")), - cast(ProgramMappingInput, task_program.get("bindings")), - ) - merged["setup"] = [ - *program_list_items( - cast(ProgramListInput, program.get("setup")), "program.setup" - ), - *program_list_items( - cast(ProgramListInput, task_program.get("setup")), "task.program.setup" - ), - ] - if kind == "command": - merged["args"] = [ - *program_list_items( - cast(ProgramListInput, program.get("args")), "program.args" - ), - *program_list_items( - cast(ProgramListInput, task_program.get("args")), "task.program.args" - ), - ] - return merged - - -def merge_task_sandbox(sandbox_config: SandboxConfig, task: Task) -> SandboxConfig: - task_sandbox = task.sandbox_config() - if task_sandbox is None: - validate_program_sandbox_scope(sandbox_config) - return sandbox_config - config = SandboxConfig.model_validate( - {**sandbox_config.data(), **task_sandbox.data(fill_defaults=False)} - ) - validate_program_sandbox_scope(config) - return config - - -def merge_program_mapping_option( - program_value: ProgramMappingInput, task_value: ProgramMappingInput, key: str -) -> ConfigData: - if key == "artifacts": - program_mapping = ArtifactsConfig.model_validate(program_value or {}).data( - "program.artifacts" - ) - task_mapping = ArtifactsConfig.model_validate(task_value or {}).data( - "task.program.artifacts" - ) - else: - program_mapping = program_option_mapping(program_value, f"program.{key}") - task_mapping = program_option_mapping(task_value, f"task.program.{key}") - duplicate = sorted(set(program_mapping) & set(task_mapping)) - if duplicate: - raise ValueError( - f"program.{key} and task.program.{key} define the same keys: {duplicate}." - ) - return {**program_mapping, **task_mapping} - - -def merge_program_bindings( - program_value: ProgramMappingInput, task_value: ProgramMappingInput -) -> ConfigData: - program_bindings = BindingsConfig.model_validate(program_value or {}).entries( - "program.bindings", allow_objects=False - ) - task_bindings = BindingsConfig.model_validate(task_value or {}).entries( - "task.program.bindings", allow_objects=False - ) - duplicate = sorted(set(program_bindings) & set(task_bindings)) - if duplicate: - raise ValueError( - "program.bindings and task.program.bindings define the same keys: " - f"{duplicate}." - ) - return string_mapping({**program_bindings, **task_bindings}) - - -def program_option_mapping( - value: ProgramMappingInput, field_name: str -) -> dict[str, ProgramValue]: - if value is None: - return {} - if not isinstance(value, dict): - raise TypeError(f"{field_name} must be a mapping.") - try: - return cast(dict[str, ProgramValue], string_mapping(value)) - except TypeError as exc: - raise TypeError(f"{field_name} keys must be strings.") from exc - - -def program_list_items(value: ProgramListInput, field_name: str) -> list[ProgramValue]: - if value is None: - return [] - if isinstance(value, str): - return [value] - if isinstance(value, tuple): - raise TypeError(f"{field_name} must be a string, mapping, or list.") - if not isinstance(value, list): - return [cast(ProgramValue, value)] - return [cast(ProgramValue, item) for item in value] - - -def program_channel_setup( - program: ConfigData, -) -> list[tuple[ProgramChannel, ProgramValue, int]]: - channels = program.get("channels") - if channels is None or isinstance(channels, str): - return [] - if isinstance(channels, list): - result: list[tuple[ProgramChannel, ProgramValue, int]] = [] - for item in channels: - result.extend(program_channel_setup({"channels": item})) - return result - if not isinstance(channels, dict): - validate_program_channels(channels) - return [] - channel_names = validate_program_channels(channels) - channels_map = cast(ConfigData, channels) - priority = cast(int, channels_map.get("priority", -100)) - result: list[tuple[ProgramChannel, ProgramValue, int]] = [] - for channel in channel_names: - value = channels_map[channel] - if value is None or value is True: - continue - if value is False: - raise ValueError( - "program.channels setup should be removed instead of false." - ) - items = value if isinstance(value, list) else [value] - for item in items: - result.append((channel, cast(ProgramValue, item), priority)) - return result diff --git a/verifiers/v1/utils/prompt_utils.py b/verifiers/v1/utils/prompt_utils.py index 63bef46e0b..599e4f30ce 100644 --- a/verifiers/v1/utils/prompt_utils.py +++ b/verifiers/v1/utils/prompt_utils.py @@ -1,21 +1,21 @@ import importlib.util from dataclasses import dataclass from pathlib import Path -from typing import TYPE_CHECKING, Literal, TypeAlias, cast +from typing import TYPE_CHECKING, Literal, TypeAlias -from pydantic import model_validator +from pydantic import Field, TypeAdapter, model_validator from typing_extensions import Self from verifiers.types import Messages, SystemMessage -from verifiers.utils.message_utils import normalize_messages from ..config import Config from ..types import JsonData, PromptInput from .config_utils import current_config_ref_module if TYPE_CHECKING: - from ..state import State from ..task import Task +_MESSAGES_ADAPTER = TypeAdapter(Messages) + SystemPromptStrategy = Literal["REJECT", "TH", "HT", "T", "H", "T_OR_H", "H_OR_T"] SystemPromptTasksetSource = Literal["task", "taskset"] @@ -28,10 +28,12 @@ class SystemPromptResolution: taskset_source: SystemPromptTasksetSource | None def apply_strategy(self, strategy: SystemPromptStrategy) -> list[JsonData]: + harness = [dict(message) for message in self.harness] + taskset = [dict(message) for message in self.taskset] if strategy == "HT": - return [*self._copy(self.harness), *self._copy(self.taskset)] + return [*harness, *taskset] if strategy == "TH": - return [*self._copy(self.taskset), *self._copy(self.harness)] + return [*taskset, *harness] if strategy == "REJECT": if self.harness and self.taskset: raise ValueError( @@ -40,28 +42,24 @@ def apply_strategy(self, strategy: SystemPromptStrategy) -> list[JsonData]: "Set system_prompt_strategy='HT', 'TH', 'H', 'T', " "'H_OR_T', or 'T_OR_H'." ) - return [*self._copy(self.harness), *self._copy(self.taskset)] + return [*harness, *taskset] if strategy == "H_OR_T": - return self._copy(self.harness or self.taskset) + return harness or taskset if strategy == "T_OR_H": - return self._copy(self.taskset or self.harness) + return taskset or harness if strategy == "H": - return self._copy(self.harness) + return harness if strategy == "T": - return self._copy(self.taskset) + return taskset raise ValueError( "system_prompt_strategy must be one of REJECT, TH, HT, T, H, " "T_OR_H, H_OR_T." ) - @staticmethod - def _copy(messages: list[JsonData]) -> list[JsonData]: - return [dict(message) for message in messages] - class SystemPromptConfig(Config): path: str | None = None - messages: list[JsonData] = [] + messages: list[JsonData] = Field(default_factory=list) @model_validator(mode="after") def validate_one_input(self) -> Self: @@ -86,19 +84,6 @@ def load(self, field_name: str) -> PromptInput | None: SystemPrompt: TypeAlias = PromptInput | SystemPromptConfig | None -def normalize_prompt( - value: PromptInput | None, field_name: str = "prompt" -) -> list[JsonData]: - messages = normalize_messages(cast(Messages, value or []), field_name=field_name) - for message in messages: - if getattr(message, "role", None) == "system": - raise ValueError( - f"{field_name} must not contain system messages. " - "Use system_prompt instead." - ) - return dump_messages(messages) - - def normalize_system_prompt( value: SystemPrompt, field_name: str = "system_prompt", @@ -108,7 +93,7 @@ def normalize_system_prompt( return [] if isinstance(value, str): return [SystemMessage(content=value).model_dump(exclude_none=True)] - messages = normalize_messages(cast(Messages, value), field_name=field_name) + messages = _MESSAGES_ADAPTER.validate_python(value) for message in messages: if getattr(message, "role", None) != "system": raise ValueError(f"{field_name} accepts only system messages.") @@ -171,7 +156,7 @@ def system_prompt_resolution( harness_system_prompt: list[JsonData], ) -> SystemPromptResolution: task_system_prompt = normalize_system_prompt( - cast(PromptInput | None, task.get("system_prompt")), + task.system_prompt, field_name="task.system_prompt", ) return SystemPromptResolution( @@ -193,57 +178,3 @@ def system_prompt_resolution( def dump_messages(messages: Messages) -> list[JsonData]: return [message.model_dump(exclude_none=True) for message in messages] - - -def task_text( - task: "Task", - state: "State", - *, - keys: tuple[str, ...] = ("instruction",), -) -> str: - for key in keys: - value = task.get(key) - if isinstance(value, str) and value: - return value - return messages_text(task.get("prompt", [])) - - -def state_system_prompt_text(task: "Task", state: "State") -> str: - return messages_text(state.get("system_prompt", [])) - - -def messages_text(messages: object) -> str: - if isinstance(messages, str): - return messages - if not isinstance(messages, list): - return str(messages or "") - parts: list[str] = [] - for message in messages: - content = getattr(message, "content", None) - if content is not None: - parts.append(content_text(content)) - elif isinstance(message, dict): - item = cast(JsonData, message) - parts.append(content_text(item.get("content"))) - else: - parts.append(str(message)) - return "\n\n".join(part for part in parts if part) - - -def content_text(content: object) -> str: - if content is None: - return "" - if isinstance(content, str): - return content - if isinstance(content, list): - text_parts: list[str] = [] - for part in content: - if isinstance(part, dict): - item = cast(JsonData, part) - text = item.get("text") - if isinstance(text, str): - text_parts.append(text) - elif isinstance(part, str): - text_parts.append(part) - return "\n".join(text_parts) - return str(content) diff --git a/verifiers/v1/utils/runtime_owner_utils.py b/verifiers/v1/utils/runtime_owner_utils.py deleted file mode 100644 index cafa7c45d6..0000000000 --- a/verifiers/v1/utils/runtime_owner_utils.py +++ /dev/null @@ -1,139 +0,0 @@ -from collections.abc import Iterable -from typing import TYPE_CHECKING, Callable, Generic, TypeVar - -from ..config import LifecycleConfig, ToolsetCollectionData -from ..artifact import Artifacts, ArtifactsConfig -from ..toolset import ( - Toolset, - ToolsetCollection, - Toolsets, - collect_toolsets, - normalize_toolset_collection, -) -from ..types import Handler, Objects -from ..user import User, UserConfig, user_from_config -from .binding_utils import BindingSources, ObjectsConfig -from .config_callable_utils import CallableKind, merge_config_handler_map - -if TYPE_CHECKING: - from ..state import State - from ..task import Task - - -_HANDLER_KINDS: tuple[CallableKind, ...] = ( - "stop", - "setup", - "update", - "metric", - "reward", - "advantage", - "cleanup", - "teardown", -) - -ConfigT = TypeVar("ConfigT", bound=LifecycleConfig) - - -class RuntimeOwnerMixin(Generic[ConfigT]): - config: ConfigT - toolsets: list[Toolset] - named_toolsets: dict[str, Toolset] - stops: list[Handler] - setups: list[Handler] - updates: list[Handler] - metrics: list[Handler] - rewards: list[Handler] - advantages: list[Handler] - cleanups: list[Handler] - teardowns: list[Handler] - bindings: BindingSources - objects: Objects - artifacts: Artifacts - runtime_refresh: Callable[[], None] | None - - def load_user(self, config: UserConfig) -> User: - return user_from_config(config) - - def load_toolsets(self, config: ConfigT) -> Toolsets: - return None - - def load_objects(self, config: ObjectsConfig) -> Objects: - return config.objects(f"{type(self).__name__}.objects") - - def load_artifacts(self, config: ArtifactsConfig) -> Artifacts: - return config.artifacts(f"{type(self).__name__}.artifacts") - - async def get_object(self, name: str, task: "Task", state: "State") -> object: - return await state._runtime().resolve_owner_object(self, name, task, state) - - def initialize_runtime_refresh(self) -> None: - self.runtime_refresh = None - - def initialize_runtime_user(self, user: UserConfig | None) -> None: - self.user = None if user is None else self.load_user(user) - - def initialize_runtime_toolsets( - self, config: ConfigT, toolsets: ToolsetCollectionData - ) -> None: - self.toolsets, self.named_toolsets = collect_toolsets( - self.load_toolsets(config), toolsets - ) - - def initialize_runtime_handlers(self) -> None: - defaults: dict[CallableKind, Iterable[Handler]] = { - kind: () for kind in _HANDLER_KINDS - } - handlers = merge_config_handler_map(defaults, self.config) - self.stops = handlers["stop"] - self.setups = handlers["setup"] - self.updates = handlers["update"] - self.metrics = handlers["metric"] - self.rewards = handlers["reward"] - self.advantages = handlers["advantage"] - self.cleanups = handlers["cleanup"] - self.teardowns = handlers["teardown"] - - def refresh_runtime(self) -> None: - if self.runtime_refresh is not None: - self.runtime_refresh() - - def add_metric(self, fn: Handler) -> None: - self.metrics.append(fn) - self.refresh_runtime() - - def add_reward(self, fn: Handler) -> None: - self.rewards.append(fn) - self.refresh_runtime() - - def add_advantage(self, fn: Handler) -> None: - self.advantages.append(fn) - self.refresh_runtime() - - def add_toolset(self, toolset: ToolsetCollection) -> None: - toolsets, named_toolsets = normalize_toolset_collection(toolset) - duplicate = set(self.named_toolsets) & set(named_toolsets) - if duplicate: - raise ValueError(f"Toolsets are defined twice: {sorted(duplicate)}.") - self.toolsets.extend(toolsets) - self.named_toolsets.update(named_toolsets) - self.refresh_runtime() - - def add_stop(self, fn: Handler) -> None: - self.stops.append(fn) - self.refresh_runtime() - - def add_setup(self, fn: Handler) -> None: - self.setups.append(fn) - self.refresh_runtime() - - def add_update(self, fn: Handler) -> None: - self.updates.append(fn) - self.refresh_runtime() - - def add_cleanup(self, fn: Handler) -> None: - self.cleanups.append(fn) - self.refresh_runtime() - - def add_teardown(self, fn: Handler) -> None: - self.teardowns.append(fn) - self.refresh_runtime() diff --git a/verifiers/v1/utils/runtime_registry.py b/verifiers/v1/utils/runtime_registry.py deleted file mode 100644 index eebf8ed229..0000000000 --- a/verifiers/v1/utils/runtime_registry.py +++ /dev/null @@ -1,33 +0,0 @@ -import weakref -from typing import TYPE_CHECKING, cast - -if TYPE_CHECKING: - from ..runtime import Runtime - from ..state import State - -_RUNTIME_REGISTRY: weakref.WeakValueDictionary[str, object] = ( - weakref.WeakValueDictionary() -) - - -def register_runtime(runtime_id: str, runtime: object) -> None: - _RUNTIME_REGISTRY[runtime_id] = runtime - - -def unregister_runtime(runtime_id: str) -> None: - _RUNTIME_REGISTRY.pop(runtime_id, None) - - -def load_runtime(runtime_id: str) -> "Runtime": - runtime = _RUNTIME_REGISTRY.get(runtime_id) - if runtime is None: - raise RuntimeError(f"No live v1 runtime registered for id {runtime_id!r}.") - return cast("Runtime", runtime) - - -def load_runtime_from_state(state: "State") -> "Runtime": - runtime_state = state.runtime_state() - runtime_id = runtime_state.get("runtime_id") - if not isinstance(runtime_id, str) or not runtime_id: - raise RuntimeError("State has no live runtime id.") - return load_runtime(runtime_id) diff --git a/verifiers/v1/utils/sandbox_program_utils.py b/verifiers/v1/utils/sandbox_program_utils.py deleted file mode 100644 index 9cec1ded2c..0000000000 --- a/verifiers/v1/utils/sandbox_program_utils.py +++ /dev/null @@ -1,588 +0,0 @@ -import importlib.machinery -import importlib.util -import json -import shlex -import sys -import sysconfig -from dataclasses import dataclass -from pathlib import Path -from typing import cast - -from verifiers.errors import Error -from verifiers.utils.error_utils import error_from_data, validate_error_data -from verifiers.utils.interception_utils import serialize_tool_defs - -from ..runtime import Runtime -from ..sandbox import SandboxConfig -from ..state import State -from ..task import Task -from .serialization_utils import serializable -from .sandbox_utils import ( - VF_STATE_INPUT_PATH_KEY, - read_sandbox_artifact, - run_sandbox_command, -) -from .sandbox_python_utils import ( - python_package_list, - python_package_install_command, - python_runtime_command, - python_runtime_setup_command, -) -from .program_utils import ( - ProgramListInput, - ProgramMappingInput, - program_list_items, - program_option_mapping, -) -from ..types import ConfigData - -TASK_PATH = "/tmp/vf_task.json" -STATE_INPUT_PATH = "/tmp/vf_state_in.json" -STATE_OUTPUT_PATH = "/tmp/vf_state_out.json" -RUNNER_CONFIG_PATH = "/tmp/vf_runner_config.json" -TOOL_DEFS_PATH = "/tmp/vf_tool_defs.json" -TOOL_DEFS_BY_PROTOCOL_PATH = "/tmp/vf_tool_defs_by_protocol.json" -RUNNER_PATH = "/tmp/vf_program_runner.py" -PYTHON_PROGRAM_PACKAGES = ("openai", "anthropic", "requests") -PACKAGE_ROOT = "/tmp/vf_program_package" - - -def python_program_sandbox(sandbox_config: ConfigData) -> ConfigData: - config = dict(sandbox_config) - packages = python_package_list(config.get("packages")) - for package in PYTHON_PROGRAM_PACKAGES: - if not any(is_python_package(existing, package) for existing in packages): - packages.append(package) - config["packages"] = packages - return config - - -def is_python_package(requirement: str, package: str) -> bool: - return ( - requirement == package - or requirement.startswith(f"{package}[") - or requirement.startswith(f"{package}=") - or requirement.startswith(f"{package}<") - or requirement.startswith(f"{package}>") - or requirement.startswith(f"{package}~") - or requirement.startswith(f"{package}!") - ) - - -async def run_sandbox_python_program( - program: ConfigData, - sandbox_config: SandboxConfig, - task: Task, - state: State, - runtime: Runtime, - mode: str, - fn_ref: str | None, - max_turns: int, -) -> State: - runner_program = sandbox_runner_program( - program=program, - task=task, - state=state, - mode=mode, - fn_ref=fn_ref, - max_turns=max_turns, - tool_defs=runtime.tool_defs(state), - ) - command_record = state.get("command") - await run_sandbox_command(runner_program, sandbox_config, task, state, runtime) - lease = runtime.active_program_sandbox_lease(state) - if lease is None: - raise RuntimeError("Sandbox Python program has no active sandbox lease.") - output = json.loads( - await read_sandbox_artifact(lease.client, lease.id, STATE_OUTPUT_PATH) - ) - if not isinstance(output, dict): - raise RuntimeError("Sandbox Python program did not return state.") - patch = dict(cast(ConfigData, output)) - apply_internal_state_patch(state, patch, mode=mode) - patch_artifacts = patch.pop("artifacts", None) - if isinstance(patch_artifacts, dict): - state.setdefault("artifacts", {}) - state["artifacts"].update(dict(patch_artifacts)) - state.update(patch) - if command_record is not None: - state["command"] = command_record - return state - - -def apply_internal_state_patch(state: State, patch: ConfigData, *, mode: str) -> None: - for key in State.INTERNAL_KEYS: - if key not in patch: - continue - value = patch.pop(key) - if value == state.get(key): - continue - if mode != "base" or key == "is_completed": - raise RuntimeError( - f"Sandbox Python program cannot set framework-managed state key {key!r}." - ) - if key == "stop_condition": - state._set_stop_condition(cast(str | None, value), overwrite=True) - elif key == "is_truncated": - state._set_truncated(bool(value), overwrite=True) - elif key == "error": - state._set_error(state_error(value)) - else: - raise RuntimeError( - f"Sandbox Python program cannot set framework-managed state key {key!r}." - ) - - -def state_error(value: object) -> Error | None: - if value is None: - return None - if not isinstance(value, dict): - raise TypeError("Sandbox Python program error patch must be a mapping or None.") - return error_from_data(validate_error_data(value)) - - -def sandbox_runner_program( - program: ConfigData, - task: Task, - state: State, - mode: str, - fn_ref: str | None, - max_turns: int, - tool_defs: object, -) -> ConfigData: - package = sandbox_program_package(mode=mode, fn_ref=fn_ref) - if package is not None: - program = sandbox_program_with_package(program, package) - files = program_option_mapping( - cast(ProgramMappingInput, program.get("files")), "program.files" - ) - files[TASK_PATH] = json.dumps(task) - files[TOOL_DEFS_PATH] = json.dumps( - serializable(serialize_tool_defs(tool_defs or [], "openai_chat_completions")) - ) - files[TOOL_DEFS_BY_PROTOCOL_PATH] = json.dumps( - { - protocol: serializable(serialize_tool_defs(tool_defs or [], protocol)) - for protocol in ( - "vf", - "openai_chat_completions", - "openai_responses", - "anthropic_messages", - ) - } - ) - files[RUNNER_PATH] = runner_source() - files[RUNNER_CONFIG_PATH] = json.dumps({"max_turns": max_turns}) - command = python_runtime_command( - RUNNER_PATH, - *([mode] if fn_ref is None else [mode, fn_ref]), - ) - package_setup = [] if package is None else [package.install_command] - return { - **dict(program), - "files": files, - "command": command, - "env": program_option_mapping( - cast(ProgramMappingInput, program.get("env")), "program.env" - ), - "setup": [ - python_runtime_setup_command(), - *package_setup, - *program_list_items( - cast(ProgramListInput, program.get("setup")), "program.setup" - ), - ], - VF_STATE_INPUT_PATH_KEY: STATE_INPUT_PATH, - } - - -@dataclass(frozen=True) -class SandboxPackage: - local_root: Path - remote_root: str = PACKAGE_ROOT - - @property - def install_command(self) -> str: - return python_package_install_command(shlex.quote(self.remote_root)) - - -def sandbox_program_package(*, mode: str, fn_ref: str | None) -> SandboxPackage | None: - if mode != "fn" or fn_ref is None: - return None - module_name, _, _ = fn_ref.partition(":") - if not module_name: - raise ValueError("program.fn must include a module path.") - spec = importlib.util.find_spec(module_name) - if spec is None: - raise ImportError(f"Cannot resolve program.fn module {module_name!r}.") - roots = package_roots_for_module(module_name, spec) - if not roots: - return None - if len(roots) != 1: - raise ValueError( - f"program.fn {fn_ref!r} resolved to multiple package roots: " - f"{sorted(str(root) for root in roots)}." - ) - return SandboxPackage(local_root=next(iter(roots))) - - -def sandbox_program_with_package( - program: ConfigData, package: SandboxPackage -) -> ConfigData: - merged = dict(program) - dirs = program_option_mapping( - cast(ProgramMappingInput, merged.get("dirs")), "program.dirs" - ) - if package.remote_root in dirs: - raise ValueError( - f"program.dirs already defines internal package path {package.remote_root!r}." - ) - dirs[package.remote_root] = str(package.local_root) - merged["dirs"] = dirs - return merged - - -def package_roots_for_module( - module_name: str, spec: importlib.machinery.ModuleSpec -) -> set[Path]: - roots = set() - for path in module_source_paths(spec): - if is_external_import_path(path): - continue - root = module_package_root(path) - if root is None: - raise ValueError( - f"Sandboxed program.fn {module_name!r} resolves to local source " - f"{path}, but no pyproject.toml was found beside the resolved " - "environment module or package." - ) - roots.add(root) - return roots - - -def module_source_paths(spec: importlib.machinery.ModuleSpec) -> list[Path]: - origin = spec.origin - if spec.submodule_search_locations: - return [Path(path).resolve() for path in spec.submodule_search_locations] - if origin in {None, "built-in", "frozen"}: - return [] - return [Path(origin).resolve()] - - -def module_package_root(path: Path) -> Path | None: - root = path if path.is_dir() else path.parent - if (root / "pyproject.toml").is_file(): - return root - return None - - -def is_external_import_path(path: Path) -> bool: - parts = set(path.parts) - if "site-packages" in parts or "dist-packages" in parts: - return True - for prefix in interpreter_prefixes(): - try: - path.relative_to(prefix) - except ValueError: - continue - return True - return False - - -def interpreter_prefixes() -> list[Path]: - prefixes: list[Path] = [] - for key in ("stdlib", "platstdlib"): - value = sysconfig.get_path(key) - if value: - prefixes.append(Path(value).resolve()) - for value in (sys.base_prefix, sys.base_exec_prefix): - if value: - prefixes.append(Path(value).resolve()) - unique: list[Path] = [] - for prefix in prefixes: - if prefix not in unique: - unique.append(prefix) - return unique - - -def runner_source() -> str: - return r""" -import asyncio -import importlib -import inspect -import json -import os -import sys - -import requests -from anthropic import AsyncAnthropic -from openai import AsyncOpenAI - -TASK_PATH = "/tmp/vf_task.json" -STATE_INPUT_PATH = "/tmp/vf_state_in.json" -STATE_OUTPUT_PATH = "/tmp/vf_state_out.json" -RUNNER_CONFIG_PATH = "/tmp/vf_runner_config.json" -TOOL_DEFS_PATH = "/tmp/vf_tool_defs.json" -TOOL_DEFS_BY_PROTOCOL_PATH = "/tmp/vf_tool_defs_by_protocol.json" - - -class Client: - def __init__(self, state): - self.openai = AsyncOpenAI( - api_key=endpoint_token(), - base_url=os.environ.get("OPENAI_BASE_URL") - or state["endpoint_base_url"], - ) - self.anthropic = AsyncAnthropic( - api_key=os.environ.get("ANTHROPIC_API_KEY") - or endpoint_token(), - base_url=os.environ.get("ANTHROPIC_BASE_URL") - or state["endpoint_root_url"], - ) - self.chat = self.openai.chat - self.responses = self.openai.responses - self.messages = self.anthropic.messages - - async def close(self): - await self.openai.close() - await self.anthropic.close() - - -def endpoint_token(): - return os.environ.get("OPENAI_API_KEY") or os.environ.get("ANTHROPIC_API_KEY") or "intercepted" - - -def endpoint_headers(): - return { - "Content-Type": "application/json", - "Accept": "application/json", - "User-Agent": "python-requests/2.32.3", - "Authorization": f"Bearer {endpoint_token()}", - } - - -def vf_url(state, path): - return f"{state['endpoint_root_url'].rstrip('/')}/vf/{path}" - - -CONTROL_ENDPOINT_TIMEOUT = 300.0 - - -def post_json(url, payload, headers=None, timeout=CONTROL_ENDPOINT_TIMEOUT): - response = requests.post( - url, - json=payload, - headers=headers or endpoint_headers(), - timeout=timeout, - ) - try: - response.raise_for_status() - except requests.HTTPError as exc: - raise RuntimeError(response.text) from exc - if not response.content: - return {} - return response.json() - - -async def vf_post(state, path, payload, timeout=CONTROL_ENDPOINT_TIMEOUT): - return await asyncio.to_thread( - post_json, vf_url(state, path), payload, endpoint_headers(), timeout - ) - - -async def call_tool(state, name, arguments): - payload = await vf_post(state, f"tools/{name}", {"arguments": arguments}) - if "error" in payload: - raise RuntimeError(str(payload["error"])) - return payload.get("result") - - -async def call_user(state, transcript): - payload = await vf_post(state, "user", {"transcript": transcript}) - if "error" in payload: - raise RuntimeError(str(payload["error"])) - return payload.get("messages") or [] - - -def set_stop_condition(state, value, *, overwrite=False): - if overwrite or state.get("stop_condition") is None: - state["stop_condition"] = value - - -async def check_stop(state): - payload = await vf_post(state, "stop", {}) - if "error" in payload: - raise RuntimeError(str(payload["error"])) - if payload.get("done"): - if payload.get("stop_condition"): - set_stop_condition(state, payload["stop_condition"]) - return True - return False - - -async def maybe_call(fn, **objects): - sig = inspect.signature(fn) - if any(param.kind == param.VAR_KEYWORD for param in sig.parameters.values()): - result = fn(**objects) - else: - result = fn(**{key: value for key, value in objects.items() if key in sig.parameters}) - if inspect.isawaitable(result): - return await result - return result - - -def import_ref(ref): - module_name, _, attr_path = ref.partition(":") - obj = importlib.import_module(module_name) - for part in attr_path.split("."): - obj = getattr(obj, part) - return obj - - -def tool_call_name(tool_call): - # The /vf/model bridge returns canonical vf tool calls ({id, name, arguments}). - return tool_call["name"] - - -def tool_call_arguments(tool_call): - raw = tool_call.get("arguments") or "{}" - if isinstance(raw, str): - return json.loads(raw) - return raw - - -def tool_error_content(error): - return str(error) - - -def is_tool_content_parts(value): - # Mirror is_valid_tool_content_parts (can't import into the lean runner): - # pass {"type": text|image_url} part lists (e.g. screenshots) through, not str(). - if not isinstance(value, list): - return False - return all( - isinstance(part, dict) and part.get("type") in ("text", "image_url") - for part in value - ) - - -def load_tool_defs(protocol): - defs = json.loads(open(TOOL_DEFS_BY_PROTOCOL_PATH).read()) - return defs.get(protocol) or [] - - -class ToolProxy: - def __init__(self, state, name, description=None): - self.state = state - self.name = name - self.__name__ = name - self.__doc__ = description or "" - - async def __call__(self, **arguments): - return await call_tool(self.state, self.name, arguments) - - -def load_tools(state): - return { - tool["name"]: ToolProxy(state, tool["name"], tool.get("description")) - for tool in load_tool_defs("vf") - } - - -async def create_model_message(state, messages): - # The sandbox sends canonical Messages over the /vf/model bridge; the host - # resolves the bound client, tokenizes, and records the trajectory step, then - # returns the assistant message. The sandbox never formats a provider payload. - payload = await vf_post(state, "model", {"messages": messages}, timeout=None) - if "error" in payload: - raise RuntimeError(str(payload["error"])) - return payload["message"] - - -async def run_base(task, state): - prompt_messages = [*(state.get("system_prompt") or []), *(state.get("prompt") or [])] - messages = list(prompt_messages) - config = json.loads(open(RUNNER_CONFIG_PATH).read()) - max_turns = int(config["max_turns"]) - turn = 0 - while max_turns <= 0 or turn < max_turns: - if await check_stop(state): - break - message = await create_model_message(state, messages) - turn += 1 - messages.append(message) - tool_calls = list(message.get("tool_calls") or []) - if not tool_calls: - user_messages = await call_user(state, messages) - if user_messages: - messages.extend(user_messages) - continue - set_stop_condition(state, "no_tools") - break - for tool_call in tool_calls: - try: - result = await call_tool( - state, tool_call_name(tool_call), tool_call_arguments(tool_call) - ) - content = result if is_tool_content_parts(result) else str(result) - except Exception as exc: - content = tool_error_content(exc) - messages.append( - { - "role": "tool", - "tool_call_id": tool_call["id"], - "content": content, - } - ) - if await check_stop(state): - completed = True - break - else: - completed = False - if completed: - break - state["completion"] = messages[len(prompt_messages):] - set_stop_condition(state, "max_turns_reached") - return state - - -async def main(): - mode = sys.argv[1] - task = json.loads(open(TASK_PATH).read()) - state = json.loads(open(STATE_INPUT_PATH).read()) - original_state = json.loads(json.dumps(state)) - if mode == "base": - # Base loop talks to the host over /vf/model; no provider SDK client. - result = await run_base(task, state) - elif mode == "fn": - # fn-mode authors call the model via the OpenAI/Anthropic SDK, which the - # interception server transparently handles, so they still need a client. - client = Client(state) - try: - result = await maybe_call( - import_ref(sys.argv[2]), - task=task, - state=state, - client=client, - tools=load_tools(state), - tool_defs=load_tool_defs("vf"), - ) - finally: - await client.close() - else: - raise ValueError(f"Unknown sandbox program mode: {mode}") - if result is not None: - if not isinstance(result, dict): - raise TypeError("Sandbox Python program must return None or a mapping.") - state.update(result) - patch = { - key: value - for key, value in state.items() - if key not in original_state or original_state[key] != value - } - with open(STATE_OUTPUT_PATH, "w") as f: - json.dump(patch, f) - - -asyncio.run(main()) -""" diff --git a/verifiers/v1/utils/sandbox_python_utils.py b/verifiers/v1/utils/sandbox_python_utils.py deleted file mode 100644 index a853608751..0000000000 --- a/verifiers/v1/utils/sandbox_python_utils.py +++ /dev/null @@ -1,99 +0,0 @@ -import shlex - -SANDBOX_PYTHON_VERSION = "3.11" -SANDBOX_BIN_DIR = "/tmp/verifiers/bin" -SANDBOX_PYTHON_ROOT = "/tmp/verifiers/python" -SANDBOX_PYTHON_BIN_DIR = f"{SANDBOX_PYTHON_ROOT}/bin" -SANDBOX_PYTHON = f"{SANDBOX_PYTHON_BIN_DIR}/python3" -SANDBOX_UV = f"{SANDBOX_BIN_DIR}/uv" -SANDBOX_DEFAULT_PATH = ( - f"{SANDBOX_PYTHON_BIN_DIR}:" - f"{SANDBOX_BIN_DIR}:" - "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin" -) - - -def uv_setup_command() -> str: - return sandbox_runtime_script() + "vf_ensure_uv\n" - - -def python_runtime_setup_command() -> str: - return sandbox_runtime_script() + "vf_ensure_python\n" - - -def python_package_install_command(package_args: str = "") -> str: - command = python_runtime_setup_command() - if package_args: - command += f"vf_python_install {package_args}\n" - return command - - -def sandbox_python_path_command(command: str) -> str: - return f"export PATH={shlex.quote(SANDBOX_DEFAULT_PATH)}:$PATH\n{command}" - - -def python_package_list(value: object, field: str = "sandbox.packages") -> list[str]: - if value is None: - return [] - if isinstance(value, str): - return shlex.split(value) - if isinstance(value, list): - return [str(item) for item in value] - raise TypeError(f"{field} must be a list or string.") - - -def sandbox_runtime_script() -> str: - return ( - "set -e\n" - "export UV_NO_PROGRESS=1\n" - f"VF_BIN_DIR={shlex.quote(SANDBOX_BIN_DIR)}\n" - f"VF_PYTHON_ROOT={shlex.quote(SANDBOX_PYTHON_ROOT)}\n" - f"VF_PYTHON={shlex.quote(SANDBOX_PYTHON)}\n" - f"VF_UV={shlex.quote(SANDBOX_UV)}\n" - f"VF_PYTHON_VERSION={shlex.quote(SANDBOX_PYTHON_VERSION)}\n" - 'export PATH="$VF_PYTHON_ROOT/bin:$VF_BIN_DIR:$PATH"\n' - 'mkdir -p "$VF_BIN_DIR"\n' - "vf_ensure_uv() {\n" - 'if [ ! -x "$VF_UV" ]; then\n' - " if command -v uv >/dev/null 2>&1; then\n" - ' cp "$(command -v uv)" "$VF_UV"\n' - " else\n" - " if ! command -v curl >/dev/null 2>&1; then\n" - " if ! command -v apt-get >/dev/null 2>&1; then\n" - " echo 'curl or apt-get is required to install uv' >&2; exit 127\n" - " fi\n" - " apt-get -o Acquire::Retries=3 update && " - "apt-get -o Acquire::Retries=3 install -y curl ca-certificates\n" - " fi\n" - " if ! curl -LsSf https://astral.sh/uv/install.sh -o /tmp/vf-uv-install.sh; then\n" - " if ! command -v apt-get >/dev/null 2>&1; then exit 1; fi\n" - " apt-get -o Acquire::Retries=3 update && " - "apt-get -o Acquire::Retries=3 install -y ca-certificates\n" - " curl -LsSf https://astral.sh/uv/install.sh -o /tmp/vf-uv-install.sh\n" - " fi\n" - ' env UV_INSTALL_DIR="$VF_BIN_DIR" sh /tmp/vf-uv-install.sh\n' - " rm -f /tmp/vf-uv-install.sh\n" - " fi\n" - "fi\n" - "}\n" - "vf_ensure_python() {\n" - "vf_ensure_uv\n" - 'if [ ! -x "$VF_PYTHON" ]; then\n' - ' "$VF_UV" venv --seed --python "$VF_PYTHON_VERSION" "$VF_PYTHON_ROOT"\n' - "fi\n" - 'if [ ! -x "$VF_PYTHON" ] && [ -x "$VF_PYTHON_ROOT/bin/python" ]; then\n' - ' ln -sfn "$VF_PYTHON_ROOT/bin/python" "$VF_PYTHON"\n' - "fi\n" - 'if [ ! -x "$VF_PYTHON_ROOT/bin/python" ] && [ -x "$VF_PYTHON" ]; then\n' - ' ln -sfn "$VF_PYTHON" "$VF_PYTHON_ROOT/bin/python"\n' - "fi\n" - "}\n" - "vf_python_install() {\n" - "vf_ensure_python\n" - ' "$VF_UV" pip install --python "$VF_PYTHON" "$@"\n' - "}\n" - ) - - -def python_runtime_command(script_path: str, *args: str) -> list[str]: - return [SANDBOX_PYTHON, script_path, *args] diff --git a/verifiers/v1/utils/sandbox_utils.py b/verifiers/v1/utils/sandbox_utils.py deleted file mode 100644 index 7f0d09874a..0000000000 --- a/verifiers/v1/utils/sandbox_utils.py +++ /dev/null @@ -1,1140 +0,0 @@ -import asyncio -import hashlib -import importlib.resources as resources -import json -import logging -import shlex -import tarfile -import tempfile -import uuid -from collections.abc import Awaitable, Callable -from importlib.abc import Traversable -from pathlib import Path -from typing import TYPE_CHECKING, Literal, Protocol, TypeVar, cast - -import tenacity as tc - -from verifiers.decorators import setup as setup_handler -from verifiers.errors import Error, SandboxError -from verifiers.utils.async_utils import maybe_call_with_named_args - -from .program_utils import command_argv, command_env, float_config, int_config -from .program_utils import program_channels -from .program_utils import ( - ProgramMappingInput, - program_option_mapping, - program_channel_setup, -) -from .program_utils import resolve_program_value -from .program_utils import validate_program_bindings -from .sandbox_python_utils import ( - python_package_install_command, - python_package_list, - sandbox_python_path_command, -) -from ..runtime import Runtime -from ..sandbox import SandboxConfig -from ..program import ProgramValue -from ..state import State -from ..task import Task -from ..types import ConfigData, Handler - -if TYPE_CHECKING: - from ..toolset import Toolset - -VF_STATE_INPUT_PATH_KEY = "_vf_state_input_path" -SANDBOX_RETRY_ATTEMPTS = 6 -SANDBOX_WAIT_FOR_CREATION_ATTEMPTS = 120 -SANDBOX_TRANSFER_RETRY_TOKENS = ( - "502", - "503", - "504", - "ConnectError", - "Read file timed out", - "Temporary failure in name resolution", -) -T = TypeVar("T") -logger = logging.getLogger(__name__) - - -class SandboxRecord(Protocol): - id: object - - -class RetryLogger(Protocol): - def log( - self, level: int, msg: str, /, *args: object, **kwargs: object - ) -> object: ... - - -class SandboxCommandResult(Protocol): - exit_code: int - stdout: str | None - stderr: str | None - - -class SandboxOwner(Protocol): - @property - def sandbox(self) -> SandboxConfig | Literal["program"] | None: ... - - -class SandboxClient(Protocol): - async def create(self, request: object) -> SandboxRecord: ... - - async def wait_for_creation( - self, - sandbox_id: str, - *, - max_attempts: int = SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) -> object: ... - - async def delete(self, sandbox_id: str) -> object: ... - - async def aclose(self) -> object: ... - - async def execute_command( - self, - sandbox_id: str, - command: str, - *, - timeout: int | None = None, - working_dir: str | None = None, - env: dict[str, str] | None = None, - ) -> SandboxCommandResult: ... - - async def upload_bytes( - self, - sandbox_id: str, - file_path: str, - file_bytes: bytes, - *, - filename: str | None = None, - ) -> object: ... - - async def upload_file( - self, - sandbox_id: str, - file_path: str, - local_file_path: str, - *, - timeout: int | None = None, - ) -> object: ... - - async def download_file( - self, - sandbox_id: str, - file_path: str, - local_file_path: str, - *, - timeout: int | None = None, - ) -> object: ... - - async def read_file(self, sandbox_id: str, path: str) -> object: ... - - async def run_background_job( - self, - sandbox_id: str, - command: str, - *, - timeout: int | None = None, - working_dir: str | None = None, - env: dict[str, str] | None = None, - poll_interval: int = 3, - ) -> SandboxCommandResult: ... - - -async def with_sandbox_retry(operation: Callable[[], Awaitable[T]]) -> T: - retry_logger = cast(RetryLogger, logger) - async for attempt in tc.AsyncRetrying( - stop=tc.stop_after_attempt(SANDBOX_RETRY_ATTEMPTS), - wait=tc.wait_exponential_jitter(initial=0.5, max=30, jitter=1e-3), - before_sleep=tc.before_sleep_log(retry_logger, logging.WARNING), - sleep=asyncio.sleep, - reraise=True, - ): - with attempt: - return await operation() - raise AssertionError("sandbox retry loop exited without running") - - -def is_retryable_sandbox_transfer_error(exc: BaseException) -> bool: - name = type(exc).__name__ - if name.endswith(("UploadTimeoutError", "DownloadTimeoutError")): - return True - if not name.endswith("APIError"): - return False - text = str(exc) - return text.strip() in {"Upload failed", "Upload failed:"} or any( - token in text for token in SANDBOX_TRANSFER_RETRY_TOKENS - ) - - -async def with_sandbox_transfer_retry(operation: Callable[[], Awaitable[T]]) -> T: - retry_logger = cast(RetryLogger, logger) - async for attempt in tc.AsyncRetrying( - stop=tc.stop_after_attempt(SANDBOX_RETRY_ATTEMPTS), - wait=tc.wait_exponential_jitter(initial=0.5, max=30, jitter=1e-3), - retry=tc.retry_if_exception(is_retryable_sandbox_transfer_error), - before_sleep=tc.before_sleep_log(retry_logger, logging.WARNING), - sleep=asyncio.sleep, - reraise=True, - ): - with attempt: - return await operation() - raise AssertionError("sandbox transfer retry loop exited without running") - - -async def close_sandbox_client(client: SandboxClient) -> None: - teardown = getattr(client, "teardown", None) - if callable(teardown): - teardown() - return - aclose = getattr(client, "aclose", None) - if callable(aclose): - await aclose() - - -def sandbox_failure_kind(exc: BaseException) -> str | None: - if isinstance(exc, TimeoutError): - return "timeout" - name = type(exc).__name__ - text = str(exc) - if name == "SandboxOOMError" or "OOM" in text or "OOM_KILLED" in text: - return "oom" - if name in {"SandboxTimeoutError", "CommandTimeoutError"} or "timed out" in text: - return "timeout" - return None - - -def mark_sandbox_failure( - state: State, - lease: "SandboxLease | None", - exc: BaseException, - *, - phase: str | None = None, -) -> str | None: - kind = sandbox_failure_kind(exc) - if kind == "oom": - state["sandbox_oom"] = True - elif kind == "timeout": - state["sandbox_timeout"] = True - if kind is not None: - failure: ConfigData = { - "kind": kind, - "type": type(exc).__name__, - "message": str(exc), - } - if phase is not None: - failure["phase"] = phase - if lease is not None: - failure["sandbox_id"] = lease.id - failure["scope"] = lease.scope - state.setdefault("sandbox_failures", []).append(failure) - return kind - - -class SandboxLease: - def __init__( - self, - client: SandboxClient, - sandbox_id: str, - scope: str, - key: str, - *, - owns_client: bool = True, - ): - self.client = client - self.id = sandbox_id - self.scope = scope - self.key = key - self.owns_client = owns_client - self.scope_key: str | None = None - self.deleted = False - self.lock = asyncio.Lock() - self.delete_lock = asyncio.Lock() - - async def execute( - self, - command: str, - timeout: int | None = None, - working_dir: str | None = None, - env: dict[str, str] | None = None, - ) -> SandboxCommandResult: - result = await maybe_call_with_named_args( - self.client.execute_command, - sandbox_id=self.id, - command=command, - timeout=timeout, - working_dir=working_dir, - env=env, - ) - return cast(SandboxCommandResult, result) - - async def upload_bytes( - self, path: str, content: bytes, filename: str | None = None - ) -> object: - return await maybe_call_with_named_args( - self.client.upload_bytes, - sandbox_id=self.id, - file_path=path, - file_bytes=content, - filename=filename or path.rsplit("/", 1)[-1] or "file", - ) - - async def upload_file( - self, path: str, local_path: str, timeout: int | None = None - ) -> object: - return await maybe_call_with_named_args( - self.client.upload_file, - sandbox_id=self.id, - file_path=path, - local_file_path=local_path, - timeout=timeout, - ) - - async def download_file( - self, path: str, local_path: str, timeout: int | None = None - ) -> object: - return await maybe_call_with_named_args( - self.client.download_file, - sandbox_id=self.id, - file_path=path, - local_file_path=local_path, - timeout=timeout, - ) - - async def read_file(self, path: str) -> object: - return await maybe_call_with_named_args( - self.client.read_file, - sandbox_id=self.id, - path=path, - ) - - async def run_background_job( - self, - command: str, - timeout: int | None = None, - working_dir: str | None = None, - env: dict[str, str] | None = None, - poll_interval: int = 3, - ) -> SandboxCommandResult: - call_args: ConfigData = { - "sandbox_id": self.id, - "command": command, - "working_dir": working_dir, - "env": env, - "poll_interval": poll_interval, - } - if timeout is not None: - call_args["timeout"] = timeout - return await maybe_call_with_named_args( - getattr(self.client, "run_background_job"), - **call_args, - ) - - async def delete(self) -> None: - async with self.delete_lock: - if self.deleted: - return - self.deleted = True - try: - await with_sandbox_retry(lambda: self.client.delete(self.id)) - except BaseException: - self.deleted = False - raise - if self.owns_client: - await close_sandbox_client(self.client) - - -class SandboxHandle: - def __init__(self, lease: SandboxLease, state: State): - self.lease = lease - self.state = state - self.id = lease.id - self.scope = lease.scope - self.key = lease.key - attach_sandbox_ref(state, lease) - - async def execute( - self, - command: str, - timeout: int | None = None, - working_dir: str | None = None, - env: dict[str, str] | None = None, - ) -> SandboxCommandResult: - try: - result = await self.lease.execute( - command=command, - timeout=timeout, - working_dir=working_dir, - env=env, - ) - except Error: - raise - except Exception as exc: - kind = mark_sandbox_failure(self.state, self.lease, exc, phase="execute") - if kind is not None: - raise SandboxError( - f"Sandbox {self.lease.id} failed during execute ({kind}): {exc}" - ) from exc - raise - record_tool_sandbox_command(self.state, self.lease, command, result) - return result - - async def upload_bytes( - self, path: str, content: bytes, filename: str | None = None - ) -> object: - return await self.lease.upload_bytes(path, content, filename) - - async def upload_file( - self, path: str, local_path: str, timeout: int | None = None - ) -> object: - return await self.lease.upload_file(path, local_path, timeout) - - async def download_file( - self, path: str, local_path: str, timeout: int | None = None - ) -> object: - return await self.lease.download_file(path, local_path, timeout) - - async def read_file(self, path: str) -> object: - return await self.lease.read_file(path) - - async def run_background_job( - self, - command: str, - timeout: int | None = None, - working_dir: str | None = None, - env: dict[str, str] | None = None, - poll_interval: int = 3, - ) -> SandboxCommandResult: - try: - result = await self.lease.run_background_job( - command=command, - timeout=timeout, - working_dir=working_dir, - env=env, - poll_interval=poll_interval, - ) - except Error: - raise - except Exception as exc: - kind = mark_sandbox_failure( - self.state, self.lease, exc, phase="background_job" - ) - if kind is not None: - raise SandboxError( - f"Sandbox {self.lease.id} failed during background job ({kind}): {exc}" - ) from exc - raise - record_tool_sandbox_command(self.state, self.lease, command, result) - return result - - async def delete(self) -> None: - await self.lease.delete() - - -async def create_sandbox_lease( - sandbox_config: SandboxConfig, - key: str, - client: SandboxClient | None = None, -) -> SandboxLease: - sandbox_data = sandbox_config.data() - owns_client = client is None - if client is None: - from verifiers.utils.threaded_sandbox_client import ThreadedAsyncSandboxClient - - client = cast(SandboxClient, ThreadedAsyncSandboxClient()) - sandbox_id = await create_sandbox(client, sandbox_data, owns_client=owns_client) - lease = SandboxLease( - client, sandbox_id, sandbox_config.scope, key, owns_client=owns_client - ) - try: - await setup_sandbox(lease, sandbox_data) - except BaseException: - await asyncio.shield(lease.delete()) - raise - return lease - - -async def create_scoped_sandbox_lease( - owner: SandboxOwner, - key: str | None = None, - client: SandboxClient | None = None, -) -> SandboxLease: - sandbox = owner.sandbox - if not isinstance(sandbox, SandboxConfig): - raise TypeError("Sandbox owner must define a sandbox config.") - return await create_sandbox_lease(sandbox, key or sandbox_owner_key(owner), client) - - -async def run_sandbox_command( - program: ConfigData, - sandbox_config: SandboxConfig, - task: Task, - state: State, - runtime: Runtime, -) -> State: - validate_program_bindings(program) - sandbox_data = sandbox_config.data() - try: - lease = await runtime.resolve_program_sandbox(sandbox_config, task, state) - except Exception as exc: - mark_sandbox_failure(state, None, exc, phase="create") - raise - async with lease.lock: - state["sandbox_id"] = lease.id - runtime_state = state.runtime_state() - lease_scope_key = lease.scope_key or runtime.scope_key(lease.scope, state) - lease.scope_key = lease_scope_key - runtime_state["sandbox"] = { - "id": lease.id, - "scope": lease.scope, - "key": lease.key, - "lease_key": [lease_scope_key, lease.key], - } - handle = SandboxHandle(lease, state) - use_sandbox_python_path = bool( - python_package_list(sandbox_data.get("packages")) - ) - try: - await runtime.setup_rollout( - task, - state, - setup_handlers=program_setup_handlers( - lease, - program, - runtime, - use_sandbox_python_path=use_sandbox_python_path, - ), - sandbox=handle, - ) - except Exception as exc: - mark_sandbox_failure(state, lease, exc, phase="setup") - raise - workdir = sandbox_config.workdir - if workdir: - await lease.client.execute_command( - lease.id, f"mkdir -p {shlex.quote(workdir)}" - ) - argv = await command_argv(program, task, state, runtime) - env = await command_env(program, task, state, runtime, include_base=False) - command = shlex.join(argv) - if use_sandbox_python_path or "mcp" in program_channels(program): - command = sandbox_python_path_command(command) - command_timeout = sandbox_config.command_timeout - try: - result = await lease.run_background_job( - command, - timeout=command_timeout, - working_dir=workdir, - env=env, - poll_interval=int_config(sandbox_data, "poll_interval", 3), - ) - except Error: - raise - except Exception as exc: - kind = mark_sandbox_failure(state, lease, exc, phase="command") - if kind is not None: - raise SandboxError( - f"Sandbox {lease.id} failed during command ({kind}): {exc}" - ) from exc - raise - state["command"] = { - "argv": argv, - "returncode": result.exit_code, - "stdout": result.stdout or "", - "stderr": result.stderr or "", - } - state["completion"] = [ - {"role": "assistant", "content": state["command"]["stdout"].strip()} - ] - if result.exit_code: - raise SandboxError( - f"Sandbox command exited with {result.exit_code}: {result.stderr}" - ) - state._set_stop_condition("command_completed") - return state - - -def program_setup_handlers( - lease: SandboxLease, - program: ConfigData, - runtime: Runtime, - *, - use_sandbox_python_path: bool = False, -) -> list[Handler]: - handlers: list[Handler] = [ - _program_setup_handler( - lease, - program, - runtime, - upload_program_files, - "program_upload_files", - 200, - ), - _program_setup_handler( - lease, - program, - runtime, - upload_program_dirs, - "program_upload_dirs", - 190, - ), - _program_setup_handler( - lease, - program, - runtime, - run_program_setup, - "program_setup", - 100, - use_sandbox_python_path=use_sandbox_python_path, - ), - _program_setup_handler( - lease, - program, - runtime, - upload_state_input, - "program_state_input", - -50, - ), - ] - for channel, setup_item, priority in program_channel_setup(program): - handlers.append( - _program_channel_setup_handler( - lease, - program, - runtime, - str(channel), - setup_item, - priority, - use_sandbox_python_path=use_sandbox_python_path, - ) - ) - return handlers - - -def _program_setup_handler( - lease: SandboxLease, - program: ConfigData, - runtime: Runtime, - fn: Callable[..., Awaitable[None]], - name: str, - priority: int, - use_sandbox_python_path: bool = False, -) -> Handler: - async def handler(task: Task, state: State) -> None: - try: - await maybe_call_with_named_args( - fn, - client=lease.client, - sandbox_id=lease.id, - program=program, - task=task, - state=state, - runtime=runtime, - use_sandbox_python_path=use_sandbox_python_path, - ) - except Error: - raise - except Exception as exc: - raise SandboxError(f"Sandbox setup handler {name} failed: {exc}") from exc - - handler.__name__ = name - return setup_handler(handler, priority=priority) - - -def _program_channel_setup_handler( - lease: SandboxLease, - program: ConfigData, - runtime: Runtime, - channel: str, - setup_item: ProgramValue, - priority: int, - use_sandbox_python_path: bool = False, -) -> Handler: - name = f"program_{channel}_channel_setup" - - async def handler(task: Task, state: State) -> None: - try: - await run_program_items( - lease.client, - lease.id, - program, - task, - state, - runtime, - items=[setup_item], - error_prefix=f"Program {channel} channel setup failed", - use_sandbox_python_path=use_sandbox_python_path, - ) - except Error: - raise - except Exception as exc: - raise SandboxError( - f"Sandbox {channel} channel setup handler {name} failed: {exc}" - ) from exc - - handler.__name__ = name - return setup_handler(handler, priority=priority) - - -async def create_sandbox( - client: SandboxClient, - sandbox_config: ConfigData, - *, - owns_client: bool = False, -) -> str: - from prime_sandboxes import CreateSandboxRequest - - labels = sandbox_config.get("labels") - gpu_count = int_config(sandbox_config, "gpu_count", 0) - vm = sandbox_config.get("vm") - environment_vars = sandbox_config.get("environment_vars") - secrets = sandbox_config.get("secrets") - request = CreateSandboxRequest( - name=f"vf-v1-{uuid.uuid4().hex[:8]}", - docker_image=str(sandbox_config.get("image") or "python:3.11-slim"), - start_command=str(sandbox_config.get("start_command") or "tail -f /dev/null"), - cpu_cores=float_config(sandbox_config, "cpu_cores", 1.0), - memory_gb=float_config(sandbox_config, "memory_gb", 2.0), - disk_size_gb=float_config(sandbox_config, "disk_size_gb", 5.0), - gpu_count=gpu_count, - gpu_type=str(sandbox_config["gpu_type"]) - if sandbox_config.get("gpu_type") is not None - else None, - vm=bool(vm) if vm is not None else gpu_count > 0, - network_access=bool(sandbox_config.get("network_access", True)), - timeout_minutes=int_config(sandbox_config, "timeout_minutes", 60), - labels=[str(label) for label in labels] if isinstance(labels, list) else [], - environment_vars={ - str(key): str(value) for key, value in environment_vars.items() - } - if isinstance(environment_vars, dict) and environment_vars - else None, - secrets={str(key): str(value) for key, value in secrets.items()} - if isinstance(secrets, dict) and secrets - else None, - team_id=str(sandbox_config["team_id"]) - if sandbox_config.get("team_id") is not None - else None, - region=str(sandbox_config["region"]) - if sandbox_config.get("region") is not None - else None, - registry_credentials_id=str(sandbox_config["registry_credentials_id"]) - if sandbox_config.get("registry_credentials_id") is not None - else None, - guaranteed=bool(sandbox_config.get("guaranteed", False)), - ) - create_task = asyncio.create_task( - with_sandbox_retry(lambda: client.create(request)) - ) - try: - create_waiter = asyncio.shield(create_task) - if sandbox_config.get("create_timeout") is not None: - sandbox = await asyncio.wait_for( - create_waiter, int_config(sandbox_config, "create_timeout", 0) - ) - else: - sandbox = await create_waiter - except (asyncio.CancelledError, TimeoutError): - try: - sandbox = cast(SandboxRecord, await asyncio.shield(create_task)) - except BaseException: - if owns_client: - await close_sandbox_client(client) - raise - await asyncio.shield( - delete_sandbox_id( - client, - str(sandbox.id), - close_client=owns_client, - reason="cancelled creation", - ) - ) - raise - except BaseException: - if owns_client: - await close_sandbox_client(client) - raise - sandbox_id = str(sandbox.id) - try: - wait = client.wait_for_creation( - sandbox_id, - max_attempts=SANDBOX_WAIT_FOR_CREATION_ATTEMPTS, - ) - if sandbox_config.get("wait_timeout") is not None: - await asyncio.wait_for(wait, int_config(sandbox_config, "wait_timeout", 0)) - else: - await wait - except BaseException: - delete_task = asyncio.create_task( - delete_sandbox_id( - client, - sandbox_id, - close_client=owns_client, - reason="creation failure", - ) - ) - await asyncio.shield(delete_task) - raise - return sandbox_id - - -async def delete_sandbox_id( - client: SandboxClient, - sandbox_id: str, - *, - close_client: bool, - reason: str, -) -> None: - try: - await with_sandbox_retry(lambda: client.delete(sandbox_id)) - except Exception as cleanup_exc: - logger.warning( - "Failed to delete sandbox %s after %s: %s", - sandbox_id, - reason, - cleanup_exc, - exc_info=True, - ) - finally: - if close_client: - await close_sandbox_client(client) - - -async def setup_sandbox(handle: SandboxLease, sandbox_config: ConfigData) -> None: - packages = python_package_list(sandbox_config.get("packages")) - if packages: - package_args = " ".join(shlex.quote(str(package)) for package in packages) - try: - result = await handle.execute( - python_package_install_command(package_args), - timeout=int_config(sandbox_config, "install_timeout", 300), - ) - except Error: - raise - except Exception as exc: - raise SandboxError(f"Sandbox package install failed: {exc}") from exc - if result.exit_code: - raise SandboxError(f"Sandbox package install failed: {result.stderr}") - commands = sandbox_config.get("setup_commands") or [] - if isinstance(commands, str): - commands = [commands] - if not isinstance(commands, list): - raise TypeError("sandbox.setup_commands must be a list or string.") - use_sandbox_python_path = bool(packages) - for command in commands: - command = str(command) - if use_sandbox_python_path: - command = sandbox_python_path_command(command) - try: - result = await handle.execute( - command, - timeout=int_config(sandbox_config, "setup_timeout", 300), - ) - except Error: - raise - except Exception as exc: - raise SandboxError(f"Sandbox setup command failed: {exc}") from exc - if result.exit_code: - raise SandboxError(f"Sandbox setup command failed: {result.stderr}") - - -def attach_sandbox_ref(state: State, lease: SandboxLease) -> None: - sandboxes = cast(ConfigData, state.runtime_state().setdefault("sandboxes", {})) - sandboxes[lease.key] = {"id": lease.id, "scope": lease.scope} - - -def record_tool_sandbox_command( - state: State, lease: SandboxLease, command: str, result: SandboxCommandResult -) -> None: - command_record: ConfigData = { - "command": command, - "returncode": result.exit_code, - "stdout": result.stdout or "", - "stderr": result.stderr or "", - } - commands = cast(list[ConfigData], state.setdefault("sandbox_commands", [])) - commands.append(command_record) - sandboxes = cast(ConfigData, state.runtime_state().setdefault("sandboxes", {})) - tool_state = cast( - ConfigData, - sandboxes.setdefault(lease.key, {"id": lease.id, "scope": lease.scope}), - ) - tool_commands = cast(list[ConfigData], tool_state.setdefault("commands", [])) - tool_commands.append(command_record) - - -def tool_sandbox_key(toolset: "Toolset") -> str: - from ..toolset import MCPTool, flatten_toolsets, tool_name - - names = [ - tool_name(tool) - for tool in flatten_toolsets((toolset,)) - if not isinstance(tool, MCPTool) - ] - if names: - return "tools:" + ",".join(sorted(names)) - return f"toolset:{id(toolset)}" - - -def program_sandbox_key(sandbox_config: SandboxConfig) -> str: - try: - fingerprint = json.dumps(sandbox_config.data(), sort_keys=True) - except TypeError as exc: - raise TypeError("Program sandbox config must be JSON-serializable.") from exc - digest = hashlib.sha256(fingerprint.encode()).hexdigest()[:12] - return f"program:{digest}" - - -def sandbox_owner_key(owner: object) -> str: - return f"sandbox:{id(owner)}" - - -async def upload_program_files( - client: SandboxClient, - sandbox_id: str, - program: ConfigData, - task: Task, - state: State, - runtime: Runtime, -) -> None: - from prime_sandboxes import APIError, UploadTimeoutError - - files = program_option_mapping( - cast(ProgramMappingInput, program.get("files")), "program.files" - ) - for path, source in files.items(): - content = await resolve_program_value(source, task, state, runtime, program) - if not isinstance(content, str): - content = str(content) - try: - await with_sandbox_transfer_retry( - lambda: maybe_call_with_named_args( - getattr(client, "upload_bytes"), - sandbox_id=sandbox_id, - file_path=path, - file_bytes=content.encode(), - filename=path.rsplit("/", 1)[-1] or "file", - ) - ) - except (APIError, UploadTimeoutError) as exc: - raise SandboxError( - f"Program file upload failed for {path!r} in sandbox {sandbox_id}: {exc}" - ) from exc - - -async def upload_program_dirs( - client: SandboxClient, - sandbox_id: str, - program: ConfigData, - task: Task, - state: State, - runtime: Runtime, -) -> None: - dirs = program_option_mapping( - cast(ProgramMappingInput, program.get("dirs")), "program.dirs" - ) - for path, source in dirs.items(): - local_source = await resolve_program_value( - source, task, state, runtime, program - ) - if local_source is None: - continue - if isinstance(local_source, str): - local_source = Path(local_source) - if not isinstance(local_source, (Path, Traversable)): - raise TypeError("program.dirs values must resolve to paths.") - remote_tar = f"/tmp/_vf_upload_{path.strip('/').replace('/', '_')}.tar.gz" - archive_path = await runtime.cached_upload_archive(local_source, path) - await with_sandbox_transfer_retry( - lambda: maybe_call_with_named_args( - client.upload_file, - sandbox_id=sandbox_id, - file_path=remote_tar, - local_file_path=str(archive_path), - ) - ) - result = await maybe_call_with_named_args( - client.execute_command, - sandbox_id=sandbox_id, - command=( - f"mkdir -p {shlex.quote(str(Path(path).parent))} && " - f"tar -xzf {shlex.quote(remote_tar)} -C / && " - f"rm -f {shlex.quote(remote_tar)}" - ), - ) - if result.exit_code: - raise SandboxError(f"Program dir upload failed: {result.stderr}") - - -def build_dir_archive(local_source: Path | Traversable, remote_path: str) -> Path: - with tempfile.NamedTemporaryFile(suffix=".tar.gz", delete=False) as tmp_file: - tar_path = Path(tmp_file.name) - arcname = remote_path.lstrip("/") - try: - with tarfile.open(tar_path, "w:gz") as tar: - if isinstance(local_source, Path): - tar.add(local_source, arcname=arcname, filter=upload_tar_filter) - else: - with resources.as_file(local_source) as local_path: - tar.add(local_path, arcname=arcname, filter=upload_tar_filter) - except BaseException: - tar_path.unlink(missing_ok=True) - raise - return tar_path - - -UPLOAD_IGNORE_PARTS = { - ".git", - ".venv", - "__pycache__", - ".mypy_cache", - ".pytest_cache", - ".ruff_cache", - "node_modules", -} - - -def upload_tar_filter(tarinfo: tarfile.TarInfo) -> tarfile.TarInfo | None: - if any(part in UPLOAD_IGNORE_PARTS for part in Path(tarinfo.name).parts): - return None - return tarinfo - - -async def run_program_setup( - client: SandboxClient, - sandbox_id: str, - program: ConfigData, - task: Task, - state: State, - runtime: Runtime, - use_sandbox_python_path: bool = False, -) -> None: - await run_program_commands( - client, - sandbox_id, - program, - task, - state, - runtime, - key="setup", - error_prefix="Program setup failed", - use_sandbox_python_path=use_sandbox_python_path, - ) - - -async def upload_state_input( - client: SandboxClient, - sandbox_id: str, - program: ConfigData, - state: State, -) -> None: - path = program.get(VF_STATE_INPUT_PATH_KEY) - if path is None: - return - if not isinstance(path, str): - raise TypeError(f"{VF_STATE_INPUT_PATH_KEY} must be a string.") - await with_sandbox_transfer_retry( - lambda: maybe_call_with_named_args( - client.upload_bytes, - sandbox_id=sandbox_id, - file_path=path, - file_bytes=json.dumps(state).encode(), - filename=path.rsplit("/", 1)[-1] or "file", - ) - ) - - -async def run_program_commands( - client: SandboxClient, - sandbox_id: str, - program: ConfigData, - task: Task, - state: State, - runtime: Runtime, - *, - key: str, - error_prefix: str, - use_sandbox_python_path: bool = False, -) -> None: - raw_setup = program.get(key) or [] - if isinstance(raw_setup, str): - setup: list[ProgramValue] = [raw_setup] - elif isinstance(raw_setup, list): - setup = [cast(ProgramValue, item) for item in raw_setup] - else: - setup = [cast(ProgramValue, raw_setup)] - await run_program_items( - client, - sandbox_id, - program, - task, - state, - runtime, - items=setup, - error_prefix=error_prefix, - use_sandbox_python_path=use_sandbox_python_path, - ) - - -async def run_program_items( - client: SandboxClient, - sandbox_id: str, - program: ConfigData, - task: Task, - state: State, - runtime: Runtime, - *, - items: list[ProgramValue], - error_prefix: str, - use_sandbox_python_path: bool = False, -) -> None: - env = await command_env(program, task, state, runtime, include_base=False) - timeout = int_config(program, "setup_timeout", 300) - for command in items: - command = await resolve_program_value(command, task, state, runtime, program) - command = str(command) - if use_sandbox_python_path: - command = sandbox_python_path_command(command) - result = await maybe_call_with_named_args( - client.execute_command, - sandbox_id=sandbox_id, - command=command, - env=env, - timeout=timeout, - ) - if result.exit_code: - raise SandboxError(f"{error_prefix}: {result.stderr}") - - -async def read_sandbox_artifact( - client: SandboxClient, sandbox_id: str, path: str -) -> str: - script = ( - "import glob, pathlib, sys\n" - f"matches = sorted(glob.glob({path!r}))\n" - "if not matches:\n" - " sys.exit(2)\n" - "sys.stdout.write(pathlib.Path(matches[0]).read_text())\n" - ) - command = ( - "PYTHON=$(command -v python3 || command -v python || true); " - 'if [ -z "$PYTHON" ]; then ' - "echo 'python is required to read sandbox artifacts' >&2; exit 127; " - "fi; " - f'exec "$PYTHON" -c {shlex.quote(script)}' - ) - command = sandbox_python_path_command(command) - result = await maybe_call_with_named_args( - client.execute_command, - sandbox_id=sandbox_id, - command=command, - ) - if result.exit_code == 2: - raise FileNotFoundError(f"Sandbox artifact not found: {path}") - if result.exit_code: - raise SandboxError( - f"Sandbox artifact reader failed: {result.stderr or result.stdout or ''}" - ) - return result.stdout or "" diff --git a/verifiers/v1/utils/scoring_utils.py b/verifiers/v1/utils/scoring_utils.py index 8b288c171a..49dc654cda 100644 --- a/verifiers/v1/utils/scoring_utils.py +++ b/verifiers/v1/utils/scoring_utils.py @@ -11,17 +11,16 @@ from typing_extensions import TypedDict +from verifiers.types import Messages from verifiers.utils.async_utils import maybe_call_with_named_args -from .binding_utils import ROLLOUT_FRAMEWORK_ARGS, function_name -from ..state import State +from ..runtime import Runtime +from ..state import State, Turn from ..task import Task -from ..types import ConfigData, Handler, RuntimeData +from ..types import Context, Handler, JsonData, ModelClient SignalKind = Literal["metric", "reward", "advantage"] SignalStage = Literal["rollout", "group"] -ScoringConfig = dict[str, ConfigData] -SIGNAL_CONFIG_KEYS = {"stage", "priority", "weight", "skip"} class SignalRecord(TypedDict): @@ -33,9 +32,28 @@ class SignalRecord(TypedDict): weight: float +SignalKwarg = ( + Task + | State + | list[Task] + | list[State] + | list[Turn] + | Messages + | JsonData + | dict[str, float] + | float + | int + | str + | ModelClient + | Runtime + | Context + | None +) +SignalKwargs = dict[str, SignalKwarg] + + def build_signals( owner: object | None = None, - scoring: ScoringConfig | None = None, metrics: Iterable[Handler] | None = None, rewards: Iterable[Handler] | None = None, advantages: Iterable[Handler] | None = None, @@ -50,7 +68,6 @@ def build_signals( add_reward(signals, fn) for fn in advantages or (): add_advantage(signals, fn) - apply_scoring_config(signals, scoring or {}) return sorted(signals, key=signal_sort_key) @@ -59,7 +76,7 @@ def collect_signals(*signal_lists: Iterable[SignalRecord]) -> list[SignalRecord] seen: set[str] = set() for signal_list in signal_lists: for signal in signal_list: - name = cast(str, signal["name"]) + name = signal["name"] if name in seen: raise ValueError(f"Signal {name!r} is defined twice.") seen.add(name) @@ -83,6 +100,10 @@ async def score_rollout( signals: Iterable[SignalRecord], task: Task, state: State, + runtime: Runtime | None = None, + model_client: ModelClient | None = None, + teacher: ModelClient | None = None, + context: Context | None = None, resolve_kwargs: Callable[ [ Handler, @@ -90,33 +111,41 @@ async def score_rollout( State, set[str], ], - Awaitable[RuntimeData], + Awaitable[SignalKwargs], ] | None = None, ) -> State: start_time = time.time() - reward = float_value(state.get("reward"), 0.0) - metrics = dict(cast(dict[str, float], state.get("metrics") or {})) - framework_kwargs = rollout_framework_kwargs(task, state) + reward = float(state.reward) + metrics = dict(state.metrics) + framework_kwargs = rollout_framework_kwargs( + task, + state, + runtime=runtime, + model_client=model_client, + teacher=teacher, + context=context, + ) protected_args = set(framework_kwargs) for signal in sorted(signals, key=signal_sort_key): if signal["stage"] != "rollout": continue - extra_kwargs: RuntimeData = {} + extra_kwargs: SignalKwargs = {} if resolve_kwargs is not None: extra_kwargs = await resolve_kwargs( - cast(Handler, signal["fn"]), + signal["fn"], task, state, protected_args, ) value = await call_rollout_signal(signal, framework_kwargs, extra_kwargs) - metrics[cast(str, signal["name"])] = value + metrics[signal["name"]] = value if signal["kind"] == "reward": - reward += value * cast(float, signal["weight"]) - state["metrics"] = metrics - state["reward"] = reward - state.record_scoring_timing(start_time) + reward += value * signal["weight"] + state.metrics = metrics + state.reward = reward + state.timing.scoring.start = start_time + state.timing.scoring.end = time.time() return state @@ -124,6 +153,8 @@ async def score_group( signals: Iterable[SignalRecord], tasks: list[Task], states: list[State], + model_client: ModelClient | None = None, + teacher: ModelClient | None = None, resolve_kwargs: Callable[ [ Handler, @@ -131,14 +162,16 @@ async def score_group( list[State], set[str], ], - Awaitable[RuntimeData], + Awaitable[SignalKwargs], ] | None = None, ) -> list[State]: start_time = time.time() - rewards = [float_value(state.get("reward"), 0.0) for state in states] + rewards = [float(state.reward) for state in states] advantage_signals: list[SignalRecord] = [] - framework_kwargs = group_framework_kwargs(tasks, states) + framework_kwargs = group_framework_kwargs( + tasks, states, model_client=model_client, teacher=teacher + ) protected_args = set(framework_kwargs) for signal in sorted(signals, key=signal_sort_key): if signal["stage"] != "group": @@ -146,70 +179,48 @@ async def score_group( if signal["kind"] == "advantage": advantage_signals.append(signal) continue - extra_kwargs: RuntimeData = {} + extra_kwargs: SignalKwargs = {} if resolve_kwargs is not None: extra_kwargs = await resolve_kwargs( - cast(Handler, signal["fn"]), + signal["fn"], tasks, states, protected_args, ) values = await call_group_signal(signal, framework_kwargs, extra_kwargs) for index, value in enumerate(values): - metrics = dict(cast(dict[str, float], states[index].get("metrics") or {})) - metrics[cast(str, signal["name"])] = value - states[index]["metrics"] = metrics + metrics = dict(states[index].metrics) + metrics[signal["name"]] = value + states[index].metrics = metrics if signal["kind"] == "reward": - rewards[index] += value * cast(float, signal["weight"]) - advantages: list[float] | None = None + rewards[index] += value * signal["weight"] + for index, state in enumerate(states): + state.reward = rewards[index] for signal in advantage_signals: - extra_kwargs: RuntimeData = {} + extra_kwargs: SignalKwargs = {} if resolve_kwargs is not None: extra_kwargs = await resolve_kwargs( - cast(Handler, signal["fn"]), + signal["fn"], tasks, states, protected_args, ) - advantages = await call_group_signal(signal, framework_kwargs, extra_kwargs) + await call_group_advantage_signal(signal, framework_kwargs, extra_kwargs) for index, state in enumerate(states): - state["reward"] = rewards[index] - if advantages is not None: - state["advantage"] = advantages[index] - apply_advantage_to_trajectory(state, advantages[index]) - state.record_scoring_timing(start_time) + state.reward = rewards[index] + state.timing.scoring.start = start_time + state.timing.scoring.end = time.time() return states def add_signal(signals: MutableSequence[SignalRecord], signal: SignalRecord) -> None: - name = cast(str, signal["name"]) + name = signal["name"] if any(existing["name"] == name for existing in signals): raise ValueError(f"Signal {name!r} is defined twice.") validate_signal(signal) signals.append(signal) -def apply_scoring_config( - signals: MutableSequence[SignalRecord], scoring: ScoringConfig -) -> None: - by_name = {cast(str, signal["name"]): signal for signal in signals} - for name, config in scoring.items(): - validate_signal_config(name, config) - if bool_config(config, "skip", default=False): - if name not in by_name: - raise ValueError(f"Cannot skip unknown signal {name!r}.") - signals.remove(by_name[name]) - del by_name[name] - continue - if name not in by_name: - raise ValueError(f"Config references unknown signal {name!r}.") - signal = apply_signal_config(by_name[name], config) - validate_signal(signal) - index = signals.index(by_name[name]) - signals[index] = signal - by_name[name] = signal - - def decorated_signals(owner: object) -> list[SignalRecord]: signals: list[SignalRecord] = [] for _, method in inspect.getmembers(owner, predicate=callable): @@ -235,7 +246,12 @@ def signal_from_function(fn: Handler, kind: SignalKind | None = None) -> SignalR f"Signal function {function_name(fn)!r} must be decorated or given a kind." ) priority = int(getattr(fn, f"{resolved_kind}_priority", 0)) - stage = cast(SignalStage, getattr(fn, f"{resolved_kind}_stage", "rollout")) + raw_stage = getattr(fn, f"{resolved_kind}_stage", "rollout") + if raw_stage not in ("rollout", "group"): + raise ValueError( + f"Signal function {function_name(fn)!r} has invalid stage {raw_stage!r}." + ) + stage: SignalStage = raw_stage weight = 0.0 if resolved_kind == "reward": weight = float(getattr(fn, "reward_weight", 1.0)) @@ -249,30 +265,6 @@ def signal_from_function(fn: Handler, kind: SignalKind | None = None) -> SignalR } -def apply_signal_config(signal: SignalRecord, config: ConfigData) -> SignalRecord: - kind = cast(SignalKind, signal["kind"]) - stage = get_optional_stage(config) or cast(SignalStage, signal["stage"]) - priority_value = get_optional_number(config, "priority") - priority = cast(int, signal["priority"]) - if priority_value is not None: - priority = int(priority_value) - weight = cast(float, signal["weight"]) - if kind in {"metric", "advantage"}: - weight = 0.0 - else: - weight_value = get_optional_number(config, "weight") - if weight_value is not None: - weight = float(weight_value) - return { - "fn": signal["fn"], - "name": signal["name"], - "kind": kind, - "stage": stage, - "priority": priority, - "weight": weight, - } - - def decorated_kind(fn: Handler) -> SignalKind | None: has_metric = bool(getattr(fn, "metric", False)) has_reward = bool(getattr(fn, "reward", False)) @@ -289,8 +281,7 @@ def decorated_kind(fn: Handler) -> SignalKind | None: def validate_signal(signal: SignalRecord) -> None: - fn = cast(Handler, signal["fn"]) - inspect.signature(fn) + inspect.signature(signal["fn"]) if signal["stage"] == "rollout": if signal["kind"] == "advantage": raise ValueError( @@ -300,10 +291,10 @@ def validate_signal(signal: SignalRecord) -> None: async def call_rollout_signal( signal: SignalRecord, - framework_kwargs: RuntimeData, - extra_kwargs: RuntimeData | None = None, + framework_kwargs: SignalKwargs, + extra_kwargs: SignalKwargs | None = None, ) -> float: - fn = cast(Handler, signal["fn"]) + fn = signal["fn"] kwargs = {**dict(extra_kwargs or {}), **dict(framework_kwargs)} validate_required_kwargs(fn, kwargs, signal_context(signal)) value = await maybe_call_with_named_args(fn, **kwargs) @@ -312,18 +303,23 @@ async def call_rollout_signal( async def call_group_signal( signal: SignalRecord, - framework_kwargs: RuntimeData, - extra_kwargs: RuntimeData | None = None, + framework_kwargs: SignalKwargs, + extra_kwargs: SignalKwargs | None = None, ) -> list[float]: - fn = cast(Handler, signal["fn"]) + fn = signal["fn"] kwargs = {**dict(extra_kwargs or {}), **dict(framework_kwargs)} validate_required_kwargs(fn, kwargs, signal_context(signal)) value = await maybe_call_with_named_args(fn, **kwargs) - name = cast(str, signal["name"]) + name = signal["name"] if not isinstance(value, Sequence) or isinstance(value, str | bytes): raise TypeError(f"Group signal {name!r} must return a list of floats.") values = [float(item) for item in value] - states = cast(list[State], framework_kwargs["states"]) + states_value = framework_kwargs["states"] + if not isinstance(states_value, list) or not all( + isinstance(state, State) for state in states_value + ): + raise TypeError("Group signal framework kwargs must include states.") + states = states_value if len(values) != len(states): raise ValueError( f"Group signal {name!r} returned {len(values)} values for " @@ -332,21 +328,70 @@ async def call_group_signal( return values -def rollout_framework_kwargs(task: Task, state: State) -> RuntimeData: - kwargs: RuntimeData = {"task": task, "state": state} - for name in sorted(ROLLOUT_FRAMEWORK_ARGS - {"task", "state"}): - if name in state: - kwargs[name] = state[name] - elif name in task: - kwargs[name] = task[name] +async def call_group_advantage_signal( + signal: SignalRecord, + framework_kwargs: SignalKwargs, + extra_kwargs: SignalKwargs | None = None, +) -> None: + fn = signal["fn"] + kwargs = {**dict(extra_kwargs or {}), **dict(framework_kwargs)} + validate_required_kwargs(fn, kwargs, signal_context(signal)) + value = await maybe_call_with_named_args(fn, **kwargs) + if value is not None: + raise TypeError( + f"Group advantage signal {signal['name']!r} must mutate states in " + f"place and return None, not {type(value).__name__}." + ) + + +def rollout_framework_kwargs( + task: Task, + state: State, + *, + runtime: Runtime | None = None, + model_client: ModelClient | None = None, + teacher: ModelClient | None = None, + context: Context | None = None, +) -> SignalKwargs: + kwargs: SignalKwargs = { + "task": task, + "state": state, + "extras": state.extras, + "transcript": state.transcript, + "completion": state.completion, + "metrics": state.metrics, + "reward": state.reward, + "prompt": state.prompt if state.transcript else task.prompt, + "example_id": task.row_id, + "model": model_client, + "model_name": model_client.config.model if model_client is not None else None, + "teacher": teacher, + "teacher_name": teacher.config.model if teacher is not None else None, + "context": context, + } + if runtime is not None: + kwargs["runtime"] = runtime return kwargs -def group_framework_kwargs(tasks: list[Task], states: list[State]) -> RuntimeData: - return {"tasks": tasks, "states": states} +def group_framework_kwargs( + tasks: list[Task], + states: list[State], + *, + model_client: ModelClient | None = None, + teacher: ModelClient | None = None, +) -> SignalKwargs: + return { + "tasks": tasks, + "states": states, + "model": model_client, + "model_name": model_client.config.model if model_client is not None else None, + "teacher": teacher, + "teacher_name": teacher.config.model if teacher is not None else None, + } -def validate_required_kwargs(fn: Handler, kwargs: RuntimeData, context: str) -> None: +def validate_required_kwargs(fn: Handler, kwargs: SignalKwargs, context: str) -> None: signature = inspect.signature(fn) missing: list[str] = [] for parameter in signature.parameters.values(): @@ -369,70 +414,17 @@ def signal_context(signal: SignalRecord) -> str: return f"{signal['kind']} signal {signal['name']!r}" -def validate_signal_config(name: str, config: ConfigData) -> None: - unknown_keys = set(config) - SIGNAL_CONFIG_KEYS - if unknown_keys: - unknown = ", ".join(sorted(unknown_keys)) - raise ValueError(f"Signal config {name!r} has unknown keys: {unknown}.") - - -def get_optional_str(config: ConfigData, key: str) -> str | None: - value = config.get(key) - if value is None: - return None - if not isinstance(value, str): - raise TypeError(f"Signal config key {key!r} must be a string.") - return value - - -def get_optional_stage(config: ConfigData) -> SignalStage | None: - value = get_optional_str(config, "stage") - if value is None: - return None - if value not in {"rollout", "group"}: - raise ValueError("Signal stage must be 'rollout' or 'group'.") - return cast(SignalStage, value) - - -def get_optional_number(config: ConfigData, key: str) -> int | float | None: - value = config.get(key) - if value is None: - return None - if not isinstance(value, int | float): - raise TypeError(f"Signal config key {key!r} must be a number.") - return value - - -def bool_config(config: ConfigData, key: str, default: bool) -> bool: - value = config.get(key, default) - if not isinstance(value, bool): - raise TypeError(f"Signal config key {key!r} must be a boolean.") - return value - - -def float_value(value: object, default: float = 0.0) -> float: - if value is None: - return default - if isinstance(value, bool) or not isinstance(value, int | float | str): - return default - return float(value or 0.0) - - -def apply_advantage_to_trajectory(state: State, advantage: float) -> None: - trajectory = state.get("trajectory", []) - if not isinstance(trajectory, list): - return - for step in trajectory: - if isinstance(step, dict): - step = cast(ConfigData, step) - if step.get("advantage") is None: - step["advantage"] = advantage - - def signal_sort_key(signal: SignalRecord) -> tuple[int, str, str, str]: return ( - -cast(int, signal["priority"]), - cast(str, signal["name"]), - cast(str, signal["kind"]), - cast(str, signal["stage"]), + -signal["priority"], + signal["name"], + signal["kind"], + signal["stage"], ) + + +def function_name(fn: Handler) -> str: + name = getattr(fn, "__name__", None) + if isinstance(name, str) and name: + return name + return type(fn).__name__ diff --git a/verifiers/v1/utils/serialization_utils.py b/verifiers/v1/utils/serialization_utils.py deleted file mode 100644 index ae0173f318..0000000000 --- a/verifiers/v1/utils/serialization_utils.py +++ /dev/null @@ -1,11 +0,0 @@ -def serializable(value: object) -> object: - model_dump = getattr(value, "model_dump", None) - if callable(model_dump): - return model_dump(exclude_none=True) - if isinstance(value, list): - return [serializable(item) for item in value] - if isinstance(value, tuple): - return [serializable(item) for item in value] - if isinstance(value, dict): - return {str(key): serializable(item) for key, item in value.items()} - return value diff --git a/verifiers/v1/utils/taskset_utils.py b/verifiers/v1/utils/taskset_utils.py index ef404f4637..50dc49f1a9 100644 --- a/verifiers/v1/utils/taskset_utils.py +++ b/verifiers/v1/utils/taskset_utils.py @@ -1,107 +1,83 @@ import importlib import importlib.resources as resources -import json -import uuid from collections.abc import Iterable from contextlib import suppress -from copy import deepcopy from importlib.abc import Traversable from pathlib import Path -from typing import cast from datasets import Dataset -from verifiers.types import task_payload_from_info from ..task import Task from ..types import JsonData, Tasks -from .serialization_utils import serializable - - -def task_from_dataset_record(record: JsonData, taskset_id: str) -> Task: - record_data = serializable(record) - assert isinstance(record_data, dict) - record_json = cast(JsonData, record_data) - serialized_task = task_payload_from_info(record_json.get("info")) - if serialized_task is not None: - record_json = serialized_task - data = deepcopy(dict(record_json)) - if "prompt" not in data: - question = data.get("question") - data["prompt"] = ( - [{"role": "user", "content": str(question)}] if question is not None else [] - ) - return prepare_task(Task(cast(JsonData, data)), taskset_id) - - -def prepare_task(task: Task, taskset_id: str) -> Task: +from .json_utils import json_data + + +def task_from_dataset_record(record: JsonData, task_type: type[Task] = Task) -> Task: + return prepare_task(task_type.model_validate(json_data(record)), task_type) + + +def prepare_task(task: Task, task_type: type[Task] = Task) -> Task: if not isinstance(task, Task): raise TypeError("v1 task loaders must return Task objects.") - prepared = Task(cast(JsonData, dict(task))) - prepared["taskset_id"] = taskset_id - if prepared.get("task_id") is not None: - prepared["task_id"] = str(prepared["task_id"]) - else: - prepared["task_id"] = uuid.uuid4().hex - return prepared.freeze() + if isinstance(task, task_type): + return task + data = json_data(task.model_dump(mode="json", exclude_none=True)) + return task_type.model_validate(data) def dataset_record_from_task( task: Task, - taskset_id: str, index: int, - record: JsonData | None = None, ) -> JsonData: - data = Task(cast(JsonData, dict(task))) - data["example_id"] = index - normalized = prepare_task(data, taskset_id) - task_payload = dict(normalized) - dataset_record = deepcopy(dict(record or {})) - dataset_record["prompt"] = task_payload["prompt"] - dataset_record["example_id"] = task_payload["example_id"] - info = dataset_record.get("info") - if not isinstance(info, dict): - info = {} - dataset_record["info"] = {**info, "task": json.dumps(task_payload)} - if "answer" in normalized: - dataset_record["answer"] = normalized["answer"] - return cast(JsonData, dataset_record) - - -def dataset_records_from_tasks( - tasks: Iterable[Task], taskset_id: str -) -> list[JsonData]: + data = json_data(task.model_dump(mode="json", exclude_none=True)) + data["row_id"] = index + normalized = prepare_task(type(task).model_validate(data), type(task)) + task_payload = json_data(normalized.model_dump(mode="json", exclude_none=True)) + task_payload["example_id"] = index + return task_payload + + +def dataset_records_from_tasks(tasks: Iterable[Task]) -> list[JsonData]: dataset_records: list[JsonData] = [] for index, task in enumerate(tasks): - dataset_records.append(dataset_record_from_task(task, taskset_id, index)) + dataset_records.append(dataset_record_from_task(task, index)) return dataset_records -def dataset_from_result(result: Tasks, taskset_id: str) -> Dataset: +def dataset_from_result(result: Tasks) -> Dataset: + return dataset_from_result_typed(result, Task) + + +def dataset_from_result_typed(result: Tasks, task_type: type[Task]) -> Dataset: if isinstance(result, Dataset): records: list[JsonData] = [] for index, record in enumerate(result): - row = cast(JsonData, dict(record)) + row = json_data(dict(record)) row["example_id"] = index - task = task_from_dataset_record(row, taskset_id) - records.append(dataset_record_from_task(task, taskset_id, index, row)) + task = task_from_dataset_record(row, task_type) + records.append(dataset_record_from_task(task, index)) return Dataset.from_list(records) - tasks = tasks_from_result(result, taskset_id) - return Dataset.from_list(dataset_records_from_tasks(tasks, taskset_id)) + tasks = tasks_from_result_typed(result, task_type) + return Dataset.from_list(dataset_records_from_tasks(tasks)) + + +def tasks_from_result(result: Tasks) -> list[Task]: + return tasks_from_result_typed(result, Task) -def tasks_from_result(result: Tasks, taskset_id: str) -> list[Task]: +def tasks_from_result_typed(result: Tasks, task_type: type[Task]) -> list[Task]: if isinstance(result, Dataset): return [ - task_from_dataset_record(cast(JsonData, dict(record)), taskset_id) + task_from_dataset_record(json_data(dict(record)), task_type) for record in result ] if isinstance(result, Iterable): tasks: list[Task] = [] for item in result: if isinstance(item, Task): - tasks.append(prepare_task(item, taskset_id)) + tasks.append(prepare_task(item, task_type)) elif isinstance(item, dict): - tasks.append(task_from_dataset_record(cast(JsonData, item), taskset_id)) + tasks.append(task_from_dataset_record(json_data(item), task_type)) else: raise TypeError( "Task loader iterables must contain Task objects or JSON task " diff --git a/verifiers/v1/utils/tool_utils.py b/verifiers/v1/utils/tool_utils.py deleted file mode 100644 index 6f6d39287d..0000000000 --- a/verifiers/v1/utils/tool_utils.py +++ /dev/null @@ -1,58 +0,0 @@ -import inspect -from typing import cast - -from verifiers.types import Tool -from verifiers.v1.state import State -from verifiers.v1.toolset import Toolset, tool_name -from ..types import ConfigData, RuntimeCallable - - -def load_tools_from_state(state: State) -> dict[str, RuntimeCallable]: - runtime = state._runtime() - task = runtime.task_for_state(state) - return runtime.tool_calls(task, state) - - -def tool_error_content(error: Exception) -> str: - return str(error) - - -def tool_visible(toolset: Toolset, name: str) -> bool: - if toolset.show is not None and name not in toolset.show: - return False - if toolset.hide is not None and name in toolset.hide: - return False - return True - - -def toolset_object_scope(toolset: Toolset) -> str: - if toolset.scope is not None: - return toolset.scope - return "rollout" if toolset.write else "global" - - -def schema_callable(tool: object, signature: inspect.Signature) -> RuntimeCallable: - def call_for_schema() -> None: - return None - - call_for_schema.__name__ = tool_name(tool) - call_for_schema.__doc__ = inspect.getdoc(tool) - setattr(call_for_schema, "__signature__", signature) - return call_for_schema - - -def tool_schema(tool: Tool, hidden_args: set[str]) -> Tool: - parameters = dict(tool.parameters) - properties = dict(cast(ConfigData, parameters.get("properties") or {})) - for arg_name in hidden_args: - properties.pop(arg_name, None) - parameters["properties"] = properties - required = parameters.get("required") - if isinstance(required, list): - parameters["required"] = [arg for arg in required if arg not in hidden_args] - return Tool( - name=tool.name, - description=tool.description, - parameters=parameters, - strict=tool.strict, - ) diff --git a/verifiers/v1/utils/toolset_utils.py b/verifiers/v1/utils/toolset_utils.py deleted file mode 100644 index 1b8296610e..0000000000 --- a/verifiers/v1/utils/toolset_utils.py +++ /dev/null @@ -1,217 +0,0 @@ -import inspect -from collections.abc import Iterable -from typing import TYPE_CHECKING, cast - -from pydantic import BaseModel -from verifiers.types import Tool - -from ..types import ConfigData, Handler -from .config_utils import coerce_config, resolved_config_data, resolve_config_object - -if TYPE_CHECKING: - from ..toolset import ToolEntry, Toolset - - -def flatten_toolsets( - toolsets: Iterable["ToolEntry"], apply_visibility: bool = False -) -> list["ToolEntry"]: - from ..toolset import Toolset - - flat: list["ToolEntry"] = [] - for item in toolsets: - if isinstance(item, Toolset): - tools = flatten_toolsets(item.tools, apply_visibility) - if apply_visibility and item.show is not None: - show = set(item.show) - tools = [tool for tool in tools if tool_name(tool) in show] - if apply_visibility and item.hide is not None: - hide = set(item.hide) - tools = [tool for tool in tools if tool_name(tool) not in hide] - flat.extend(tools) - else: - flat.append(item) - return flat - - -def iter_toolsets(toolsets: Iterable["ToolEntry"]) -> list["Toolset"]: - from ..toolset import Toolset - - groups: list["Toolset"] = [] - for item in toolsets: - if isinstance(item, Toolset): - groups.append(item) - groups.extend(iter_toolsets(item.tools)) - return groups - - -def normalize_toolsets(toolsets: Iterable["ToolEntry"]) -> list["Toolset"]: - return [normalize_toolset(toolset) for toolset in toolsets] - - -def collect_toolsets( - values: object, - config: object, -) -> tuple[list["Toolset"], dict[str, "Toolset"]]: - value_toolsets, value_named = normalize_toolset_collection(values) - config_toolsets, config_named = normalize_toolset_collection(config) - duplicate = set(value_named) & set(config_named) - if duplicate: - raise ValueError(f"Toolsets are defined twice: {sorted(duplicate)}.") - return [*value_toolsets, *config_toolsets], {**value_named, **config_named} - - -def normalize_toolset_collection( - value: object, -) -> tuple[list["Toolset"], dict[str, "Toolset"]]: - from ..toolset import ToolEntry - - if value is None: - return [], {} - if isinstance(value, dict): - named: dict[str, Toolset] = {} - for key, item in value.items(): - if not isinstance(key, str): - raise TypeError("Toolset names must be strings.") - if key in named: - raise ValueError(f"Toolset {key!r} is defined twice.") - named[key] = named_toolset(key, item) - return list(named.values()), named - if isinstance(value, str): - return [normalize_toolset(value)], {} - if not isinstance(value, Iterable): - return [normalize_toolset(value)], {} - return normalize_toolsets(cast(Iterable[ToolEntry], value)), {} - - -def named_toolset(name: str, value: object) -> "Toolset": - from ..toolset import Toolset - - value = resolve_config_object(value) - if isinstance(value, Toolset): - return value - if isinstance(value, BaseModel): - value = value.model_dump(exclude_none=True) - if isinstance(value, dict): - spec = cast(ConfigData, value) - if "fn" in spec: - return toolset_from_factory(name, spec) - return toolset_from_mapping(spec) - if callable(value): - return call_toolset_factory(name, cast(Handler, value), {}) - return normalize_toolset(value) - - -def toolset_from_factory(name: str, spec: ConfigData) -> "Toolset": - fn = resolve_config_object(spec.get("fn")) - if not callable(fn): - raise TypeError(f"Toolset {name!r} requires callable fn.") - kwargs = {key: value for key, value in spec.items() if key != "fn"} - return call_toolset_factory(name, cast(Handler, fn), kwargs) - - -def call_toolset_factory(name: str, fn: Handler, kwargs: ConfigData) -> "Toolset": - result = fn(**kwargs) - if inspect.isawaitable(result): - raise TypeError(f"Toolset {name!r} fn must be synchronous.") - toolsets = normalize_toolset_result(result) - if len(toolsets) != 1: - raise ValueError(f"Toolset {name!r} fn must return exactly one Toolset.") - return toolsets[0] - - -def normalize_toolset_result(value: object) -> list["Toolset"]: - from ..toolset import ToolEntry, Toolset - - value = resolve_config_object(value) - if value is None: - return [] - if isinstance(value, Toolset | dict | str): - return [normalize_toolset(value)] - if not isinstance(value, Iterable): - return [normalize_toolset(value)] - return normalize_toolsets(cast(Iterable[ToolEntry], value)) - - -def normalize_toolset(value: object) -> "Toolset": - from ..toolset import ToolEntry, Toolset - - value = resolve_config_object(value) - if isinstance(value, Toolset): - return value - if isinstance(value, dict): - return toolset_from_mapping(cast(ConfigData, value)) - return Toolset(tools=[cast(ToolEntry, value)]) - - -def toolset_from_mapping(spec: ConfigData) -> "Toolset": - from ..toolset import Toolset, ToolsetConfig - - extra_keys = set(spec) - set(ToolsetConfig.model_fields) - if extra_keys: - raise ValueError(f"Unknown toolset config keys: {sorted(extra_keys)}.") - return Toolset(config=ToolsetConfig.model_validate(spec)) - - -def tool_items(value: object) -> list["ToolEntry"]: - if value is None: - return [] - if isinstance(value, str) or isinstance(value, dict): - return [tool_item(value)] - if not isinstance(value, Iterable): - return [tool_item(value)] - return [tool_item(item) for item in value] - - -def tool_item(value: object) -> "ToolEntry": - from ..toolset import MCPTool, MCPToolConfig, Toolset - - value = resolve_config_object(value) - if isinstance(value, Toolset | MCPTool | Tool): - return value - if isinstance(value, MCPToolConfig): - return MCPTool( - command=value.command, - args=value.args, - env=value.env, - cwd=value.cwd, - ) - if isinstance(value, dict): - if "command" in value: - config = coerce_config(MCPToolConfig, cast(ConfigData, value)) - return MCPTool( - command=config.command, - args=config.args, - env=config.env, - cwd=config.cwd, - ) - raise TypeError("Tool mapping specs require command.") - if not callable(value): - raise TypeError( - "Tool entries must be callables, Tools, Toolsets, or MCP tool specs." - ) - return cast(Handler, value) - - -def toolset_config_mapping(config: BaseModel | ConfigData | None) -> ConfigData: - from ..toolset import ToolsetConfig - - if config is None: - return {} - return resolved_config_data(coerce_config(ToolsetConfig, config)) - - -def optional_string(value: object) -> str | None: - if value is None: - return None - if not isinstance(value, str): - raise TypeError("Toolset scope must be a string.") - return value - - -def tool_name(tool: object) -> str: - if isinstance(tool, Tool): - return tool.name - name = getattr(tool, "__name__", None) or getattr(tool, "name", None) - if not isinstance(name, str) or not name: - raise ValueError("Tools require a stable __name__ or name.") - return name diff --git a/verifiers/v1/utils/trajectory_utils.py b/verifiers/v1/utils/trajectory_utils.py deleted file mode 100644 index 99b78d7953..0000000000 --- a/verifiers/v1/utils/trajectory_utils.py +++ /dev/null @@ -1,82 +0,0 @@ -from collections.abc import Sequence -from typing import cast - -from ..state import State -from verifiers.types import Message - -from ..runtime_handles import ResolvedRuntimeHandlesConfig -from ..types import ConfigData, JsonData, PromptMessage - - -def sync_trajectory( - state: State, - trajectory: Sequence[JsonData] | None = None, -) -> State: - if trajectory is not None: - state["trajectory"] = [dict(step) for step in trajectory] - - steps = state.get("trajectory") or [] - if not isinstance(steps, list): - raise TypeError("state.trajectory must be a list.") - state["num_model_requests"] = len(steps) - state._set_truncated( - any(bool(cast(ConfigData, step).get("is_truncated", False)) for step in steps) - ) - - if not steps: - return state - - state["prompt"] = message_list(steps[0], "prompt") - state["completion"] = merge_existing_completion( - completion_from_trajectory(steps), state.get("completion") - ) - return state - - -def has_borrowed_trajectory(state: State) -> bool: - runtime = state.runtime_state() - resolved = ResolvedRuntimeHandlesConfig.model_validate( - runtime.get("resolved") or {} - ) - return resolved.trajectory is not None - - -def completion_from_trajectory(steps: Sequence[JsonData]) -> list[PromptMessage]: - if not steps: - return [] - first_prompt = message_list(steps[0], "prompt") - last_prompt = message_list(steps[-1], "prompt") - last_completion = message_list(steps[-1], "completion") - last_trace = [*last_prompt, *last_completion] - if last_trace[: len(first_prompt)] == first_prompt: - return last_trace[len(first_prompt) :] - return last_trace - - -def merge_existing_completion( - trajectory_completion: list[PromptMessage], existing: object -) -> list[PromptMessage]: - if not isinstance(existing, list): - return trajectory_completion - if existing[: len(trajectory_completion)] == trajectory_completion: - return [cast(PromptMessage, message) for message in existing] - return trajectory_completion - - -def message_list(step: object, field: str) -> list[PromptMessage]: - if not isinstance(step, dict): - raise TypeError("trajectory steps must be mappings.") - value = cast(ConfigData, step).get(field) - if value is None: - return [] - if not isinstance(value, list): - raise TypeError(f"trajectory step {field} must be a list.") - messages: list[PromptMessage] = [] - for item in value: - if isinstance(item, dict): - messages.append(cast(JsonData, item)) - elif hasattr(item, "role") and hasattr(item, "content"): - messages.append(cast(Message, item)) - else: - raise TypeError(f"trajectory step {field} items must be messages.") - return messages diff --git a/verifiers/v1/utils/usage_utils.py b/verifiers/v1/utils/usage_utils.py deleted file mode 100644 index 19895c7e65..0000000000 --- a/verifiers/v1/utils/usage_utils.py +++ /dev/null @@ -1,34 +0,0 @@ -from typing import cast - -from verifiers.types import Response -from verifiers.utils.usage_utils import usage_tokens - -from ..state import State - - -def record_response_usage(state: State, response: Response) -> None: - if response.usage is None: - return - input_tokens, output_tokens = usage_tokens(response.usage) - usage = state.setdefault("token_usage", {"input_tokens": 0.0, "output_tokens": 0.0}) - if not isinstance(usage, dict): - raise TypeError("state.token_usage must be a mapping.") - usage = cast(dict[str, float], usage) - usage["input_tokens"] = float(usage.get("input_tokens", 0.0)) + float(input_tokens) - usage["output_tokens"] = float(usage.get("output_tokens", 0.0)) + float( - output_tokens - ) - # Context ("final") token metrics, accumulated at write time from the live - # Response. v1 serializes trajectory responses to plain dicts, so they can't - # be recomputed from the trajectory afterward (the isinstance(Response) gate - # in compute_context_token_metrics fails). Mirror that helper's formula for a - # linear rollout: final_output is the running sum of completions; final_input - # is the latest step's full context minus that sum. - usage["final_output_tokens"] = float(usage.get("final_output_tokens", 0.0)) + float( - output_tokens - ) - last_step_total = float(input_tokens) + float(output_tokens) - usage["final_input_tokens"] = max( - 0.0, last_step_total - usage["final_output_tokens"] - ) - state["usage"] = usage