diff --git a/ergon_builtins/ergon_builtins/benchmarks/gdpeval/benchmark.py b/ergon_builtins/ergon_builtins/benchmarks/gdpeval/benchmark.py index 3924339c5..37c01b8f1 100644 --- a/ergon_builtins/ergon_builtins/benchmarks/gdpeval/benchmark.py +++ b/ergon_builtins/ergon_builtins/benchmarks/gdpeval/benchmark.py @@ -5,7 +5,7 @@ :class:`Benchmark` interface. """ -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from typing import Any, ClassVar from ergon_core.api import Benchmark, BenchmarkRequirements, Task @@ -92,22 +92,7 @@ def build_instances(self) -> Mapping[str, Sequence[Task[GDPTaskConfig]]]: All tasks land in a single ``"default"`` instance since there is no multi-instance structure in the GDP dataset. """ - tasks: list[Task[GDPTaskConfig]] = [] - for payload in self._load_task_configs(): - description = extract_task_description(payload.task_id, repo_id=self.dataset_repo) - tasks.append( - GDPEvalTask( - task_slug=payload.task_id, - instance_key="default", - description=description, - task_payload=payload, - worker=self._worker_factory(), - sandbox=self._sandbox_factory(), - evaluators=(self._evaluator_factory(),), - ) - ) - - return {"default": tasks} + return {"default": [sample.tasks[0] for sample in self._environment().all_samples()]} def evaluator_requirements(self) -> Sequence[str]: return () @@ -130,3 +115,29 @@ def _load_task_configs(self) -> list[GDPTaskConfig]: ) ) return configs + + # Compatibility loader surface for GDPEvalEnvironment. PR11 deletes + # benchmark-centered authoring and these wrappers. + def iter_rows(self) -> Iterator[GDPTaskConfig]: + yield from self._load_task_configs() + + def load_rows(self) -> Sequence[GDPTaskConfig]: + return self._load_task_configs() + + def _environment(self) -> Any: + from ergon_builtins.environments.gdpeval import GDPEvalEnvironment + + return GDPEvalEnvironment( + name=self.name, + dataset_repo=self.dataset_repo, + split=self.split, + limit=self.limit, + loader=self, + task_description=lambda row: extract_task_description( + row.task_id, + repo_id=self.dataset_repo, + ), + worker=lambda _row: self._worker_factory(), + sandbox=lambda _row: self._sandbox_factory(), + evaluators=lambda _row: (self._evaluator_factory(),), + ) diff --git a/ergon_builtins/ergon_builtins/benchmarks/minif2f/benchmark.py b/ergon_builtins/ergon_builtins/benchmarks/minif2f/benchmark.py index 986c8b875..dc1135f17 100644 --- a/ergon_builtins/ergon_builtins/benchmarks/minif2f/benchmark.py +++ b/ergon_builtins/ergon_builtins/benchmarks/minif2f/benchmark.py @@ -7,7 +7,7 @@ import json import logging -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from pathlib import Path from typing import Any, ClassVar @@ -76,33 +76,7 @@ def __init__( # ------------------------------------------------------------------ def build_instances(self) -> Mapping[str, Sequence[Task[MiniF2FTaskPayload]]]: - problems = self._load_problems() - tasks: list[Task[MiniF2FTaskPayload]] = [] - for problem in problems: - payload = MiniF2FTaskPayload( - name=problem.name, - informal_statement=problem.informal_statement, - formal_statement=problem.formal_statement, - header=problem.header, - ) - description = ( - f"{problem.informal_statement}\n\n" - f"Your task: prove the following theorem in Lean 4.\n\n" - f"{problem.header}\n" - f"{problem.formal_statement}" - ) - tasks.append( - MiniF2FTask( - task_slug=problem.name, - instance_key="default", - description=description, - task_payload=payload, - worker=self._worker_factory(), - sandbox=self._sandbox_factory(), - evaluators=(self._evaluator_factory(),), - ) - ) - return {"default": tasks} + return {"default": [sample.tasks[0] for sample in self._environment().all_samples()]} def evaluator_requirements(self) -> Sequence[str]: return () @@ -144,3 +118,24 @@ def _load_problems(self) -> list[MiniF2FProblem]: logger.info("Loaded %d MiniF2F-v2c problems", len(problems)) return problems + + # Compatibility loader surface for MiniF2FEnvironment. PR11 deletes + # benchmark-centered authoring and these wrappers. + def iter_rows(self) -> Iterator[MiniF2FProblem]: + yield from self._load_problems() + + def load_rows(self) -> Sequence[MiniF2FProblem]: + return self._load_problems() + + def _environment(self) -> Any: + from ergon_builtins.environments.minif2f import MiniF2FEnvironment + + return MiniF2FEnvironment( + name=self.name, + limit=self.limit, + data_dir=self.data_dir, + loader=self, + worker=lambda _row: self._worker_factory(), + sandbox=lambda _row: self._sandbox_factory(), + evaluators=lambda _row: (self._evaluator_factory(),), + ) diff --git a/ergon_builtins/ergon_builtins/benchmarks/researchrubrics/benchmark.py b/ergon_builtins/ergon_builtins/benchmarks/researchrubrics/benchmark.py index 76675d633..1e28e4fe3 100644 --- a/ergon_builtins/ergon_builtins/benchmarks/researchrubrics/benchmark.py +++ b/ergon_builtins/ergon_builtins/benchmarks/researchrubrics/benchmark.py @@ -4,7 +4,7 @@ whether agents know when and what to ask stakeholders. """ -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from typing import Any, ClassVar from datasets import load_dataset @@ -15,7 +15,6 @@ from ergon_core.core.shared.settings import settings from ergon_builtins.benchmarks.researchrubrics.sandbox import ResearchE2BSandbox -from ergon_builtins.benchmarks.researchrubrics.rubric import ResearchRubricsRubric from ergon_builtins.benchmarks.researchrubrics.task_schemas import ( ResearchRubricsTaskPayload, RubricCriterion, @@ -83,28 +82,7 @@ def __init__( # ------------------------------------------------------------------ def build_instances(self) -> Mapping[str, Sequence[Task[ResearchRubricsTaskPayload]]]: - payloads = self._load_rows() - tasks: list[Task[ResearchRubricsTaskPayload]] = [] - for payload in payloads: - evaluator = self._evaluator_factory() - if isinstance(evaluator, ResearchRubricsRubric) and not evaluator.rubric_criteria: - evaluator = ResearchRubricsRubric( - name=evaluator.name, - metadata=evaluator.metadata, - rubric_criteria=tuple(payload.rubrics), - ) - tasks.append( - ResearchRubricsTask( - task_slug=payload.sample_id, - instance_key="default", - description=payload.prompt, - task_payload=payload, - worker=self._worker_factory(), - sandbox=self._sandbox_factory(), - evaluators=(evaluator,), - ) - ) - return {"default": tasks} + return {"default": [sample.tasks[0] for sample in self._environment().all_samples()]} def evaluator_requirements(self) -> Sequence[str]: return () @@ -125,6 +103,27 @@ def _load_rows(self) -> list[ResearchRubricsTaskPayload]: return [_payload_from_row(train_ds[idx]) for idx in range(len(train_ds))] + # Compatibility loader surface for ResearchRubricsEnvironment. PR11 + # deletes benchmark-centered authoring and these wrappers. + def iter_rows(self) -> Iterator[ResearchRubricsTaskPayload]: + yield from self._load_rows() + + def load_rows(self) -> Sequence[ResearchRubricsTaskPayload]: + return self._load_rows() + + def _environment(self) -> Any: + from ergon_builtins.environments.researchrubrics import ResearchRubricsEnvironment + + return ResearchRubricsEnvironment( + name=self.name, + dataset_name=self.dataset_name, + limit=self.limit, + loader=self, + worker=lambda _row: self._worker_factory(), + sandbox=lambda _row: self._sandbox_factory(), + evaluators=lambda _row: (self._evaluator_factory(),), + ) + def _payload_from_row( row: Mapping[str, Any], # slopcop: ignore[no-typing-any] diff --git a/ergon_builtins/ergon_builtins/benchmarks/swebench_verified/benchmark.py b/ergon_builtins/ergon_builtins/benchmarks/swebench_verified/benchmark.py index 0ebe53239..2e661e447 100644 --- a/ergon_builtins/ergon_builtins/benchmarks/swebench_verified/benchmark.py +++ b/ergon_builtins/ergon_builtins/benchmarks/swebench_verified/benchmark.py @@ -6,7 +6,7 @@ """ import logging -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from typing import Any, ClassVar from datasets import load_dataset @@ -77,27 +77,33 @@ def __init__( self._evaluator_factory = evaluator_factory def build_instances(self) -> Mapping[str, Sequence[Task[SWEBenchTaskPayload]]]: - instances = _load_rows(limit=self.limit) - tasks: list[Task[SWEBenchTaskPayload]] = [] - for instance in instances: - payload = SWEBenchTaskPayload.from_instance(instance) - tasks.append( - SweBenchTask( - task_slug=instance.instance_id, - instance_key="default", - description=payload.build_worker_description(), - task_payload=payload, - worker=self._worker_factory(), - sandbox=self._sandbox_factory(), - evaluators=(self._evaluator_factory(),), - ) - ) - logger.info("Loaded %d SWE-Bench Verified instances", len(tasks)) - return {"default": tasks} + samples = self._environment().all_samples() + logger.info("Loaded %d SWE-Bench Verified instances", len(samples)) + return {"default": [sample.tasks[0] for sample in samples]} def evaluator_requirements(self) -> Sequence[str]: return () + # Compatibility loader surface for SweBenchVerifiedEnvironment. PR11 + # deletes benchmark-centered authoring and these wrappers. + def iter_rows(self) -> Iterator[SWEBenchInstance]: + yield from _load_rows(limit=self.limit) + + def load_rows(self) -> Sequence[SWEBenchInstance]: + return _load_rows(limit=self.limit) + + def _environment(self) -> Any: + from ergon_builtins.environments.swebench_verified import SweBenchVerifiedEnvironment + + return SweBenchVerifiedEnvironment( + name=self.name, + limit=self.limit, + loader=self, + worker=lambda _row: self._worker_factory(), + sandbox=lambda _row: self._sandbox_factory(), + evaluators=lambda _row: (self._evaluator_factory(),), + ) + def _load_rows(*, limit: int | None = None) -> list[SWEBenchInstance]: """Load and validate SWE-Bench instances from HuggingFace.""" diff --git a/ergon_builtins/ergon_builtins/environments/__init__.py b/ergon_builtins/ergon_builtins/environments/__init__.py new file mode 100644 index 000000000..3a58aaec9 --- /dev/null +++ b/ergon_builtins/ergon_builtins/environments/__init__.py @@ -0,0 +1,13 @@ +"""Builtin sample-producing environments.""" + +from ergon_builtins.environments.gdpeval import GDPEvalEnvironment +from ergon_builtins.environments.minif2f import MiniF2FEnvironment +from ergon_builtins.environments.researchrubrics import ResearchRubricsEnvironment +from ergon_builtins.environments.swebench_verified import SweBenchVerifiedEnvironment + +__all__ = [ + "GDPEvalEnvironment", + "MiniF2FEnvironment", + "ResearchRubricsEnvironment", + "SweBenchVerifiedEnvironment", +] diff --git a/ergon_builtins/ergon_builtins/environments/_resolve.py b/ergon_builtins/ergon_builtins/environments/_resolve.py new file mode 100644 index 000000000..0458f73c6 --- /dev/null +++ b/ergon_builtins/ergon_builtins/environments/_resolve.py @@ -0,0 +1,33 @@ +"""Typed component resolution helpers for builtin environments.""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from typing import TypeVar, cast + +from ergon_core.api.rubric import Evaluator +from ergon_core.api.sandbox import Sandbox +from ergon_core.api.worker import Worker + +RowT = TypeVar("RowT") + + +def resolve_worker(value: Worker | Callable[[RowT], Worker], row: RowT) -> Worker: + if isinstance(value, Worker): + return value + return cast(Callable[[RowT], Worker], value)(row) + + +def resolve_sandbox(value: Sandbox | Callable[[RowT], Sandbox], row: RowT) -> Sandbox: + if isinstance(value, Sandbox): + return value + return cast(Callable[[RowT], Sandbox], value)(row) + + +def resolve_evaluators( + value: Sequence[Evaluator] | Callable[[RowT], Sequence[Evaluator]], + row: RowT, +) -> tuple[Evaluator, ...]: + if callable(value): + return tuple(cast(Callable[[RowT], Sequence[Evaluator]], value)(row)) + return tuple(value) diff --git a/ergon_builtins/ergon_builtins/environments/gdpeval.py b/ergon_builtins/ergon_builtins/environments/gdpeval.py new file mode 100644 index 000000000..8ff092573 --- /dev/null +++ b/ergon_builtins/ergon_builtins/environments/gdpeval.py @@ -0,0 +1,133 @@ +"""GDPEval sample-producing environment.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterator, Sequence +from dataclasses import dataclass +from typing import Any, Literal, Protocol, cast + +from ergon_core.api import Environment, Sample, Task +from ergon_core.api.rubric import Evaluator +from ergon_core.api.sandbox import Sandbox +from ergon_core.api.worker import Worker +from pydantic import Field + +from ergon_builtins.benchmarks.gdpeval.benchmark import ( + GDPEvalTask, + _default_gdpeval_sandbox, +) +from ergon_builtins.benchmarks.gdpeval.loader import ( + HF_REPO_ID, + extract_task_description, + find_reference_files, + load_task_ids, +) +from ergon_builtins.benchmarks.gdpeval.task_schemas import GDPTaskConfig +from ergon_builtins.benchmarks.gdpeval.worker_factory import ( + make_gdpeval_rubric, + make_gdpeval_worker, +) +from ergon_builtins.environments._resolve import ( + resolve_evaluators, + resolve_sandbox, + resolve_worker, +) + + +class GDPEvalRowLoader(Protocol): + def iter_rows(self) -> Iterator[GDPTaskConfig]: ... + + def load_rows(self) -> Sequence[GDPTaskConfig]: ... + + +@dataclass(frozen=True) +class DefaultGDPEvalRowLoader: + dataset_repo: str = HF_REPO_ID + split: str = "train" + limit: int | None = None + + def iter_rows(self) -> Iterator[GDPTaskConfig]: + yield from self.load_rows() + + def load_rows(self) -> Sequence[GDPTaskConfig]: + configs: list[GDPTaskConfig] = [] + for task_id in load_task_ids(split=self.split, repo_id=self.dataset_repo, limit=self.limit): + ref_files = find_reference_files(task_id, repo_id=self.dataset_repo) + configs.append( + GDPTaskConfig( + task_id=task_id, + workflow_type="document_processing", + reference_files=[str(path) for path in ref_files], + ) + ) + return configs + + +class GDPEvalEnvironment(Environment): + name: str = "gdpeval" + dataset_repo: str = HF_REPO_ID + split: str = "train" + limit: int | None = None + source_mode: Literal["materialized"] = "materialized" + loader: Any | None = None + task_description: Callable[[GDPTaskConfig], str] | None = None + worker: Worker | Callable[[GDPTaskConfig], Worker] = Field(default_factory=make_gdpeval_worker) + evaluators: Sequence[Evaluator] | Callable[[GDPTaskConfig], Sequence[Evaluator]] = Field( + default_factory=lambda: (make_gdpeval_rubric(),) + ) + sandbox: Sandbox | Callable[[GDPTaskConfig], Sandbox] = Field( + default_factory=_default_gdpeval_sandbox + ) + + def iter_samples(self) -> Iterator[Sample]: + for row in self._limit_rows(self._loader().iter_rows()): + yield self._sample_from_row(row) + + def all_samples(self) -> Sequence[Sample]: + if self.source_mode == "streaming": + raise NotImplementedError("Streaming environments may not materialize all samples") + return [self._sample_from_row(row) for row in self._limit_rows(self._loader().load_rows())] + + def _loader(self) -> GDPEvalRowLoader: + return self.loader or DefaultGDPEvalRowLoader( + dataset_repo=self.dataset_repo, + split=self.split, + limit=self.limit, + ) + + def _limit_rows( + self, + rows: Sequence[GDPTaskConfig] | Iterator[GDPTaskConfig], + ) -> Iterator[GDPTaskConfig]: + for index, row in enumerate(rows): + if self.limit is not None and index >= self.limit: + break + yield row + + def _sample_from_row(self, row: GDPTaskConfig) -> Sample: + task = GDPEvalTask( + task_slug=row.task_id, + instance_key="default", + description=self._description_for(row), + task_payload=row, + worker=resolve_worker(self.worker, row), + sandbox=resolve_sandbox(self.sandbox, row), + evaluators=resolve_evaluators(self.evaluators, row), + ) + return Sample.from_tasks( + name=f"{self.name}:{row.task_id}", + sample_key=row.task_id, + environment_name=self.name, + sample_ref={"task_id": row.task_id, "split": self.split}, + source_metadata={ + "provider": "ergon-builtin:gdpeval", + "dataset_repo": self.dataset_repo, + "split": self.split, + }, + tasks=[cast(Task, task)], + ) + + def _description_for(self, row: GDPTaskConfig) -> str: + if self.task_description is not None: + return self.task_description(row) + return extract_task_description(row.task_id, repo_id=self.dataset_repo) diff --git a/ergon_builtins/ergon_builtins/environments/minif2f.py b/ergon_builtins/ergon_builtins/environments/minif2f.py new file mode 100644 index 000000000..3858e5b34 --- /dev/null +++ b/ergon_builtins/ergon_builtins/environments/minif2f.py @@ -0,0 +1,137 @@ +"""MiniF2F sample-producing environment.""" + +from __future__ import annotations + +import json +from collections.abc import Callable, Iterator, Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Literal, Protocol, cast + +from ergon_core.api import Environment, Sample, Task +from ergon_core.api.rubric import Evaluator +from ergon_core.api.sandbox import Sandbox +from ergon_core.api.worker import Worker +from huggingface_hub import hf_hub_download +from pydantic import Field + +from ergon_builtins.benchmarks.minif2f.benchmark import HF_FILENAME, HF_REPO_ID, MiniF2FTask +from ergon_builtins.benchmarks.minif2f.sandbox import LeanSandbox +from ergon_builtins.benchmarks.minif2f.task_schemas import MiniF2FProblem, MiniF2FTaskPayload +from ergon_builtins.benchmarks.minif2f.worker_factory import ( + make_minif2f_rubric, + make_minif2f_worker, +) +from ergon_builtins.environments._resolve import ( + resolve_evaluators, + resolve_sandbox, + resolve_worker, +) + + +class MiniF2FRowLoader(Protocol): + def iter_rows(self) -> Iterator[MiniF2FProblem]: ... + + def load_rows(self) -> Sequence[MiniF2FProblem]: ... + + +@dataclass(frozen=True) +class DefaultMiniF2FRowLoader: + data_dir: Path | None = None + limit: int | None = None + + def iter_rows(self) -> Iterator[MiniF2FProblem]: + yield from self.load_rows() + + def load_rows(self) -> Sequence[MiniF2FProblem]: + cache_dir = str(self.data_dir) if self.data_dir else None + jsonl_path = Path( + hf_hub_download( + repo_id=HF_REPO_ID, + filename=HF_FILENAME, + repo_type="dataset", + cache_dir=cache_dir, + ) + ) + rows: list[MiniF2FProblem] = [] + with jsonl_path.open(encoding="utf-8") as f: + for line in f: + stripped = line.strip() + if not stripped: + continue + raw = json.loads(stripped) + rows.append( + MiniF2FProblem( + name=raw["name"], + informal_statement=raw["informal_statement"], + formal_statement=raw["formal_statement"], + header=raw["header"], + ) + ) + if self.limit is not None and len(rows) >= self.limit: + break + return rows + + +class MiniF2FEnvironment(Environment): + name: str = "minif2f" + split: str = "validation" + limit: int | None = None + source_mode: Literal["materialized"] = "materialized" + data_dir: Path | None = None + loader: Any | None = None + worker: Worker | Callable[[MiniF2FProblem], Worker] = Field(default_factory=make_minif2f_worker) + evaluators: Sequence[Evaluator] | Callable[[MiniF2FProblem], Sequence[Evaluator]] = Field( + default_factory=lambda: (make_minif2f_rubric(),) + ) + sandbox: Sandbox | Callable[[MiniF2FProblem], Sandbox] = Field(default_factory=LeanSandbox) + + def iter_samples(self) -> Iterator[Sample]: + for row in self._limit_rows(self._loader().iter_rows()): + yield self._sample_from_row(row) + + def all_samples(self) -> Sequence[Sample]: + if self.source_mode == "streaming": + raise NotImplementedError("Streaming environments may not materialize all samples") + return [self._sample_from_row(row) for row in self._limit_rows(self._loader().load_rows())] + + def _loader(self) -> MiniF2FRowLoader: + return self.loader or DefaultMiniF2FRowLoader(data_dir=self.data_dir, limit=self.limit) + + def _limit_rows( + self, + rows: Sequence[MiniF2FProblem] | Iterator[MiniF2FProblem], + ) -> Iterator[MiniF2FProblem]: + for index, row in enumerate(rows): + if self.limit is not None and index >= self.limit: + break + yield row + + def _sample_from_row(self, row: MiniF2FProblem) -> Sample: + payload = MiniF2FTaskPayload( + name=row.name, + informal_statement=row.informal_statement, + formal_statement=row.formal_statement, + header=row.header, + ) + task = MiniF2FTask( + task_slug=row.name, + instance_key="default", + description=( + f"{row.informal_statement}\n\n" + f"Your task: prove the following theorem in Lean 4.\n\n" + f"{row.header}\n{row.formal_statement}" + ), + task_payload=payload, + worker=resolve_worker(self.worker, row), + sandbox=resolve_sandbox(self.sandbox, row), + evaluators=resolve_evaluators(self.evaluators, row), + ) + return Sample.from_tasks( + name=f"{self.name}:{row.name}", + sample_key=row.name, + environment_name=self.name, + sample_ref={"problem_id": row.name, "split": self.split}, + source_metadata={"provider": "ergon-builtin:minif2f", "split": self.split}, + tasks=[cast(Task, task)], + ) diff --git a/ergon_builtins/ergon_builtins/environments/researchrubrics.py b/ergon_builtins/ergon_builtins/environments/researchrubrics.py new file mode 100644 index 000000000..2aa6d1b9f --- /dev/null +++ b/ergon_builtins/ergon_builtins/environments/researchrubrics.py @@ -0,0 +1,144 @@ +"""ResearchRubrics sample-producing environment.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterator, Mapping, Sequence +from dataclasses import dataclass +from typing import Any, Literal, Protocol, cast + +from datasets import load_dataset +from ergon_core.api import Environment, Sample, Task +from ergon_core.api.rubric import Evaluator +from ergon_core.api.sandbox import Sandbox +from ergon_core.api.worker import Worker +from ergon_core.core.shared.settings import settings +from pydantic import Field + +from ergon_builtins.benchmarks.researchrubrics.benchmark import ( + ResearchRubricsBenchmark, + ResearchRubricsTask, + _default_research_sandbox, + _payload_from_row, +) +from ergon_builtins.benchmarks.researchrubrics.rubric import ResearchRubricsRubric +from ergon_builtins.benchmarks.researchrubrics.task_schemas import ResearchRubricsTaskPayload +from ergon_builtins.benchmarks.researchrubrics.worker_factory import ( + make_research_rubric, + make_research_worker, +) +from ergon_builtins.environments._resolve import ( + resolve_evaluators, + resolve_sandbox, + resolve_worker, +) + + +class ResearchRubricsRowLoader(Protocol): + def iter_rows(self) -> Iterator[ResearchRubricsTaskPayload]: ... + + def load_rows(self) -> Sequence[ResearchRubricsTaskPayload]: ... + + +@dataclass(frozen=True) +class DefaultResearchRubricsRowLoader: + dataset_name: str = ResearchRubricsBenchmark.dataset_name + split: str = "train" + limit: int | None = None + + def iter_rows(self) -> Iterator[ResearchRubricsTaskPayload]: + yield from self.load_rows() + + def load_rows(self) -> Sequence[ResearchRubricsTaskPayload]: + ds = load_dataset(self.dataset_name, token=settings.hf_api_key) + split_ds = ds[self.split] + if self.limit is not None: + split_ds = split_ds.select(range(min(self.limit, len(split_ds)))) + return [_payload_from_row(split_ds[idx]) for idx in range(len(split_ds))] + + +class ResearchRubricsEnvironment(Environment): + name: str = "researchrubrics" + dataset_name: str = ResearchRubricsBenchmark.dataset_name + split: str = "train" + limit: int | None = None + source_mode: Literal["materialized"] = "materialized" + loader: Any | None = None + worker: Worker | Callable[[ResearchRubricsTaskPayload], Worker] = Field( + default_factory=make_research_worker + ) + evaluators: ( + Sequence[Evaluator] | Callable[[ResearchRubricsTaskPayload], Sequence[Evaluator]] + ) = Field(default_factory=lambda: (make_research_rubric(),)) + sandbox: Sandbox | Callable[[ResearchRubricsTaskPayload], Sandbox] = Field( + default_factory=_default_research_sandbox + ) + + def iter_samples(self) -> Iterator[Sample]: + for row in self._limit_rows(self._loader().iter_rows()): + yield self._sample_from_row(row) + + def all_samples(self) -> Sequence[Sample]: + if self.source_mode == "streaming": + raise NotImplementedError("Streaming environments may not materialize all samples") + return [self._sample_from_row(row) for row in self._limit_rows(self._loader().load_rows())] + + def _loader(self) -> ResearchRubricsRowLoader: + return self.loader or DefaultResearchRubricsRowLoader( + dataset_name=self.dataset_name, + split=self.split, + limit=self.limit, + ) + + def _limit_rows( + self, + rows: Sequence[ResearchRubricsTaskPayload] | Iterator[ResearchRubricsTaskPayload], + ) -> Iterator[ResearchRubricsTaskPayload]: + for index, row in enumerate(rows): + if self.limit is not None and index >= self.limit: + break + yield row + + def _sample_from_row(self, row: ResearchRubricsTaskPayload | Mapping[str, Any]) -> Sample: + payload = row if isinstance(row, ResearchRubricsTaskPayload) else _payload_from_row(row) + evaluators = resolve_evaluators(self.evaluators, payload) + evaluators = tuple( + self._bind_payload_rubric(evaluator, payload) for evaluator in evaluators + ) + task = ResearchRubricsTask( + task_slug=payload.sample_id, + instance_key="default", + description=payload.prompt, + task_payload=payload, + worker=resolve_worker(self.worker, payload), + sandbox=resolve_sandbox(self.sandbox, payload), + evaluators=evaluators, + ) + return Sample.from_tasks( + name=f"{self.name}:{payload.sample_id}", + sample_key=payload.sample_id, + environment_name=self.name, + sample_ref={ + "sample_id": payload.sample_id, + "domain": payload.domain, + "split": self.split, + }, + source_metadata={ + "provider": "ergon-builtin:researchrubrics", + "dataset_name": self.dataset_name, + "split": self.split, + }, + tasks=[cast(Task, task)], + ) + + @staticmethod + def _bind_payload_rubric( + evaluator: Evaluator, + payload: ResearchRubricsTaskPayload, + ) -> Evaluator: + if isinstance(evaluator, ResearchRubricsRubric) and not evaluator.rubric_criteria: + return ResearchRubricsRubric( + name=evaluator.name, + metadata=evaluator.metadata, + rubric_criteria=tuple(payload.rubrics), + ) + return evaluator diff --git a/ergon_builtins/ergon_builtins/environments/swebench_verified.py b/ergon_builtins/ergon_builtins/environments/swebench_verified.py new file mode 100644 index 000000000..e0d89d1f7 --- /dev/null +++ b/ergon_builtins/ergon_builtins/environments/swebench_verified.py @@ -0,0 +1,143 @@ +"""SWE-Bench Verified sample-producing environment.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterator, Sequence +from dataclasses import dataclass +from typing import Any, Literal, Protocol, cast + +from datasets import load_dataset +from ergon_core.api import Environment, Sample, Task +from ergon_core.api.rubric import Evaluator +from ergon_core.api.sandbox import Sandbox +from ergon_core.api.worker import Worker +from pydantic import Field, model_validator + +from ergon_builtins.benchmarks.swebench_verified.benchmark import ( + HF_DATASET_ID, + HF_SPLIT, + SweBenchTask, + _default_swebench_sandbox, +) +from ergon_builtins.benchmarks.swebench_verified.task_schemas import ( + SWEBenchInstance, + SWEBenchTaskPayload, +) +from ergon_builtins.benchmarks.swebench_verified.worker_factory import ( + make_swebench_rubric, + make_swebench_worker, +) +from ergon_builtins.environments._resolve import ( + resolve_evaluators, + resolve_sandbox, + resolve_worker, +) + + +class SweBenchVerifiedRowLoader(Protocol): + def iter_rows(self) -> Iterator[SWEBenchInstance]: ... + + def load_rows(self) -> Sequence[SWEBenchInstance]: ... + + +@dataclass(frozen=True) +class DefaultSweBenchVerifiedRowLoader: + dataset_id: str = HF_DATASET_ID + split: str = HF_SPLIT + limit: int | None = None + streaming: bool = False + + def iter_rows(self) -> Iterator[SWEBenchInstance]: + ds = load_dataset(self.dataset_id, split=self.split, streaming=self.streaming) + for index, row in enumerate(ds): + if self.limit is not None and index >= self.limit: + break + yield SWEBenchInstance.from_raw(row) + + def load_rows(self) -> Sequence[SWEBenchInstance]: + ds = load_dataset(self.dataset_id, split=self.split) + if self.limit is not None: + ds = ds.select(range(min(self.limit, len(ds)))) + return [SWEBenchInstance.from_raw(row) for row in ds] + + +class SweBenchVerifiedEnvironment(Environment): + name: str = "swebench-verified" + split: str = HF_SPLIT + dataset_id: str = HF_DATASET_ID + limit: int | None = None + source_mode: Literal["materialized", "streaming"] = "materialized" + streaming: bool = False + loader: Any | None = None + worker: Worker | Callable[[SWEBenchInstance], Worker] = Field( + default_factory=make_swebench_worker + ) + evaluators: Sequence[Evaluator] | Callable[[SWEBenchInstance], Sequence[Evaluator]] = Field( + default_factory=lambda: (make_swebench_rubric(),) + ) + sandbox: Sandbox | Callable[[SWEBenchInstance], Sandbox] = Field( + default_factory=_default_swebench_sandbox + ) + + @model_validator(mode="after") + def _sync_streaming_flag(self) -> "SweBenchVerifiedEnvironment": + if self.streaming: + self.source_mode = "streaming" + elif self.source_mode == "streaming": + self.streaming = True + return self + + def iter_samples(self) -> Iterator[Sample]: + for row in self._limit_rows(self._loader().iter_rows()): + yield self._sample_from_row(row) + + def all_samples(self) -> Sequence[Sample]: + if self.source_mode == "streaming": + raise NotImplementedError("Streaming environments may not materialize all samples") + return [self._sample_from_row(row) for row in self._limit_rows(self._loader().load_rows())] + + def _loader(self) -> SweBenchVerifiedRowLoader: + return self.loader or DefaultSweBenchVerifiedRowLoader( + dataset_id=self.dataset_id, + split=self.split, + limit=self.limit, + streaming=self.source_mode == "streaming", + ) + + def _limit_rows( + self, + rows: Sequence[SWEBenchInstance] | Iterator[SWEBenchInstance], + ) -> Iterator[SWEBenchInstance]: + for index, row in enumerate(rows): + if self.limit is not None and index >= self.limit: + break + yield row + + def _sample_from_row(self, row: SWEBenchInstance) -> Sample: + payload = SWEBenchTaskPayload.from_instance(row) + task = SweBenchTask( + task_slug=row.instance_id, + instance_key="default", + description=payload.build_worker_description(), + task_payload=payload, + worker=resolve_worker(self.worker, row), + sandbox=resolve_sandbox(self.sandbox, row), + evaluators=resolve_evaluators(self.evaluators, row), + ) + return Sample.from_tasks( + name=f"{self.name}:{row.instance_id}", + sample_key=row.instance_id, + environment_name=self.name, + sample_ref={ + "instance_id": row.instance_id, + "repo": row.repo, + "base_commit": row.base_commit, + "split": self.split, + }, + source_metadata={ + "provider": "ergon-builtin:swebench-verified", + "dataset_id": self.dataset_id, + "split": self.split, + }, + tasks=[cast(Task, task)], + ) diff --git a/ergon_builtins/tests/unit/environments/conftest.py b/ergon_builtins/tests/unit/environments/conftest.py new file mode 100644 index 000000000..18b3d7bc1 --- /dev/null +++ b/ergon_builtins/tests/unit/environments/conftest.py @@ -0,0 +1,285 @@ +from __future__ import annotations + +from collections.abc import Iterable, Iterator, Sequence +from dataclasses import dataclass + +import pytest + +from ergon_builtins.benchmarks.gdpeval.task_schemas import GDPTaskConfig +from ergon_builtins.benchmarks.minif2f.task_schemas import MiniF2FProblem +from ergon_builtins.benchmarks.researchrubrics.task_schemas import ( + ResearchRubricsTaskPayload, + RubricCriterion, +) +from ergon_builtins.benchmarks.swebench_verified.task_schemas import SWEBenchInstance +from ergon_core.api.benchmark import Task +from ergon_core.api.criterion import CriterionOutcome +from ergon_core.api.rubric import Evaluator, TaskEvaluationResult +from ergon_core.test_support.task_factory import TestSandbox, TestWorker + + +class TestEvaluator(Evaluator): + type_slug = "test-evaluator" + + def criteria_for(self, task: Task) -> Iterable: + return () + + def aggregate_task( + self, + task: Task, + criterion_results: Iterable[CriterionOutcome], + ) -> TaskEvaluationResult: + return TaskEvaluationResult( + task_slug=task.task_slug, + score=1.0, + passed=True, + evaluator_name=self.name, + criterion_results=list(criterion_results), + ) + + +@dataclass(frozen=True) +class FakeLoader: + rows: Sequence[object] + + def iter_rows(self) -> Iterator[object]: + yield from self.rows + + def load_rows(self) -> Sequence[object]: + return list(self.rows) + + +class FakeSweBenchRow(SWEBenchInstance): + needs_ui: bool = False + + +@pytest.fixture +def worker() -> TestWorker: + return TestWorker(name="worker", model="test:none") + + +@pytest.fixture +def robot_worker() -> TestWorker: + return TestWorker(name="robot", model="test:none") + + +@pytest.fixture +def researcher_worker() -> TestWorker: + return TestWorker(name="researcher", model="test:none") + + +@pytest.fixture +def evaluator() -> TestEvaluator: + return TestEvaluator(name="judge") + + +@pytest.fixture +def ui_eval() -> TestEvaluator: + return TestEvaluator(name="ui-judge") + + +@pytest.fixture +def patch_eval() -> TestEvaluator: + return TestEvaluator(name="patch-judge") + + +@pytest.fixture +def sandbox() -> TestSandbox: + return TestSandbox() + + +@pytest.fixture +def browser_sandbox() -> TestSandbox: + return TestSandbox(env={"mode": "browser"}) + + +@pytest.fixture +def repo_sandbox() -> TestSandbox: + return TestSandbox(env={"mode": "repo"}) + + +def fake_minif2f_rows() -> list[MiniF2FProblem]: + return [ + MiniF2FProblem( + name="mini-1", + informal_statement="Prove one equals one.", + formal_statement="theorem mini_1 : 1 = 1 := by", + header="import Mathlib\n", + ), + MiniF2FProblem( + name="mini-2", + informal_statement="Prove two equals two.", + formal_statement="theorem mini_2 : 2 = 2 := by", + header="import Mathlib\n", + ), + ] + + +def fake_swebench_rows() -> list[FakeSweBenchRow]: + return [ + FakeSweBenchRow( + instance_id="swe-1", + repo="org/repo", + base_commit="abcdef123456", + problem_statement="Fix the parser.", + version="1.0", + fail_to_pass=["tests/test_parser.py::test_fix"], + pass_to_pass=[], + environment_setup_commit="abcdef123456", + test_patch="diff --git a/tests/test_parser.py b/tests/test_parser.py\n", + ), + FakeSweBenchRow( + instance_id="swe-2", + repo="org/repo", + base_commit="abcdef123456", + problem_statement="Fix the lexer.", + version="1.0", + fail_to_pass=["tests/test_lexer.py::test_fix"], + pass_to_pass=[], + environment_setup_commit="abcdef123456", + test_patch="diff --git a/tests/test_lexer.py b/tests/test_lexer.py\n", + ), + ] + + +def fake_researchrubrics_rows() -> list[ResearchRubricsTaskPayload]: + return [ + ResearchRubricsTaskPayload( + sample_id="research-1", + domain="analysis", + prompt="Write a concise market analysis.", + rubrics=[ + RubricCriterion( + criterion="Includes a clear conclusion.", + axis="Communication Quality", + weight=1.0, + ) + ], + ), + ResearchRubricsTaskPayload( + sample_id="research-2", + domain="analysis", + prompt="Write a concise technical analysis.", + rubrics=[ + RubricCriterion( + criterion="Cites relevant evidence.", + axis="References & Citation Quality", + weight=1.0, + ) + ], + ), + ] + + +def fake_gdpeval_rows() -> list[GDPTaskConfig]: + return [ + GDPTaskConfig( + task_id="gdp-1", + workflow_type="document_processing", + reference_files=["/tmp/reference-1.pdf"], + ), + GDPTaskConfig( + task_id="gdp-2", + workflow_type="document_processing", + reference_files=["/tmp/reference-2.pdf"], + ), + ] + + +def make_fake_minif2f_environment( + *, + limit: int | None = None, + worker: TestWorker | None = None, + evaluators: Sequence[TestEvaluator] | None = None, + sandbox: TestSandbox | None = None, + rows: Sequence[MiniF2FProblem] | None = None, +): + from ergon_builtins.environments.minif2f import MiniF2FEnvironment + + return MiniF2FEnvironment( + limit=limit, + loader=FakeLoader(rows or fake_minif2f_rows()), + worker=worker or TestWorker(name="worker", model="test:none"), + evaluators=evaluators or [TestEvaluator(name="judge")], + sandbox=sandbox or TestSandbox(), + ) + + +def make_fake_swebench_environment( + *, + limit: int | None = None, + streaming: bool = False, + worker=None, + evaluators=None, + sandbox=None, + rows: Sequence[FakeSweBenchRow] | None = None, +): + from ergon_builtins.environments.swebench_verified import SweBenchVerifiedEnvironment + + return SweBenchVerifiedEnvironment( + limit=limit, + streaming=streaming, + loader=FakeLoader(rows or fake_swebench_rows()), + worker=worker or TestWorker(name="worker", model="test:none"), + evaluators=evaluators or [TestEvaluator(name="judge")], + sandbox=sandbox or TestSandbox(), + ) + + +def make_fake_researchrubrics_environment( + *, + limit: int | None = None, + worker: TestWorker | None = None, + evaluators: Sequence[TestEvaluator] | None = None, + sandbox: TestSandbox | None = None, + rows: Sequence[ResearchRubricsTaskPayload] | None = None, +): + from ergon_builtins.environments.researchrubrics import ResearchRubricsEnvironment + + return ResearchRubricsEnvironment( + limit=limit, + loader=FakeLoader(rows or fake_researchrubrics_rows()), + worker=worker or TestWorker(name="worker", model="test:none"), + evaluators=evaluators or [TestEvaluator(name="judge")], + sandbox=sandbox or TestSandbox(), + ) + + +def make_fake_gdpeval_environment( + *, + limit: int | None = None, + worker: TestWorker | None = None, + evaluators: Sequence[TestEvaluator] | None = None, + sandbox: TestSandbox | None = None, + rows: Sequence[GDPTaskConfig] | None = None, +): + from ergon_builtins.environments.gdpeval import GDPEvalEnvironment + + return GDPEvalEnvironment( + limit=limit, + loader=FakeLoader(rows or fake_gdpeval_rows()), + task_description=lambda row: f"Process {row.task_id}.", + worker=worker or TestWorker(name="worker", model="test:none"), + evaluators=evaluators or [TestEvaluator(name="judge")], + sandbox=sandbox or TestSandbox(), + ) + + +@pytest.fixture +def minif2f_environment_factory(): + return make_fake_minif2f_environment + + +@pytest.fixture +def swebench_environment_factory(): + return make_fake_swebench_environment + + +@pytest.fixture +def researchrubrics_environment_factory(): + return make_fake_researchrubrics_environment + + +@pytest.fixture +def gdpeval_environment_factory(): + return make_fake_gdpeval_environment diff --git a/ergon_builtins/tests/unit/environments/test_environment_exports.py b/ergon_builtins/tests/unit/environments/test_environment_exports.py new file mode 100644 index 000000000..d49982028 --- /dev/null +++ b/ergon_builtins/tests/unit/environments/test_environment_exports.py @@ -0,0 +1,26 @@ +from __future__ import annotations + + +def test_all_builtin_environment_exports_exist() -> None: + from ergon_builtins.environments import ( + GDPEvalEnvironment, + MiniF2FEnvironment, + ResearchRubricsEnvironment, + SweBenchVerifiedEnvironment, + ) + + assert MiniF2FEnvironment.__name__ == "MiniF2FEnvironment" + assert SweBenchVerifiedEnvironment.__name__ == "SweBenchVerifiedEnvironment" + assert ResearchRubricsEnvironment.__name__ == "ResearchRubricsEnvironment" + assert GDPEvalEnvironment.__name__ == "GDPEvalEnvironment" + + +def test_legacy_benchmark_composition_is_not_exported_from_catalog() -> None: + from ergon_builtins.benchmarks import catalog + + exported = set(dir(catalog)) + + assert "MiniF2FBenchmark" not in exported + assert "SweBenchVerifiedBenchmark" not in exported + assert "ResearchRubricsBenchmark" not in exported + assert "GDPEvalBenchmark" not in exported diff --git a/ergon_builtins/tests/unit/environments/test_gdpeval_environment.py b/ergon_builtins/tests/unit/environments/test_gdpeval_environment.py new file mode 100644 index 000000000..8e31bbf50 --- /dev/null +++ b/ergon_builtins/tests/unit/environments/test_gdpeval_environment.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +import json + +from ergon_builtins.benchmarks.gdpeval.benchmark import GDPEvalTask +from ergon_builtins.benchmarks.gdpeval.task_schemas import GDPTaskConfig +from ergon_core.api import Sample + + +def test_gdpeval_environment_returns_samples(gdpeval_environment_factory) -> None: + env = gdpeval_environment_factory(limit=1) + + sample = next(iter(env.iter_samples())) + + assert isinstance(sample, Sample) + assert sample.sample_key == "gdp-1" + assert sample.environment_name == env.name + assert sample.tasks + assert sample.tasks[0].worker is not None + assert sample.tasks[0].sandbox is not None + assert sample.tasks[0].evaluators + assert isinstance(sample.tasks[0], GDPEvalTask) + assert isinstance(sample.tasks[0].task_payload, GDPTaskConfig) + + +def test_gdpeval_materialized_environment_supports_all_samples(gdpeval_environment_factory) -> None: + env = gdpeval_environment_factory(limit=2) + + samples = env.all_samples() + + assert len(samples) == 2 + assert [sample.sample_key for sample in samples] == ["gdp-1", "gdp-2"] + + +def test_gdpeval_sample_provenance_is_json_safe(gdpeval_environment_factory) -> None: + sample = next(iter(gdpeval_environment_factory(limit=1).iter_samples())) + + json.dumps(sample.sample_ref) + json.dumps(sample.source_metadata) diff --git a/ergon_builtins/tests/unit/environments/test_minif2f_environment.py b/ergon_builtins/tests/unit/environments/test_minif2f_environment.py new file mode 100644 index 000000000..4325872c8 --- /dev/null +++ b/ergon_builtins/tests/unit/environments/test_minif2f_environment.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +import json + +from ergon_builtins.benchmarks.minif2f.benchmark import MiniF2FTask +from ergon_builtins.benchmarks.minif2f.task_schemas import MiniF2FTaskPayload +from ergon_core.api import Sample + + +def test_minif2f_environment_returns_samples(minif2f_environment_factory) -> None: + env = minif2f_environment_factory(limit=1) + + sample = next(iter(env.iter_samples())) + + assert isinstance(sample, Sample) + assert sample.sample_key == "mini-1" + assert sample.environment_name == env.name + assert sample.tasks + assert sample.tasks[0].worker is not None + assert sample.tasks[0].sandbox is not None + assert sample.tasks[0].evaluators + assert isinstance(sample.tasks[0], MiniF2FTask) + assert isinstance(sample.tasks[0].task_payload, MiniF2FTaskPayload) + + +def test_minif2f_materialized_environment_supports_all_samples(minif2f_environment_factory) -> None: + env = minif2f_environment_factory(limit=2) + + samples = env.all_samples() + + assert len(samples) == 2 + assert [sample.sample_key for sample in samples] == ["mini-1", "mini-2"] + + +def test_minif2f_sample_provenance_is_json_safe(minif2f_environment_factory) -> None: + sample = next(iter(minif2f_environment_factory(limit=1).iter_samples())) + + json.dumps(sample.sample_ref) + json.dumps(sample.source_metadata) diff --git a/ergon_builtins/tests/unit/environments/test_researchrubrics_environment.py b/ergon_builtins/tests/unit/environments/test_researchrubrics_environment.py new file mode 100644 index 000000000..f93e60f41 --- /dev/null +++ b/ergon_builtins/tests/unit/environments/test_researchrubrics_environment.py @@ -0,0 +1,43 @@ +from __future__ import annotations + +import json + +from ergon_builtins.benchmarks.researchrubrics.benchmark import ResearchRubricsTask +from ergon_builtins.benchmarks.researchrubrics.task_schemas import ResearchRubricsTaskPayload +from ergon_core.api import Sample + + +def test_researchrubrics_environment_returns_samples(researchrubrics_environment_factory) -> None: + env = researchrubrics_environment_factory(limit=1) + + sample = next(iter(env.iter_samples())) + + assert isinstance(sample, Sample) + assert sample.sample_key == "research-1" + assert sample.environment_name == env.name + assert sample.tasks + assert sample.tasks[0].worker is not None + assert sample.tasks[0].sandbox is not None + assert sample.tasks[0].evaluators + assert isinstance(sample.tasks[0], ResearchRubricsTask) + assert isinstance(sample.tasks[0].task_payload, ResearchRubricsTaskPayload) + + +def test_researchrubrics_materialized_environment_supports_all_samples( + researchrubrics_environment_factory, +) -> None: + env = researchrubrics_environment_factory(limit=2) + + samples = env.all_samples() + + assert len(samples) == 2 + assert [sample.sample_key for sample in samples] == ["research-1", "research-2"] + + +def test_researchrubrics_sample_provenance_is_json_safe( + researchrubrics_environment_factory, +) -> None: + sample = next(iter(researchrubrics_environment_factory(limit=1).iter_samples())) + + json.dumps(sample.sample_ref) + json.dumps(sample.source_metadata) diff --git a/ergon_builtins/tests/unit/environments/test_swebench_verified_environment.py b/ergon_builtins/tests/unit/environments/test_swebench_verified_environment.py new file mode 100644 index 000000000..f0b216d8b --- /dev/null +++ b/ergon_builtins/tests/unit/environments/test_swebench_verified_environment.py @@ -0,0 +1,113 @@ +from __future__ import annotations + +import json +from itertools import islice + +import pytest + +from ergon_builtins.benchmarks.swebench_verified.benchmark import SweBenchTask +from ergon_builtins.benchmarks.swebench_verified.task_schemas import ( + SWEBenchInstance, + SWEBenchTaskPayload, +) +from ergon_core.api import Sample + + +class FakeSweBenchRow(SWEBenchInstance): + needs_ui: bool = False + + +def test_swebench_environment_returns_samples(swebench_environment_factory) -> None: + env = swebench_environment_factory(limit=1) + + sample = next(iter(env.iter_samples())) + + assert isinstance(sample, Sample) + assert sample.sample_key == "swe-1" + assert sample.environment_name == env.name + assert sample.tasks + assert sample.tasks[0].worker is not None + assert sample.tasks[0].sandbox is not None + assert sample.tasks[0].evaluators + assert isinstance(sample.tasks[0], SweBenchTask) + assert isinstance(sample.tasks[0].task_payload, SWEBenchTaskPayload) + + +def test_swebench_streaming_environment_yields_samples_without_all_samples( + worker, + evaluator, + sandbox, + swebench_environment_factory, +) -> None: + env = swebench_environment_factory( + limit=2, + streaming=True, + worker=worker, + evaluators=[evaluator], + sandbox=sandbox, + ) + + samples = list(islice(env.iter_samples(), 2)) + + assert len(samples) == 2 + with pytest.raises(NotImplementedError): + env.all_samples() + + +def test_swebench_environment_resolves_row_dependent_components( + robot_worker, + researcher_worker, + ui_eval, + patch_eval, + browser_sandbox, + repo_sandbox, + swebench_environment_factory, +) -> None: + rows = [ + FakeSweBenchRow( + instance_id="ui", + repo="org/repo", + base_commit="abcdef123456", + problem_statement="Fix the UI.", + version="1.0", + fail_to_pass=["tests/test_ui.py::test_fix"], + pass_to_pass=[], + environment_setup_commit="abcdef123456", + test_patch="diff --git a/tests/test_ui.py b/tests/test_ui.py\n", + needs_ui=True, + ), + FakeSweBenchRow( + instance_id="patch", + repo="org/repo", + base_commit="abcdef123456", + problem_statement="Fix the patch.", + version="1.0", + fail_to_pass=["tests/test_patch.py::test_fix"], + pass_to_pass=[], + environment_setup_commit="abcdef123456", + test_patch="diff --git a/tests/test_patch.py b/tests/test_patch.py\n", + needs_ui=False, + ), + ] + env = swebench_environment_factory( + rows=rows, + worker=lambda row: robot_worker if row.needs_ui else researcher_worker, + evaluators=lambda row: [ui_eval] if row.needs_ui else [patch_eval], + sandbox=lambda row: browser_sandbox if row.needs_ui else repo_sandbox, + ) + + samples = list(env.iter_samples()) + + assert samples[0].tasks[0].worker is robot_worker + assert samples[0].tasks[0].evaluators == (ui_eval,) + assert samples[0].tasks[0].sandbox is browser_sandbox + assert samples[1].tasks[0].worker is researcher_worker + assert samples[1].tasks[0].evaluators == (patch_eval,) + assert samples[1].tasks[0].sandbox is repo_sandbox + + +def test_swebench_sample_provenance_is_json_safe(swebench_environment_factory) -> None: + sample = next(iter(swebench_environment_factory(limit=1).iter_samples())) + + json.dumps(sample.sample_ref) + json.dumps(sample.source_metadata)