From 3d572cddba1473ab0423aaba722d16c763e45363 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 14:52:22 +0100 Subject: [PATCH 1/7] Persist experiment candidate pools --- .../ergon_core/api/experiment/sample.py | 2 - .../application/experiments/candidate_pool.py | 180 ++++++++++++++ .../application/experiments/repositories.py | 77 ++++++ .../core/persistence/experiments/__init__.py | 15 ++ .../core/persistence/experiments/models.py | 80 +++++++ ergon_core/migrations/env.py | 1 + .../versions/00000000_initial_v2.py | 39 ++- .../00000001_add_experiment_persistence.py | 186 +++++++++++++++ .../test_experiment_persistence_roundtrip.py | 132 +++++++++++ .../test_application_domain_boundaries.py | 8 +- .../test_definition_identity_naming.py | 60 +++++ .../architecture/test_single_alembic_head.py | 5 +- .../experiments/test_candidate_pool.py | 222 ++++++++++++++++++ .../test_experiment_persistence.py | 115 +++++++++ 14 files changed, 1116 insertions(+), 6 deletions(-) create mode 100644 ergon_core/ergon_core/core/application/experiments/candidate_pool.py create mode 100644 ergon_core/ergon_core/core/application/experiments/repositories.py create mode 100644 ergon_core/ergon_core/core/persistence/experiments/__init__.py create mode 100644 ergon_core/ergon_core/core/persistence/experiments/models.py create mode 100644 ergon_core/migrations/versions/00000001_add_experiment_persistence.py create mode 100644 ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py create mode 100644 ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py create mode 100644 ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py diff --git a/ergon_core/ergon_core/api/experiment/sample.py b/ergon_core/ergon_core/api/experiment/sample.py index f93aef7bc..d2d836e58 100644 --- a/ergon_core/ergon_core/api/experiment/sample.py +++ b/ergon_core/ergon_core/api/experiment/sample.py @@ -7,8 +7,6 @@ from pydantic import BaseModel, ConfigDict, Field, JsonValue, model_validator from ergon_core.api.benchmark import Task - - class Sample(BaseModel): """A selected runnable unit containing concrete public ``Task`` objects.""" diff --git a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py new file mode 100644 index 000000000..1b5aba420 --- /dev/null +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -0,0 +1,180 @@ +"""Candidate-pool persistence for public experiment samples.""" + +from collections.abc import Sequence +from uuid import UUID, uuid4 + +from sqlmodel import Session, select + +from ergon_core.api.experiment.experiment import Experiment, ExperimentRef +from ergon_core.api.experiment.sample import Sample +from ergon_core.api.benchmark import Task +from ergon_core.core.persistence.experiments.models import ( + ExperimentEnvironmentRow, + ExperimentSamplePoolEntryRow, +) +from ergon_core.core.shared.utils import utcnow + + +class SampleCandidatePool: + def __init__( + self, + session: Session, + *, + max_duplicate_pulls_per_environment: int = 1_000, + ) -> None: + self._session = session + self._max_duplicate_pulls_per_environment = max_duplicate_pulls_per_environment + + def fill( + self, + *, + experiment: Experiment, + handle: ExperimentRef, + candidate_pool_size: int, + ) -> list[ExperimentSamplePoolEntryRow]: + entries = self._pending_unselected_entries(handle.experiment_id) + if len(entries) >= candidate_pool_size: + return entries[:candidate_pool_size] + + env_rows = { + environment.name: self._environment_row(handle.experiment_id, environment.name) + for environment in experiment.environments + } + known_keys = self._known_keys_by_environment(handle.experiment_id) + iterators = { + environment.name: iter(environment.iter_samples()) + for environment in experiment.environments + } + active_names = [environment.name for environment in experiment.environments] + duplicate_pulls = dict.fromkeys(active_names, 0) + + while active_names and len(entries) < candidate_pool_size: + for environment_name in list(active_names): + try: + sample = next(iterators[environment_name]) + except StopIteration: + active_names.remove(environment_name) + continue + + if sample.sample_key in known_keys.setdefault(environment_name, set()): + duplicate_pulls[environment_name] += 1 + if ( + duplicate_pulls[environment_name] + >= self._max_duplicate_pulls_per_environment + ): + active_names.remove(environment_name) + continue + + duplicate_pulls[environment_name] = 0 + env_row = env_rows[environment_name] + entries.append( + self._record_candidate( + handle=handle, + environment_id=env_row.id, + sample=sample, + ) + ) + known_keys[environment_name].add(sample.sample_key) + if len(entries) >= candidate_pool_size: + break + return entries + + def mark_selected( + self, + entries: Sequence[ExperimentSamplePoolEntryRow], + *, + sampler_invocation_id: UUID, + ) -> None: + selected_at = utcnow() + for entry in entries: + entry.selected = True + entry.selected_at = selected_at + entry.sampler_invocation_id = sampler_invocation_id + self._session.add(entry) + self._session.flush() + + def _pending_unselected_entries( + self, + experiment_id: UUID, + ) -> list[ExperimentSamplePoolEntryRow]: + return self._session.exec( + select(ExperimentSamplePoolEntryRow) + .where(ExperimentSamplePoolEntryRow.experiment_id == experiment_id) + .where(ExperimentSamplePoolEntryRow.selected.is_(False)) + .where(ExperimentSamplePoolEntryRow.discarded.is_(False)) + .order_by(ExperimentSamplePoolEntryRow.created_at, ExperimentSamplePoolEntryRow.id) + ).all() + + def _environment_row( + self, + experiment_id: UUID, + environment_name: str, + ) -> ExperimentEnvironmentRow: + return self._session.exec( + select(ExperimentEnvironmentRow) + .where(ExperimentEnvironmentRow.experiment_id == experiment_id) + .where(ExperimentEnvironmentRow.name == environment_name) + ).one() + + def _known_keys_by_environment(self, experiment_id: UUID) -> dict[str, set[str]]: + rows = self._session.exec( + select(ExperimentSamplePoolEntryRow, ExperimentEnvironmentRow.name) + .join( + ExperimentEnvironmentRow, + ExperimentSamplePoolEntryRow.environment_id == ExperimentEnvironmentRow.id, + ) + .where(ExperimentSamplePoolEntryRow.experiment_id == experiment_id) + ).all() + known: dict[str, set[str]] = {} + for entry, environment_name in rows: + known.setdefault(environment_name, set()).add(entry.sample_key) + return known + + def _record_candidate( + self, + *, + handle: ExperimentRef, + environment_id: UUID, + sample: Sample, + ) -> ExperimentSamplePoolEntryRow: + row = ExperimentSamplePoolEntryRow( + experiment_id=handle.experiment_id, + environment_id=environment_id, + sample_key=sample.sample_key, + sample_json=sample.model_dump(mode="json"), + ) + self._session.add(row) + self._session.flush() + return row + + +async def sample_from_pool_entry(entry: ExperimentSamplePoolEntryRow) -> Sample: + """Rehydrate retained candidate JSON into an authored, unmaterialized Sample.""" + + payload = dict(entry.sample_json) + task_snapshots = payload.pop("tasks", []) + tasks = [await _task_from_candidate_snapshot(task_json) for task_json in task_snapshots] + return Sample.from_tasks( + name=str(payload["name"]), + sample_key=str(payload["sample_key"]), + environment_name=str(payload["environment_name"]), + tasks=tasks, + sample_ref=payload.get("sample_ref") if isinstance(payload.get("sample_ref"), dict) else None, + source_metadata=( + payload.get("source_metadata") + if isinstance(payload.get("source_metadata"), dict) + else None + ), + metadata=payload.get("metadata") if isinstance(payload.get("metadata"), dict) else None, + ) + + +async def _task_from_candidate_snapshot(task_json: object) -> Task: + if not isinstance(task_json, dict): + raise ValueError(f"Candidate task snapshot must be an object, got {type(task_json).__name__}") + task = await Task.from_definition(task_json, task_id=uuid4()) + # Candidate-pool rows are not runtime materializations. We use the existing + # `_type` dispatch path to rebuild object-bound config, then clear the + # temporary id so materialization remains the only runtime-id boundary. + task._task_id = None + return task diff --git a/ergon_core/ergon_core/core/application/experiments/repositories.py b/ergon_core/ergon_core/core/application/experiments/repositories.py new file mode 100644 index 000000000..4c17b7b5a --- /dev/null +++ b/ergon_core/ergon_core/core/application/experiments/repositories.py @@ -0,0 +1,77 @@ +"""Application repository helpers for experiment persistence.""" + +from uuid import UUID + +from pydantic import JsonValue +from sqlmodel import Session, select + +from ergon_core.api.experiment.experiment import Experiment, ExperimentRef +from ergon_core.core.persistence.experiments.models import ( + ExperimentEnvironmentRow, + ExperimentRow, + ExperimentSamplerInvocationRow, +) + + +def persist_experiment(*, session: Session, experiment: Experiment) -> ExperimentRef: + experiment.validate() + row = ExperimentRow( + name=experiment.name, + description=experiment.description, + created_by=experiment.created_by, + metadata_json=experiment.metadata, + ) + session.add(row) + session.flush() + + for environment in experiment.environments: + session.add( + ExperimentEnvironmentRow( + experiment_id=row.id, + name=environment.name, + source_mode=environment.source_mode, + source_metadata_json=environment.source_metadata, + metadata_json=environment.metadata, + ) + ) + session.flush() + + return ExperimentRef( + experiment_id=row.id, + name=row.name, + environment_ids=_environment_ids_for_experiment(session, row.id), + created_at=row.created_at, + metadata=row.metadata_json, + ) + + +def record_sampler_invocation( + *, + session: Session, + experiment_ref: ExperimentRef, + sampler_name: str, + requested_k: int, + candidate_pool_size: int, + selected_count: int = 0, + sampler_config: dict[str, JsonValue] | None = None, +) -> ExperimentSamplerInvocationRow: + row = ExperimentSamplerInvocationRow( + experiment_id=experiment_ref.experiment_id, + sampler_name=sampler_name, + requested_k=requested_k, + candidate_pool_size=candidate_pool_size, + selected_count=selected_count, + sampler_config_json=dict(sampler_config or {}), + ) + session.add(row) + session.flush() + return row + + +def _environment_ids_for_experiment(session: Session, experiment_id: UUID) -> dict[str, UUID]: + rows = session.exec( + select(ExperimentEnvironmentRow).where( + ExperimentEnvironmentRow.experiment_id == experiment_id + ) + ).all() + return {row.name: row.id for row in rows} diff --git a/ergon_core/ergon_core/core/persistence/experiments/__init__.py b/ergon_core/ergon_core/core/persistence/experiments/__init__.py new file mode 100644 index 000000000..728e76ead --- /dev/null +++ b/ergon_core/ergon_core/core/persistence/experiments/__init__.py @@ -0,0 +1,15 @@ +"""Experiment provenance and candidate-pool persistence models.""" + +from ergon_core.core.persistence.experiments.models import ( + ExperimentEnvironmentRow, + ExperimentRow, + ExperimentSamplePoolEntryRow, + ExperimentSamplerInvocationRow, +) + +__all__ = [ + "ExperimentEnvironmentRow", + "ExperimentRow", + "ExperimentSamplePoolEntryRow", + "ExperimentSamplerInvocationRow", +] diff --git a/ergon_core/ergon_core/core/persistence/experiments/models.py b/ergon_core/ergon_core/core/persistence/experiments/models.py new file mode 100644 index 000000000..44afe7c74 --- /dev/null +++ b/ergon_core/ergon_core/core/persistence/experiments/models.py @@ -0,0 +1,80 @@ +"""Experiment provenance and candidate-pool tables.""" + +from datetime import datetime +from uuid import UUID, uuid4 + +from ergon_core.core.shared.json_types import JsonValue +from ergon_core.core.shared.utils import utcnow as _utcnow +from sqlalchemy import JSON, Boolean, Column, DateTime, UniqueConstraint +from sqlmodel import Field, SQLModel + +TZDateTime = DateTime(timezone=True) + + +class ExperimentRow(SQLModel, table=True): + __tablename__ = "experiments" + + id: UUID = Field(default_factory=uuid4, primary_key=True) + name: str = Field(index=True) + description: str | None = None + created_by: str | None = Field(default=None, index=True) + metadata_json: dict[str, JsonValue] = Field(default_factory=dict, sa_column=Column(JSON)) + created_at: datetime = Field(default_factory=_utcnow, sa_type=TZDateTime) + + +class ExperimentEnvironmentRow(SQLModel, table=True): + __tablename__ = "experiment_environments" + __table_args__ = (UniqueConstraint("experiment_id", "name"),) + + id: UUID = Field(default_factory=uuid4, primary_key=True) + experiment_id: UUID = Field(foreign_key="experiments.id", index=True) + name: str = Field(index=True) + source_mode: str = Field(index=True) + source_metadata_json: dict[str, JsonValue] = Field( + default_factory=dict, + sa_column=Column(JSON), + ) + metadata_json: dict[str, JsonValue] = Field(default_factory=dict, sa_column=Column(JSON)) + + +class ExperimentSamplerInvocationRow(SQLModel, table=True): + __tablename__ = "experiment_sampler_invocations" + + id: UUID = Field(default_factory=uuid4, primary_key=True) + experiment_id: UUID = Field(foreign_key="experiments.id", index=True) + sampler_name: str = Field(index=True) + requested_k: int + candidate_pool_size: int + selected_count: int = 0 + sampler_config_json: dict[str, JsonValue] = Field( + default_factory=dict, + sa_column=Column(JSON), + ) + created_at: datetime = Field(default_factory=_utcnow, sa_type=TZDateTime) + + +class ExperimentSamplePoolEntryRow(SQLModel, table=True): + __tablename__ = "experiment_sample_pool_entries" + __table_args__ = (UniqueConstraint("experiment_id", "environment_id", "sample_key"),) + + id: UUID = Field(default_factory=uuid4, primary_key=True) + experiment_id: UUID = Field(foreign_key="experiments.id", index=True) + environment_id: UUID = Field(foreign_key="experiment_environments.id", index=True) + sampler_invocation_id: UUID | None = Field( + default=None, + foreign_key="experiment_sampler_invocations.id", + index=True, + ) + sample_key: str = Field(index=True) + sample_json: dict[str, JsonValue] = Field(default_factory=dict, sa_column=Column(JSON)) + selected: bool = Field( + default=False, + sa_column=Column(Boolean, nullable=False, default=False, index=True), + ) + discarded: bool = Field( + default=False, + sa_column=Column(Boolean, nullable=False, default=False, index=True), + ) + created_at: datetime = Field(default_factory=_utcnow, sa_type=TZDateTime) + selected_at: datetime | None = Field(default=None, sa_type=TZDateTime) + discarded_at: datetime | None = Field(default=None, sa_type=TZDateTime) diff --git a/ergon_core/migrations/env.py b/ergon_core/migrations/env.py index 5648ba02f..0683b17b7 100644 --- a/ergon_core/migrations/env.py +++ b/ergon_core/migrations/env.py @@ -7,6 +7,7 @@ from logging.config import fileConfig import ergon_core.core.persistence.definitions.models +import ergon_core.core.persistence.experiments.models import ergon_core.core.persistence.graph.models import ergon_core.core.persistence.samples.models import ergon_core.core.persistence.telemetry.models diff --git a/ergon_core/migrations/versions/00000000_initial_v2.py b/ergon_core/migrations/versions/00000000_initial_v2.py index 969bd08ce..4271f2929 100644 --- a/ergon_core/migrations/versions/00000000_initial_v2.py +++ b/ergon_core/migrations/versions/00000000_initial_v2.py @@ -26,10 +26,45 @@ branch_labels = None depends_on = None +INITIAL_TABLES = ( + "experiment_definitions", + "experiment_definition_workers", + "experiment_definition_evaluators", + "experiment_definition_instances", + "experiment_definition_tasks", + "experiment_definition_task_dependencies", + "experiment_definition_task_assignments", + "experiment_definition_task_evaluators", + "samples", + "sample_graph_nodes", + "sample_graph_edges", + "sample_status_events", + "sample_task_events", + "sample_edge_events", + "sample_worker_events", + "sample_evaluator_events", + "sample_sandbox_events", + "sample_annotation_events", + "sample_context_events", + "sample_task_attempts", + "sample_resources", + "sample_task_evaluations", + "threads", + "thread_messages", + "rollout_batches", + "rollout_batch_sample_memberships", + "sandbox_command_wal_entries", + "sandbox_events", +) + def upgrade() -> None: - SQLModel.metadata.create_all(op.get_bind()) + bind = op.get_bind() + for table_name in INITIAL_TABLES: + SQLModel.metadata.tables[table_name].create(bind, checkfirst=True) def downgrade() -> None: - SQLModel.metadata.drop_all(op.get_bind()) + bind = op.get_bind() + for table_name in reversed(INITIAL_TABLES): + SQLModel.metadata.tables[table_name].drop(bind, checkfirst=True) diff --git a/ergon_core/migrations/versions/00000001_add_experiment_persistence.py b/ergon_core/migrations/versions/00000001_add_experiment_persistence.py new file mode 100644 index 000000000..ee720b5e7 --- /dev/null +++ b/ergon_core/migrations/versions/00000001_add_experiment_persistence.py @@ -0,0 +1,186 @@ +"""Add experiment persistence and candidate pool tables. + +Revision ID: 00000001 +Revises: 00000000 +Create Date: 2026-05-26 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + + +revision = "00000001" +down_revision = "00000000" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "experiments", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("name", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("description", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("created_by", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("metadata_json", sa.JSON(), nullable=True), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index(op.f("ix_experiments_created_by"), "experiments", ["created_by"]) + op.create_index(op.f("ix_experiments_name"), "experiments", ["name"]) + + op.create_table( + "experiment_environments", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("experiment_id", sa.Uuid(), nullable=False), + sa.Column("name", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("source_mode", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("source_metadata_json", sa.JSON(), nullable=True), + sa.Column("metadata_json", sa.JSON(), nullable=True), + sa.ForeignKeyConstraint(["experiment_id"], ["experiments.id"]), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("experiment_id", "name"), + ) + op.create_index( + op.f("ix_experiment_environments_experiment_id"), + "experiment_environments", + ["experiment_id"], + ) + op.create_index( + op.f("ix_experiment_environments_name"), + "experiment_environments", + ["name"], + ) + op.create_index( + op.f("ix_experiment_environments_source_mode"), + "experiment_environments", + ["source_mode"], + ) + + op.create_table( + "experiment_sampler_invocations", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("experiment_id", sa.Uuid(), nullable=False), + sa.Column("sampler_name", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("requested_k", sa.Integer(), nullable=False), + sa.Column("candidate_pool_size", sa.Integer(), nullable=False), + sa.Column("selected_count", sa.Integer(), nullable=False), + sa.Column("sampler_config_json", sa.JSON(), nullable=True), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.ForeignKeyConstraint(["experiment_id"], ["experiments.id"]), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index( + op.f("ix_experiment_sampler_invocations_experiment_id"), + "experiment_sampler_invocations", + ["experiment_id"], + ) + op.create_index( + op.f("ix_experiment_sampler_invocations_sampler_name"), + "experiment_sampler_invocations", + ["sampler_name"], + ) + + op.create_table( + "experiment_sample_pool_entries", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("experiment_id", sa.Uuid(), nullable=False), + sa.Column("environment_id", sa.Uuid(), nullable=False), + sa.Column("sampler_invocation_id", sa.Uuid(), nullable=True), + sa.Column("sample_key", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("sample_json", sa.JSON(), nullable=True), + sa.Column("selected", sa.Boolean(), nullable=False), + sa.Column("discarded", sa.Boolean(), nullable=False), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("selected_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("discarded_at", sa.DateTime(timezone=True), nullable=True), + sa.ForeignKeyConstraint(["environment_id"], ["experiment_environments.id"]), + sa.ForeignKeyConstraint(["experiment_id"], ["experiments.id"]), + sa.ForeignKeyConstraint( + ["sampler_invocation_id"], + ["experiment_sampler_invocations.id"], + ), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("experiment_id", "environment_id", "sample_key"), + ) + op.create_index( + op.f("ix_experiment_sample_pool_entries_discarded"), + "experiment_sample_pool_entries", + ["discarded"], + ) + op.create_index( + op.f("ix_experiment_sample_pool_entries_environment_id"), + "experiment_sample_pool_entries", + ["environment_id"], + ) + op.create_index( + op.f("ix_experiment_sample_pool_entries_experiment_id"), + "experiment_sample_pool_entries", + ["experiment_id"], + ) + op.create_index( + op.f("ix_experiment_sample_pool_entries_sample_key"), + "experiment_sample_pool_entries", + ["sample_key"], + ) + op.create_index( + op.f("ix_experiment_sample_pool_entries_sampler_invocation_id"), + "experiment_sample_pool_entries", + ["sampler_invocation_id"], + ) + op.create_index( + op.f("ix_experiment_sample_pool_entries_selected"), + "experiment_sample_pool_entries", + ["selected"], + ) + + +def downgrade() -> None: + op.drop_index( + op.f("ix_experiment_sample_pool_entries_selected"), + table_name="experiment_sample_pool_entries", + ) + op.drop_index( + op.f("ix_experiment_sample_pool_entries_sampler_invocation_id"), + table_name="experiment_sample_pool_entries", + ) + op.drop_index( + op.f("ix_experiment_sample_pool_entries_sample_key"), + table_name="experiment_sample_pool_entries", + ) + op.drop_index( + op.f("ix_experiment_sample_pool_entries_experiment_id"), + table_name="experiment_sample_pool_entries", + ) + op.drop_index( + op.f("ix_experiment_sample_pool_entries_environment_id"), + table_name="experiment_sample_pool_entries", + ) + op.drop_index( + op.f("ix_experiment_sample_pool_entries_discarded"), + table_name="experiment_sample_pool_entries", + ) + op.drop_table("experiment_sample_pool_entries") + op.drop_index( + op.f("ix_experiment_sampler_invocations_sampler_name"), + table_name="experiment_sampler_invocations", + ) + op.drop_index( + op.f("ix_experiment_sampler_invocations_experiment_id"), + table_name="experiment_sampler_invocations", + ) + op.drop_table("experiment_sampler_invocations") + op.drop_index( + op.f("ix_experiment_environments_source_mode"), + table_name="experiment_environments", + ) + op.drop_index(op.f("ix_experiment_environments_name"), table_name="experiment_environments") + op.drop_index( + op.f("ix_experiment_environments_experiment_id"), + table_name="experiment_environments", + ) + op.drop_table("experiment_environments") + op.drop_index(op.f("ix_experiments_name"), table_name="experiments") + op.drop_index(op.f("ix_experiments_created_by"), table_name="experiments") + op.drop_table("experiments") diff --git a/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py new file mode 100644 index 000000000..d812db7ad --- /dev/null +++ b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py @@ -0,0 +1,132 @@ +from collections.abc import Iterator +import os +from pathlib import Path +import subprocess +import sys +from typing import Literal +from uuid import uuid4 + +import pytest +from sqlmodel import Session, SQLModel, create_engine, select +from sqlalchemy import inspect + +import ergon_core.core.persistence.definitions.models # noqa: F401 +import ergon_core.core.persistence.samples.models # noqa: F401 +import ergon_core.core.persistence.telemetry.models # noqa: F401 +from ergon_core.api import Environment, Experiment, Sample +from ergon_core.core.application.experiments.repositories import persist_experiment +from ergon_core.core.persistence.experiments.models import ExperimentEnvironmentRow +from ergon_core.test_support.task_factory import task_with_id + +ROOT = Path(__file__).resolve().parents[4] + + +def make_sample(environment_name: str, key: str) -> Sample: + return Sample.from_tasks( + name=f"{environment_name}:{key}", + sample_key=key, + environment_name=environment_name, + tasks=[ + task_with_id( + uuid4(), + task_slug=f"solve-{environment_name}-{key}", + instance_key=key, + description=f"Solve {key}", + ) + ], + ) + + +class MaterializedEnvironment(Environment): + source_mode: Literal["materialized"] = "materialized" + + def iter_samples(self) -> Iterator[Sample]: + yield make_sample(self.name, "a") + + +@pytest.fixture +def sqlite_session() -> Iterator[Session]: + engine = create_engine("sqlite:///:memory:") + SQLModel.metadata.create_all(engine) + with Session(engine) as session: + yield session + + +@pytest.fixture +def two_env_experiment() -> Experiment: + return Experiment( + name="integration smoke", + environments=[ + MaterializedEnvironment( + name="mini-validation", + source_metadata={"dataset": "mini"}, + ), + MaterializedEnvironment( + name="swe-validation", + source_metadata={"dataset": "swe"}, + ), + ], + ) + + +def test_experiment_persistence_tables_roundtrip_json_and_fks( + sqlite_session: Session, + two_env_experiment: Experiment, +) -> None: + handle = persist_experiment(session=sqlite_session, experiment=two_env_experiment) + + rows = sqlite_session.exec( + select(ExperimentEnvironmentRow).where( + ExperimentEnvironmentRow.experiment_id == handle.experiment_id + ) + ).all() + + assert {row.name for row in rows} == {"mini-validation", "swe-validation"} + assert all(isinstance(row.source_metadata_json, dict) for row in rows) + + +def test_alembic_upgrade_head_creates_experiment_persistence_tables(tmp_path: Path) -> None: + db_path = tmp_path / "experiment-persistence.sqlite" + env = { + **os.environ, + "ERGON_DATABASE_URL": f"sqlite:///{db_path}", + } + + result = subprocess.run( + [sys.executable, "-m", "alembic", "-c", "alembic.ini", "upgrade", "head"], + cwd=ROOT / "ergon_core", + env=env, + capture_output=True, + text=True, + check=False, + ) + + assert result.returncode == 0, result.stderr + tables = set(inspect(create_engine(f"sqlite:///{db_path}")).get_table_names()) + assert { + "experiments", + "experiment_environments", + "experiment_sampler_invocations", + "experiment_sample_pool_entries", + }.issubset(tables) + + +def test_alembic_offline_postgres_sql_renders_experiment_persistence_tables() -> None: + env = { + **os.environ, + "ERGON_DATABASE_URL": "postgresql://ergon:ergon@localhost:5432/ergon", + } + + result = subprocess.run( + [sys.executable, "-m", "alembic", "-c", "alembic.ini", "upgrade", "head", "--sql"], + cwd=ROOT / "ergon_core", + env=env, + capture_output=True, + text=True, + check=False, + ) + + assert result.returncode == 0, result.stderr + assert "CREATE TABLE experiments" in result.stdout + assert "CREATE TABLE experiment_sample_pool_entries" in result.stdout + assert " UUID " in result.stdout diff --git a/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py b/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py index aded6f5ed..1de86a5b4 100644 --- a/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py +++ b/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py @@ -46,7 +46,13 @@ "events": {"base.py", "runtime.py"}, # Experiments exposes cross-domain application behavior through service.py. # These files are domain-internal implementation modules, not public subfacades. - "experiments": {"definition_writer.py", "handles.py", "launch.py"}, + "experiments": { + "candidate_pool.py", + "definition_writer.py", + "handles.py", + "launch.py", + "repositories.py", + }, "ports": {"dashboard.py", "resources.py"}, "resources": {"publishing.py"}, "runtime": { diff --git a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py index bba5fe1b1..0adf7030e 100644 --- a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py +++ b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py @@ -19,6 +19,63 @@ ) EXCLUDED_FILES = {Path(__file__).resolve()} EXPERIMENT_ID_PATTERN = re.compile(r"\b(?:experiment" r"_id|experiment" r"Id)\b") +ALLOWED_EXPERIMENT_ID_PATTERNS_BY_FILE = { + ROOT / "ergon_core" / "ergon_core" / "api" / "experiment" / "experiment.py": ( + re.compile(r"experiment_id: UUID"), + ), + ROOT / "ergon_core" / "ergon_core" / "api" / "experiment" / "sampling.py": ( + re.compile(r"experiment_id: UUID"), + ), + ROOT + / "ergon_core" + / "ergon_core" + / "core" + / "application" + / "experiments" + / "candidate_pool.py": ( + re.compile(r"handle\.experiment_id"), + re.compile(r"experiment_id: UUID"), + re.compile(r"experiment_id == experiment_id"), + ), + ROOT + / "ergon_core" + / "ergon_core" + / "core" + / "application" + / "experiments" + / "repositories.py": ( + re.compile(r"experiment_id=row\.id"), + re.compile(r"experiment_ref\.experiment_id"), + re.compile(r"experiment_id: UUID"), + re.compile(r"experiment_id == experiment_id"), + ), + ROOT / "ergon_core" / "ergon_core" / "core" / "persistence" / "experiments" / "models.py": ( + re.compile(r"experiment_id"), + ), + ROOT + / "ergon_core" + / "tests" + / "integration" + / "experiments" + / "test_experiment_persistence_roundtrip.py": ( + re.compile(r"experiment_id == handle\.experiment_id"), + ), + ROOT + / "ergon_core" + / "tests" + / "unit" + / "core" + / "application" + / "experiments" + / "test_experiment_persistence.py": ( + re.compile(r"experiment_id == handle\.experiment_id"), + re.compile(r"handle\.experiment_id"), + ), + ROOT / "ergon_core" / "tests" / "unit" / "api" / "test_sampler_contract.py": ( + re.compile(r"experiment_id=uuid4"), + re.compile(r"result\.experiment_id"), + ), +} def test_definition_identity_uses_definition_name() -> None: @@ -40,6 +97,9 @@ def test_definition_identity_uses_definition_name() -> None: continue for line_number, line in enumerate(text.splitlines(), start=1): if EXPERIMENT_ID_PATTERN.search(line): + allowed = ALLOWED_EXPERIMENT_ID_PATTERNS_BY_FILE.get(path.resolve(), ()) + if any(pattern.search(line) for pattern in allowed): + continue hits.append(f"{path.relative_to(ROOT)}:{line_number}: {line.strip()}") assert hits == [] diff --git a/ergon_core/tests/unit/architecture/test_single_alembic_head.py b/ergon_core/tests/unit/architecture/test_single_alembic_head.py index 3fbc7f1ea..39316985c 100644 --- a/ergon_core/tests/unit/architecture/test_single_alembic_head.py +++ b/ergon_core/tests/unit/architecture/test_single_alembic_head.py @@ -9,7 +9,10 @@ def test_v2_has_one_initial_migration() -> None: migrations = sorted(path.name for path in VERSIONS.glob("*.py")) - assert migrations == ["00000000_initial_v2.py"] + assert migrations == [ + "00000000_initial_v2.py", + "00000001_add_experiment_persistence.py", + ] def test_initial_migration_has_no_parent() -> None: diff --git a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py new file mode 100644 index 000000000..85acb9521 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -0,0 +1,222 @@ +from collections.abc import Iterator +from typing import Literal +from uuid import uuid4 + +import pytest +from sqlmodel import Session, SQLModel, create_engine, select + +import ergon_core.core.persistence.definitions.models # noqa: F401 +import ergon_core.core.persistence.samples.models # noqa: F401 +import ergon_core.core.persistence.telemetry.models # noqa: F401 +from ergon_core.api import Environment, Experiment, Sample +from ergon_core.core.application.experiments.candidate_pool import ( + SampleCandidatePool, + sample_from_pool_entry, +) +from ergon_core.core.application.experiments.repositories import ( + persist_experiment, + record_sampler_invocation, +) +from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow +from ergon_core.test_support.task_factory import task_with_id + + +def make_sample(environment_name: str, key: str) -> Sample: + return Sample.from_tasks( + name=f"{environment_name}:{key}", + sample_key=key, + environment_name=environment_name, + tasks=[ + task_with_id( + uuid4(), + task_slug=f"solve-{environment_name}-{key}", + instance_key=key, + description=f"Solve {key}", + ) + ], + sample_ref={"key": key}, + source_metadata={"source": environment_name}, + metadata={"difficulty": "small"}, + ) + + +class StreamingEnvironment(Environment): + source_mode: Literal["streaming"] = "streaming" + total: int + + def iter_samples(self) -> Iterator[Sample]: + for index in range(self.total): + yield make_sample(self.name, str(index)) + + +class CountedStreamingEnvironment(StreamingEnvironment): + pull_count: int = 0 + + def iter_samples(self) -> Iterator[Sample]: + while self.pull_count < self.total: + index = self.pull_count + self.pull_count += 1 + yield make_sample(self.name, str(index)) + + +class DuplicateStreamingEnvironment(Environment): + source_mode: Literal["streaming"] = "streaming" + pull_count: int = 0 + + def iter_samples(self) -> Iterator[Sample]: + while True: + self.pull_count += 1 + yield make_sample(self.name, "duplicate") + + +class MaterializedEnvironment(Environment): + source_mode: Literal["materialized"] = "materialized" + keys: tuple[str, ...] + + def iter_samples(self) -> Iterator[Sample]: + for key in self.keys: + yield make_sample(self.name, key) + + +@pytest.fixture +def session() -> Iterator[Session]: + engine = create_engine("sqlite:///:memory:") + SQLModel.metadata.create_all(engine) + with Session(engine) as session: + yield session + + +@pytest.fixture +def streamed_experiment() -> Experiment: + return Experiment( + name="streamed", + environments=[StreamingEnvironment(name="stream", total=20)], + ) + + +@pytest.fixture +def counted_streaming_experiment() -> Experiment: + return Experiment( + name="counted", + environments=[CountedStreamingEnvironment(name="stream", total=20)], + ) + + +def test_candidate_pool_retains_unselected_streamed_candidates( + session: Session, + streamed_experiment: Experiment, +) -> None: + handle = persist_experiment(session=session, experiment=streamed_experiment) + pool = SampleCandidatePool(session) + + candidates = pool.fill( + experiment=streamed_experiment, + handle=handle, + candidate_pool_size=8, + ) + selected = candidates[:3] + invocation = record_sampler_invocation( + session=session, + experiment_ref=handle, + sampler_name="sequential", + requested_k=3, + candidate_pool_size=8, + selected_count=3, + sampler_config={}, + ) + pool.mark_selected(selected, sampler_invocation_id=invocation.id) + + rows = session.exec(select(ExperimentSamplePoolEntryRow)).all() + assert len(rows) == 8 + assert sum(row.selected for row in rows) == 3 + assert sum(row.discarded for row in rows) == 0 + + +def test_candidate_pool_reuses_unselected_entries_before_advancing_stream( + session: Session, + counted_streaming_experiment: Experiment, +) -> None: + handle = persist_experiment(session=session, experiment=counted_streaming_experiment) + pool = SampleCandidatePool(session) + + first = pool.fill( + experiment=counted_streaming_experiment, + handle=handle, + candidate_pool_size=8, + ) + pool.mark_selected(first[:2], sampler_invocation_id=uuid4()) + session.commit() + + second = pool.fill( + experiment=counted_streaming_experiment, + handle=handle, + candidate_pool_size=8, + ) + + assert [entry.sample_key for entry in second[:6]] == [ + entry.sample_key for entry in first[2:] + ] + assert counted_streaming_experiment.environments[0].pull_count == 10 + + +def test_candidate_pool_round_robins_new_entries_across_environments( + session: Session, +) -> None: + experiment = Experiment( + name="balanced", + environments=[ + MaterializedEnvironment(name="mini-validation", keys=("a", "b", "c")), + MaterializedEnvironment(name="swe-validation", keys=("1", "2", "3")), + ], + ) + handle = persist_experiment(session=session, experiment=experiment) + pool = SampleCandidatePool(session) + + entries = pool.fill(experiment=experiment, handle=handle, candidate_pool_size=4) + + assert [entry.sample_json["environment_name"] for entry in entries] == [ + "mini-validation", + "swe-validation", + "mini-validation", + "swe-validation", + ] + assert entries[0].sample_json["sample_key"] == "a" + assert entries[0].sample_json["tasks"][0]["task_slug"] == "solve-mini-validation-a" + + +def test_candidate_pool_stops_after_duplicate_pull_budget( + session: Session, +) -> None: + duplicate_env = DuplicateStreamingEnvironment(name="stream") + experiment = Experiment(name="duplicates", environments=[duplicate_env]) + handle = persist_experiment(session=session, experiment=experiment) + pool = SampleCandidatePool(session, max_duplicate_pulls_per_environment=3) + + entries = pool.fill(experiment=experiment, handle=handle, candidate_pool_size=2) + + assert [entry.sample_key for entry in entries] == ["duplicate"] + assert duplicate_env.pull_count == 4 + + +@pytest.mark.asyncio +async def test_candidate_pool_entry_rehydrates_object_bound_sample_without_runtime_task_id( + session: Session, +) -> None: + experiment = Experiment( + name="rehydrate", + environments=[MaterializedEnvironment(name="mini-validation", keys=("a",))], + ) + handle = persist_experiment(session=session, experiment=experiment) + entry = SampleCandidatePool(session).fill( + experiment=experiment, + handle=handle, + candidate_pool_size=1, + )[0] + + sample = await sample_from_pool_entry(entry) + + assert sample.sample_key == "a" + assert sample.environment_name == "mini-validation" + assert sample.tasks[0].task_slug == "solve-mini-validation-a" + with pytest.raises(RuntimeError, match="not been materialized"): + _ = sample.tasks[0].task_id diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py new file mode 100644 index 000000000..31ea0b4a2 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py @@ -0,0 +1,115 @@ +from collections.abc import Iterator +from typing import Literal +from uuid import uuid4 + +import pytest +from sqlmodel import Session, SQLModel, create_engine, select + +import ergon_core.core.persistence.definitions.models # noqa: F401 +import ergon_core.core.persistence.samples.models # noqa: F401 +import ergon_core.core.persistence.telemetry.models # noqa: F401 +from ergon_core.api import Environment, Experiment, Sample +from ergon_core.core.application.experiments.repositories import ( + persist_experiment, + record_sampler_invocation, +) +from ergon_core.core.persistence.experiments.models import ExperimentEnvironmentRow +from ergon_core.core.persistence.samples.models import SampleStatusEventRow, SampleTaskEventRow +from ergon_core.core.persistence.telemetry.models import SampleRecord +from ergon_core.test_support.task_factory import task_with_id + + +def make_sample(environment_name: str, key: str) -> Sample: + return Sample.from_tasks( + name=f"{environment_name}:{key}", + sample_key=key, + environment_name=environment_name, + tasks=[ + task_with_id( + uuid4(), + task_slug=f"solve-{environment_name}-{key}", + instance_key=key, + description=f"Solve {key}", + ) + ], + sample_ref={"key": key}, + source_metadata={"source": environment_name}, + metadata={"difficulty": "small"}, + ) + + +class MaterializedEnvironment(Environment): + source_mode: Literal["materialized"] = "materialized" + + def iter_samples(self) -> Iterator[Sample]: + yield make_sample(self.name, "a") + + +@pytest.fixture +def session() -> Iterator[Session]: + engine = create_engine("sqlite:///:memory:") + SQLModel.metadata.create_all(engine) + with Session(engine) as session: + yield session + + +@pytest.fixture +def two_env_experiment() -> Experiment: + return Experiment( + name="persistence smoke", + description="Persist shape only", + created_by="unit-test", + metadata={"suite": "pr04"}, + environments=[ + MaterializedEnvironment( + name="mini-validation", + source_metadata={"dataset": "mini"}, + metadata={"split": "validation"}, + ), + MaterializedEnvironment( + name="swe-validation", + source_metadata={"dataset": "swe"}, + metadata={"split": "validation"}, + ), + ], + ) + + +def test_persist_experiment_writes_environment_rows( + session: Session, + two_env_experiment: Experiment, +) -> None: + handle = persist_experiment(session=session, experiment=two_env_experiment) + + env_rows = session.exec( + select(ExperimentEnvironmentRow).where( + ExperimentEnvironmentRow.experiment_id == handle.experiment_id + ) + ).all() + + assert handle.experiment_id is not None + assert [row.name for row in env_rows] == ["mini-validation", "swe-validation"] + assert session.exec(select(SampleRecord)).all() == [] + assert session.exec(select(SampleStatusEventRow)).all() == [] + assert session.exec(select(SampleTaskEventRow)).all() == [] + + +def test_record_sampler_invocation_writes_no_runtime_state( + session: Session, + two_env_experiment: Experiment, +) -> None: + handle = persist_experiment(session=session, experiment=two_env_experiment) + + invocation = record_sampler_invocation( + session=session, + experiment_ref=handle, + sampler_name="random", + requested_k=4, + candidate_pool_size=16, + selected_count=0, + sampler_config={"seed": 1}, + ) + + assert invocation.experiment_id == handle.experiment_id + assert session.exec(select(SampleRecord)).all() == [] + assert session.exec(select(SampleTaskEventRow)).all() == [] From 83b8c8eb1b0437d751993ba1c620daf25f4df763 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 16:59:37 +0100 Subject: [PATCH 2/7] Reduce experiment persistence test suppressions --- .../test_experiment_persistence_roundtrip.py | 11 ++++++++--- .../application/experiments/test_candidate_pool.py | 11 ++++++++--- .../experiments/test_experiment_persistence.py | 11 ++++++++--- 3 files changed, 24 insertions(+), 9 deletions(-) diff --git a/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py index d812db7ad..c9af3b838 100644 --- a/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py +++ b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py @@ -1,4 +1,5 @@ from collections.abc import Iterator +from importlib import import_module import os from pathlib import Path import subprocess @@ -10,9 +11,6 @@ from sqlmodel import Session, SQLModel, create_engine, select from sqlalchemy import inspect -import ergon_core.core.persistence.definitions.models # noqa: F401 -import ergon_core.core.persistence.samples.models # noqa: F401 -import ergon_core.core.persistence.telemetry.models # noqa: F401 from ergon_core.api import Environment, Experiment, Sample from ergon_core.core.application.experiments.repositories import persist_experiment from ergon_core.core.persistence.experiments.models import ExperimentEnvironmentRow @@ -20,6 +18,13 @@ ROOT = Path(__file__).resolve().parents[4] +for module_name in ( + "ergon_core.core.persistence.definitions.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + def make_sample(environment_name: str, key: str) -> Sample: return Sample.from_tasks( diff --git a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py index 85acb9521..45513041c 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -1,13 +1,11 @@ from collections.abc import Iterator +from importlib import import_module from typing import Literal from uuid import uuid4 import pytest from sqlmodel import Session, SQLModel, create_engine, select -import ergon_core.core.persistence.definitions.models # noqa: F401 -import ergon_core.core.persistence.samples.models # noqa: F401 -import ergon_core.core.persistence.telemetry.models # noqa: F401 from ergon_core.api import Environment, Experiment, Sample from ergon_core.core.application.experiments.candidate_pool import ( SampleCandidatePool, @@ -20,6 +18,13 @@ from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow from ergon_core.test_support.task_factory import task_with_id +for module_name in ( + "ergon_core.core.persistence.definitions.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + def make_sample(environment_name: str, key: str) -> Sample: return Sample.from_tasks( diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py index 31ea0b4a2..bde9e4243 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py @@ -1,13 +1,11 @@ from collections.abc import Iterator +from importlib import import_module from typing import Literal from uuid import uuid4 import pytest from sqlmodel import Session, SQLModel, create_engine, select -import ergon_core.core.persistence.definitions.models # noqa: F401 -import ergon_core.core.persistence.samples.models # noqa: F401 -import ergon_core.core.persistence.telemetry.models # noqa: F401 from ergon_core.api import Environment, Experiment, Sample from ergon_core.core.application.experiments.repositories import ( persist_experiment, @@ -18,6 +16,13 @@ from ergon_core.core.persistence.telemetry.models import SampleRecord from ergon_core.test_support.task_factory import task_with_id +for module_name in ( + "ergon_core.core.persistence.definitions.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + def make_sample(environment_name: str, key: str) -> Sample: return Sample.from_tasks( From d6fe976f80214718d1f4bb9377b69452314d99cd Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 22:39:25 +0100 Subject: [PATCH 3/7] Separate experiment pool repository access --- .../application/experiments/candidate_pool.py | 86 ++------ .../application/experiments/repositories.py | 187 ++++++++++++++---- .../experiments/test_candidate_pool.py | 10 + 3 files changed, 170 insertions(+), 113 deletions(-) diff --git a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py index 1b5aba420..52b6f7f9c 100644 --- a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -3,16 +3,15 @@ from collections.abc import Sequence from uuid import UUID, uuid4 -from sqlmodel import Session, select +from sqlmodel import Session from ergon_core.api.experiment.experiment import Experiment, ExperimentRef from ergon_core.api.experiment.sample import Sample from ergon_core.api.benchmark import Task -from ergon_core.core.persistence.experiments.models import ( - ExperimentEnvironmentRow, - ExperimentSamplePoolEntryRow, +from ergon_core.core.application.experiments.repositories import ( + ExperimentRepository, ) -from ergon_core.core.shared.utils import utcnow +from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow class SampleCandidatePool: @@ -21,8 +20,9 @@ def __init__( session: Session, *, max_duplicate_pulls_per_environment: int = 1_000, + repository: ExperimentRepository | None = None, ) -> None: - self._session = session + self._repository = repository or ExperimentRepository(session) self._max_duplicate_pulls_per_environment = max_duplicate_pulls_per_environment def fill( @@ -32,15 +32,18 @@ def fill( handle: ExperimentRef, candidate_pool_size: int, ) -> list[ExperimentSamplePoolEntryRow]: - entries = self._pending_unselected_entries(handle.experiment_id) + entries = self._repository.pending_unselected_pool_entries(handle.experiment_id) if len(entries) >= candidate_pool_size: return entries[:candidate_pool_size] env_rows = { - environment.name: self._environment_row(handle.experiment_id, environment.name) + environment.name: self._repository.experiment_environment_row( + experiment_id=handle.experiment_id, + environment_name=environment.name, + ) for environment in experiment.environments } - known_keys = self._known_keys_by_environment(handle.experiment_id) + known_keys = self._repository.known_sample_keys_by_environment(handle.experiment_id) iterators = { environment.name: iter(environment.iter_samples()) for environment in experiment.environments @@ -68,7 +71,7 @@ def fill( duplicate_pulls[environment_name] = 0 env_row = env_rows[environment_name] entries.append( - self._record_candidate( + self._repository.record_candidate( handle=handle, environment_id=env_row.id, sample=sample, @@ -85,67 +88,10 @@ def mark_selected( *, sampler_invocation_id: UUID, ) -> None: - selected_at = utcnow() - for entry in entries: - entry.selected = True - entry.selected_at = selected_at - entry.sampler_invocation_id = sampler_invocation_id - self._session.add(entry) - self._session.flush() - - def _pending_unselected_entries( - self, - experiment_id: UUID, - ) -> list[ExperimentSamplePoolEntryRow]: - return self._session.exec( - select(ExperimentSamplePoolEntryRow) - .where(ExperimentSamplePoolEntryRow.experiment_id == experiment_id) - .where(ExperimentSamplePoolEntryRow.selected.is_(False)) - .where(ExperimentSamplePoolEntryRow.discarded.is_(False)) - .order_by(ExperimentSamplePoolEntryRow.created_at, ExperimentSamplePoolEntryRow.id) - ).all() - - def _environment_row( - self, - experiment_id: UUID, - environment_name: str, - ) -> ExperimentEnvironmentRow: - return self._session.exec( - select(ExperimentEnvironmentRow) - .where(ExperimentEnvironmentRow.experiment_id == experiment_id) - .where(ExperimentEnvironmentRow.name == environment_name) - ).one() - - def _known_keys_by_environment(self, experiment_id: UUID) -> dict[str, set[str]]: - rows = self._session.exec( - select(ExperimentSamplePoolEntryRow, ExperimentEnvironmentRow.name) - .join( - ExperimentEnvironmentRow, - ExperimentSamplePoolEntryRow.environment_id == ExperimentEnvironmentRow.id, - ) - .where(ExperimentSamplePoolEntryRow.experiment_id == experiment_id) - ).all() - known: dict[str, set[str]] = {} - for entry, environment_name in rows: - known.setdefault(environment_name, set()).add(entry.sample_key) - return known - - def _record_candidate( - self, - *, - handle: ExperimentRef, - environment_id: UUID, - sample: Sample, - ) -> ExperimentSamplePoolEntryRow: - row = ExperimentSamplePoolEntryRow( - experiment_id=handle.experiment_id, - environment_id=environment_id, - sample_key=sample.sample_key, - sample_json=sample.model_dump(mode="json"), + self._repository.mark_pool_entries_selected( + list(entries), + sampler_invocation_id=sampler_invocation_id, ) - self._session.add(row) - self._session.flush() - return row async def sample_from_pool_entry(entry: ExperimentSamplePoolEntryRow) -> Sample: diff --git a/ergon_core/ergon_core/core/application/experiments/repositories.py b/ergon_core/ergon_core/core/application/experiments/repositories.py index 4c17b7b5a..95a81387e 100644 --- a/ergon_core/ergon_core/core/application/experiments/repositories.py +++ b/ergon_core/ergon_core/core/application/experiments/repositories.py @@ -6,43 +6,156 @@ from sqlmodel import Session, select from ergon_core.api.experiment.experiment import Experiment, ExperimentRef +from ergon_core.api.experiment.sample import Sample from ergon_core.core.persistence.experiments.models import ( ExperimentEnvironmentRow, ExperimentRow, ExperimentSamplerInvocationRow, + ExperimentSamplePoolEntryRow, ) +from ergon_core.core.shared.utils import utcnow -def persist_experiment(*, session: Session, experiment: Experiment) -> ExperimentRef: - experiment.validate() - row = ExperimentRow( - name=experiment.name, - description=experiment.description, - created_by=experiment.created_by, - metadata_json=experiment.metadata, - ) - session.add(row) - session.flush() - - for environment in experiment.environments: - session.add( - ExperimentEnvironmentRow( - experiment_id=row.id, - name=environment.name, - source_mode=environment.source_mode, - source_metadata_json=environment.source_metadata, - metadata_json=environment.metadata, +class ExperimentRepository: + """Data-access boundary for experiment authoring persistence.""" + + def __init__(self, session: Session) -> None: + self._session = session + + def persist_experiment(self, experiment: Experiment) -> ExperimentRef: + experiment.validate_authoring() + row = ExperimentRow( + name=experiment.name, + description=experiment.description, + created_by=experiment.created_by, + metadata_json=experiment.metadata, + ) + self._session.add(row) + self._session.flush() + + for environment in experiment.environments: + self._session.add( + ExperimentEnvironmentRow( + experiment_id=row.id, + name=environment.name, + source_mode=environment.source_mode, + source_metadata_json=environment.source_metadata, + metadata_json=environment.metadata, + ) ) + self._session.flush() + + return ExperimentRef( + experiment_id=row.id, + name=row.name, + environment_ids=self.environment_ids_for_experiment(row.id), + created_at=row.created_at, + metadata=row.metadata_json, ) - session.flush() - - return ExperimentRef( - experiment_id=row.id, - name=row.name, - environment_ids=_environment_ids_for_experiment(session, row.id), - created_at=row.created_at, - metadata=row.metadata_json, - ) + + def record_sampler_invocation( + self, + *, + experiment_ref: ExperimentRef, + sampler_name: str, + requested_k: int, + candidate_pool_size: int, + selected_count: int = 0, + sampler_config: dict[str, JsonValue] | None = None, + ) -> ExperimentSamplerInvocationRow: + row = ExperimentSamplerInvocationRow( + experiment_id=experiment_ref.experiment_id, + sampler_name=sampler_name, + requested_k=requested_k, + candidate_pool_size=candidate_pool_size, + selected_count=selected_count, + sampler_config_json=dict(sampler_config or {}), + ) + self._session.add(row) + self._session.flush() + return row + + def pending_unselected_pool_entries( + self, + experiment_id: UUID, + ) -> list[ExperimentSamplePoolEntryRow]: + rows = self._session.exec( + select(ExperimentSamplePoolEntryRow) + .where(ExperimentSamplePoolEntryRow.experiment_id == experiment_id) + .where(ExperimentSamplePoolEntryRow.selected.is_(False)) + .where(ExperimentSamplePoolEntryRow.discarded.is_(False)) + .order_by(ExperimentSamplePoolEntryRow.created_at, ExperimentSamplePoolEntryRow.id) + ).all() + return list(rows) + + def experiment_environment_row( + self, + *, + experiment_id: UUID, + environment_name: str, + ) -> ExperimentEnvironmentRow: + return self._session.exec( + select(ExperimentEnvironmentRow) + .where(ExperimentEnvironmentRow.experiment_id == experiment_id) + .where(ExperimentEnvironmentRow.name == environment_name) + ).one() + + def known_sample_keys_by_environment(self, experiment_id: UUID) -> dict[str, set[str]]: + rows = self._session.exec( + select(ExperimentSamplePoolEntryRow, ExperimentEnvironmentRow.name) + .join( + ExperimentEnvironmentRow, + ExperimentSamplePoolEntryRow.environment_id == ExperimentEnvironmentRow.id, + ) + .where(ExperimentSamplePoolEntryRow.experiment_id == experiment_id) + ).all() + known: dict[str, set[str]] = {} + for entry, environment_name in rows: + known.setdefault(environment_name, set()).add(entry.sample_key) + return known + + def record_candidate( + self, + *, + handle: ExperimentRef, + environment_id: UUID, + sample: Sample, + ) -> ExperimentSamplePoolEntryRow: + row = ExperimentSamplePoolEntryRow( + experiment_id=handle.experiment_id, + environment_id=environment_id, + sample_key=sample.sample_key, + sample_json=sample.model_dump(mode="json"), + ) + self._session.add(row) + self._session.flush() + return row + + def mark_pool_entries_selected( + self, + entries: list[ExperimentSamplePoolEntryRow], + *, + sampler_invocation_id: UUID, + ) -> None: + selected_at = utcnow() + for entry in entries: + entry.selected = True + entry.selected_at = selected_at + entry.sampler_invocation_id = sampler_invocation_id + self._session.add(entry) + self._session.flush() + + def environment_ids_for_experiment(self, experiment_id: UUID) -> dict[str, UUID]: + rows = self._session.exec( + select(ExperimentEnvironmentRow).where( + ExperimentEnvironmentRow.experiment_id == experiment_id + ) + ).all() + return {row.name: row.id for row in rows} + + +def persist_experiment(*, session: Session, experiment: Experiment) -> ExperimentRef: + return ExperimentRepository(session).persist_experiment(experiment) def record_sampler_invocation( @@ -55,23 +168,11 @@ def record_sampler_invocation( selected_count: int = 0, sampler_config: dict[str, JsonValue] | None = None, ) -> ExperimentSamplerInvocationRow: - row = ExperimentSamplerInvocationRow( - experiment_id=experiment_ref.experiment_id, + return ExperimentRepository(session).record_sampler_invocation( + experiment_ref=experiment_ref, sampler_name=sampler_name, requested_k=requested_k, candidate_pool_size=candidate_pool_size, selected_count=selected_count, - sampler_config_json=dict(sampler_config or {}), + sampler_config=sampler_config, ) - session.add(row) - session.flush() - return row - - -def _environment_ids_for_experiment(session: Session, experiment_id: UUID) -> dict[str, UUID]: - rows = session.exec( - select(ExperimentEnvironmentRow).where( - ExperimentEnvironmentRow.experiment_id == experiment_id - ) - ).all() - return {row.name: row.id for row in rows} diff --git a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py index 45513041c..6ffb311f7 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -1,5 +1,6 @@ from collections.abc import Iterator from importlib import import_module +from pathlib import Path from typing import Literal from uuid import uuid4 @@ -203,6 +204,15 @@ def test_candidate_pool_stops_after_duplicate_pull_budget( assert duplicate_env.pull_count == 4 +def test_candidate_pool_keeps_sql_access_in_repository() -> None: + source = Path( + "ergon_core/ergon_core/core/application/experiments/candidate_pool.py" + ).read_text() + + assert "session.exec" not in source + assert "select(" not in source + + @pytest.mark.asyncio async def test_candidate_pool_entry_rehydrates_object_bound_sample_without_runtime_task_id( session: Session, From 53663fb7f5c68b6d963b48b0e080c17ccab74c0b Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 00:07:59 +0100 Subject: [PATCH 4/7] Fix experiment repository architecture --- .../application/experiments/candidate_pool.py | 2 +- .../{repositories.py => repository.py} | 27 ++++++++++--------- .../test_experiment_persistence_roundtrip.py | 2 +- .../test_definition_identity_naming.py | 10 +++---- .../experiments/test_candidate_pool.py | 2 +- .../test_experiment_persistence.py | 2 +- 6 files changed, 22 insertions(+), 23 deletions(-) rename ergon_core/ergon_core/core/application/experiments/{repositories.py => repository.py} (91%) diff --git a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py index 52b6f7f9c..64d034163 100644 --- a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -8,7 +8,7 @@ from ergon_core.api.experiment.experiment import Experiment, ExperimentRef from ergon_core.api.experiment.sample import Sample from ergon_core.api.benchmark import Task -from ergon_core.core.application.experiments.repositories import ( +from ergon_core.core.application.experiments.repository import ( ExperimentRepository, ) from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow diff --git a/ergon_core/ergon_core/core/application/experiments/repositories.py b/ergon_core/ergon_core/core/application/experiments/repository.py similarity index 91% rename from ergon_core/ergon_core/core/application/experiments/repositories.py rename to ergon_core/ergon_core/core/application/experiments/repository.py index 95a81387e..14c38d016 100644 --- a/ergon_core/ergon_core/core/application/experiments/repositories.py +++ b/ergon_core/ergon_core/core/application/experiments/repository.py @@ -45,10 +45,18 @@ def persist_experiment(self, experiment: Experiment) -> ExperimentRef: ) self._session.flush() + environment_ids = { + environment.name: environment.id + for environment in self._session.exec( + select(ExperimentEnvironmentRow).where( + ExperimentEnvironmentRow.experiment_id == row.id + ) + ).all() + } return ExperimentRef( experiment_id=row.id, name=row.name, - environment_ids=self.environment_ids_for_experiment(row.id), + environment_ids=environment_ids, created_at=row.created_at, metadata=row.metadata_json, ) @@ -145,14 +153,6 @@ def mark_pool_entries_selected( self._session.add(entry) self._session.flush() - def environment_ids_for_experiment(self, experiment_id: UUID) -> dict[str, UUID]: - rows = self._session.exec( - select(ExperimentEnvironmentRow).where( - ExperimentEnvironmentRow.experiment_id == experiment_id - ) - ).all() - return {row.name: row.id for row in rows} - def persist_experiment(*, session: Session, experiment: Experiment) -> ExperimentRef: return ExperimentRepository(session).persist_experiment(experiment) @@ -168,11 +168,14 @@ def record_sampler_invocation( selected_count: int = 0, sampler_config: dict[str, JsonValue] | None = None, ) -> ExperimentSamplerInvocationRow: - return ExperimentRepository(session).record_sampler_invocation( - experiment_ref=experiment_ref, + row = ExperimentSamplerInvocationRow( + experiment_id=experiment_ref.experiment_id, sampler_name=sampler_name, requested_k=requested_k, candidate_pool_size=candidate_pool_size, selected_count=selected_count, - sampler_config=sampler_config, + sampler_config_json=dict(sampler_config or {}), ) + session.add(row) + session.flush() + return row diff --git a/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py index c9af3b838..6f956da56 100644 --- a/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py +++ b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py @@ -12,7 +12,7 @@ from sqlalchemy import inspect from ergon_core.api import Environment, Experiment, Sample -from ergon_core.core.application.experiments.repositories import persist_experiment +from ergon_core.core.application.experiments.repository import persist_experiment from ergon_core.core.persistence.experiments.models import ExperimentEnvironmentRow from ergon_core.test_support.task_factory import task_with_id diff --git a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py index 0adf7030e..a034ad370 100644 --- a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py +++ b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py @@ -37,14 +37,10 @@ re.compile(r"experiment_id: UUID"), re.compile(r"experiment_id == experiment_id"), ), - ROOT - / "ergon_core" - / "ergon_core" - / "core" - / "application" - / "experiments" - / "repositories.py": ( + ROOT / "ergon_core" / "ergon_core" / "core" / "application" / "experiments" / "repository.py": ( re.compile(r"experiment_id=row\.id"), + re.compile(r"experiment_id == row\.id"), + re.compile(r"experiment_id=handle\.id"), re.compile(r"experiment_ref\.experiment_id"), re.compile(r"experiment_id: UUID"), re.compile(r"experiment_id == experiment_id"), diff --git a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py index 6ffb311f7..0525549d5 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -12,7 +12,7 @@ SampleCandidatePool, sample_from_pool_entry, ) -from ergon_core.core.application.experiments.repositories import ( +from ergon_core.core.application.experiments.repository import ( persist_experiment, record_sampler_invocation, ) diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py index bde9e4243..6619c2612 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py @@ -7,7 +7,7 @@ from sqlmodel import Session, SQLModel, create_engine, select from ergon_core.api import Environment, Experiment, Sample -from ergon_core.core.application.experiments.repositories import ( +from ergon_core.core.application.experiments.repository import ( persist_experiment, record_sampler_invocation, ) From 3d9f98e5fbc326a13bc41f76ca49907bb0b14c6a Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 13:15:02 +0100 Subject: [PATCH 5/7] Bind experiment persistence facade to core rows --- .../application/experiments/persistence.py | 20 ++++++++++ .../application/experiments/repository.py | 14 ++++--- .../core/application/experiments/service.py | 20 ---------- .../core/persistence/experiments/models.py | 2 + .../00000001_add_experiment_persistence.py | 11 ++++++ .../experiments/test_candidate_pool.py | 7 ++-- .../test_experiment_persistence.py | 38 ++++++++++++++++++- 7 files changed, 81 insertions(+), 31 deletions(-) create mode 100644 ergon_core/ergon_core/core/application/experiments/persistence.py diff --git a/ergon_core/ergon_core/core/application/experiments/persistence.py b/ergon_core/ergon_core/core/application/experiments/persistence.py new file mode 100644 index 000000000..0f0aff231 --- /dev/null +++ b/ergon_core/ergon_core/core/application/experiments/persistence.py @@ -0,0 +1,20 @@ +"""Concrete core adapter for the public experiment persistence facade.""" + +from __future__ import annotations + +from sqlmodel import Session + +from ergon_core.api.experiment.experiment import Experiment, ExperimentRef +from ergon_core.core.application.experiments.repository import ( + persist_experiment as persist_row_graph, +) + + +class CoreExperimentPersistencePort: + """Application-backed implementation of ``api.experiment.persist_experiment``.""" + + def __init__(self, session: Session) -> None: + self._session = session + + async def persist_experiment(self, experiment: Experiment) -> ExperimentRef: + return persist_row_graph(session=self._session, experiment=experiment) diff --git a/ergon_core/ergon_core/core/application/experiments/repository.py b/ergon_core/ergon_core/core/application/experiments/repository.py index 14c38d016..671fdc1ab 100644 --- a/ergon_core/ergon_core/core/application/experiments/repository.py +++ b/ergon_core/ergon_core/core/application/experiments/repository.py @@ -69,6 +69,7 @@ def record_sampler_invocation( requested_k: int, candidate_pool_size: int, selected_count: int = 0, + policy_version: int | None = None, sampler_config: dict[str, JsonValue] | None = None, ) -> ExperimentSamplerInvocationRow: row = ExperimentSamplerInvocationRow( @@ -77,6 +78,7 @@ def record_sampler_invocation( requested_k=requested_k, candidate_pool_size=candidate_pool_size, selected_count=selected_count, + policy_version=policy_version, sampler_config_json=dict(sampler_config or {}), ) self._session.add(row) @@ -133,6 +135,7 @@ def record_candidate( experiment_id=handle.experiment_id, environment_id=environment_id, sample_key=sample.sample_key, + sample_ref_json=sample.sample_ref, sample_json=sample.model_dump(mode="json"), ) self._session.add(row) @@ -166,16 +169,15 @@ def record_sampler_invocation( requested_k: int, candidate_pool_size: int, selected_count: int = 0, + policy_version: int | None = None, sampler_config: dict[str, JsonValue] | None = None, ) -> ExperimentSamplerInvocationRow: - row = ExperimentSamplerInvocationRow( - experiment_id=experiment_ref.experiment_id, + return ExperimentRepository(session).record_sampler_invocation( + experiment_ref=experiment_ref, sampler_name=sampler_name, requested_k=requested_k, candidate_pool_size=candidate_pool_size, selected_count=selected_count, - sampler_config_json=dict(sampler_config or {}), + policy_version=policy_version, + sampler_config=sampler_config, ) - session.add(row) - session.flush() - return row diff --git a/ergon_core/ergon_core/core/application/experiments/service.py b/ergon_core/ergon_core/core/application/experiments/service.py index 4317b7d1a..e43173baa 100644 --- a/ergon_core/ergon_core/core/application/experiments/service.py +++ b/ergon_core/ergon_core/core/application/experiments/service.py @@ -16,7 +16,6 @@ from ergon_core.api.benchmark import Benchmark from ergon_core.api.experiment.experiment import ( Experiment, - ExperimentRef, ExperimentSubmitResult, ) from ergon_core.api.experiment.sampling import Sampler @@ -38,25 +37,6 @@ async def submit( policy_version: int | None, ) -> "ExperimentSubmitResult": ... - -@runtime_checkable -class PersistExperimentPort(Protocol): - # TODO(PR04): replace this protocol with the concrete core persistence - # service once experiment/environment/candidate-pool rows exist. - async def persist_experiment(self, experiment: "Experiment") -> "ExperimentRef": ... - - -async def persist_experiment( - experiment: "Experiment", - *, - service: PersistExperimentPort, -) -> "ExperimentRef": - # TODO(PR04): move callers to the concrete core persistence entry point - # after experiment rows and environment rows are introduced. - experiment.validate_authoring() - return await service.persist_experiment(experiment) - - def persist_benchmark(benchmark: "Benchmark") -> DefinitionHandle: """Persist a configured object-bound Benchmark as an experiment definition.""" diff --git a/ergon_core/ergon_core/core/persistence/experiments/models.py b/ergon_core/ergon_core/core/persistence/experiments/models.py index 44afe7c74..41d94243e 100644 --- a/ergon_core/ergon_core/core/persistence/experiments/models.py +++ b/ergon_core/ergon_core/core/persistence/experiments/models.py @@ -46,6 +46,7 @@ class ExperimentSamplerInvocationRow(SQLModel, table=True): requested_k: int candidate_pool_size: int selected_count: int = 0 + policy_version: int | None = Field(default=None, index=True) sampler_config_json: dict[str, JsonValue] = Field( default_factory=dict, sa_column=Column(JSON), @@ -66,6 +67,7 @@ class ExperimentSamplePoolEntryRow(SQLModel, table=True): index=True, ) sample_key: str = Field(index=True) + sample_ref_json: dict[str, JsonValue] = Field(default_factory=dict, sa_column=Column(JSON)) sample_json: dict[str, JsonValue] = Field(default_factory=dict, sa_column=Column(JSON)) selected: bool = Field( default=False, diff --git a/ergon_core/migrations/versions/00000001_add_experiment_persistence.py b/ergon_core/migrations/versions/00000001_add_experiment_persistence.py index ee720b5e7..408b4f1af 100644 --- a/ergon_core/migrations/versions/00000001_add_experiment_persistence.py +++ b/ergon_core/migrations/versions/00000001_add_experiment_persistence.py @@ -66,6 +66,7 @@ def upgrade() -> None: sa.Column("requested_k", sa.Integer(), nullable=False), sa.Column("candidate_pool_size", sa.Integer(), nullable=False), sa.Column("selected_count", sa.Integer(), nullable=False), + sa.Column("policy_version", sa.Integer(), nullable=True), sa.Column("sampler_config_json", sa.JSON(), nullable=True), sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), sa.ForeignKeyConstraint(["experiment_id"], ["experiments.id"]), @@ -76,6 +77,11 @@ def upgrade() -> None: "experiment_sampler_invocations", ["experiment_id"], ) + op.create_index( + op.f("ix_experiment_sampler_invocations_policy_version"), + "experiment_sampler_invocations", + ["policy_version"], + ) op.create_index( op.f("ix_experiment_sampler_invocations_sampler_name"), "experiment_sampler_invocations", @@ -89,6 +95,7 @@ def upgrade() -> None: sa.Column("environment_id", sa.Uuid(), nullable=False), sa.Column("sampler_invocation_id", sa.Uuid(), nullable=True), sa.Column("sample_key", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("sample_ref_json", sa.JSON(), nullable=True), sa.Column("sample_json", sa.JSON(), nullable=True), sa.Column("selected", sa.Boolean(), nullable=False), sa.Column("discarded", sa.Boolean(), nullable=False), @@ -170,6 +177,10 @@ def downgrade() -> None: op.f("ix_experiment_sampler_invocations_experiment_id"), table_name="experiment_sampler_invocations", ) + op.drop_index( + op.f("ix_experiment_sampler_invocations_policy_version"), + table_name="experiment_sampler_invocations", + ) op.drop_table("experiment_sampler_invocations") op.drop_index( op.f("ix_experiment_environments_source_mode"), diff --git a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py index 0525549d5..61aa5b984 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -136,6 +136,8 @@ def test_candidate_pool_retains_unselected_streamed_candidates( assert len(rows) == 8 assert sum(row.selected for row in rows) == 3 assert sum(row.discarded for row in rows) == 0 + assert rows[0].sample_ref_json == {"key": "0"} + assert rows[0].sample_json["sample_ref"] == rows[0].sample_ref_json def test_candidate_pool_reuses_unselected_entries_before_advancing_stream( @@ -159,9 +161,7 @@ def test_candidate_pool_reuses_unselected_entries_before_advancing_stream( candidate_pool_size=8, ) - assert [entry.sample_key for entry in second[:6]] == [ - entry.sample_key for entry in first[2:] - ] + assert [entry.sample_key for entry in second[:6]] == [entry.sample_key for entry in first[2:]] assert counted_streaming_experiment.environments[0].pull_count == 10 @@ -187,6 +187,7 @@ def test_candidate_pool_round_robins_new_entries_across_environments( "swe-validation", ] assert entries[0].sample_json["sample_key"] == "a" + assert entries[0].sample_ref_json == {"key": "a"} assert entries[0].sample_json["tasks"][0]["task_slug"] == "solve-mini-validation-a" diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py index 6619c2612..9ed4913ac 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py @@ -6,12 +6,22 @@ import pytest from sqlmodel import Session, SQLModel, create_engine, select -from ergon_core.api import Environment, Experiment, Sample +from ergon_core.api import ( + Environment, + Experiment, + Sample, + persist_experiment as persist_public_experiment, +) +from ergon_core.core.application.experiments.persistence import CoreExperimentPersistencePort from ergon_core.core.application.experiments.repository import ( persist_experiment, record_sampler_invocation, ) -from ergon_core.core.persistence.experiments.models import ExperimentEnvironmentRow +from ergon_core.core.persistence.experiments.models import ( + ExperimentEnvironmentRow, + ExperimentRow, + ExperimentSamplePoolEntryRow, +) from ergon_core.core.persistence.samples.models import SampleStatusEventRow, SampleTaskEventRow from ergon_core.core.persistence.telemetry.models import SampleRecord from ergon_core.test_support.task_factory import task_with_id @@ -112,9 +122,33 @@ def test_record_sampler_invocation_writes_no_runtime_state( requested_k=4, candidate_pool_size=16, selected_count=0, + policy_version=3, sampler_config={"seed": 1}, ) assert invocation.experiment_id == handle.experiment_id + assert invocation.policy_version == 3 assert session.exec(select(SampleRecord)).all() == [] assert session.exec(select(SampleTaskEventRow)).all() == [] + + +@pytest.mark.asyncio +async def test_public_persistence_facade_wires_to_core_port( + session: Session, + two_env_experiment: Experiment, +) -> None: + ref = await persist_public_experiment( + two_env_experiment, + service=CoreExperimentPersistencePort(session), + ) + + assert ref.experiment_id + assert session.get(ExperimentRow, ref.experiment_id) is not None + + +def test_experiment_persistence_uses_sample_ref_json_without_source_alias() -> None: + tables = SQLModel.metadata.tables + + assert "sample_ref_json" in ExperimentSamplePoolEntryRow.model_fields + assert "source_sample_ref_json" not in ExperimentSamplePoolEntryRow.model_fields + assert all("source_sample_ref_json" not in table.columns for table in tables.values()) From 6a4bcfb327cc328e674877f866e007bdbfbccf80 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 14:57:39 +0100 Subject: [PATCH 6/7] Format experiment persistence files --- ergon_core/ergon_core/api/experiment/sample.py | 2 ++ .../core/application/experiments/candidate_pool.py | 8 ++++++-- .../ergon_core/core/application/experiments/service.py | 1 + 3 files changed, 9 insertions(+), 2 deletions(-) diff --git a/ergon_core/ergon_core/api/experiment/sample.py b/ergon_core/ergon_core/api/experiment/sample.py index d2d836e58..f93aef7bc 100644 --- a/ergon_core/ergon_core/api/experiment/sample.py +++ b/ergon_core/ergon_core/api/experiment/sample.py @@ -7,6 +7,8 @@ from pydantic import BaseModel, ConfigDict, Field, JsonValue, model_validator from ergon_core.api.benchmark import Task + + class Sample(BaseModel): """A selected runnable unit containing concrete public ``Task`` objects.""" diff --git a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py index 64d034163..a4e7e953c 100644 --- a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -105,7 +105,9 @@ async def sample_from_pool_entry(entry: ExperimentSamplePoolEntryRow) -> Sample: sample_key=str(payload["sample_key"]), environment_name=str(payload["environment_name"]), tasks=tasks, - sample_ref=payload.get("sample_ref") if isinstance(payload.get("sample_ref"), dict) else None, + sample_ref=payload.get("sample_ref") + if isinstance(payload.get("sample_ref"), dict) + else None, source_metadata=( payload.get("source_metadata") if isinstance(payload.get("source_metadata"), dict) @@ -117,7 +119,9 @@ async def sample_from_pool_entry(entry: ExperimentSamplePoolEntryRow) -> Sample: async def _task_from_candidate_snapshot(task_json: object) -> Task: if not isinstance(task_json, dict): - raise ValueError(f"Candidate task snapshot must be an object, got {type(task_json).__name__}") + raise ValueError( + f"Candidate task snapshot must be an object, got {type(task_json).__name__}" + ) task = await Task.from_definition(task_json, task_id=uuid4()) # Candidate-pool rows are not runtime materializations. We use the existing # `_type` dispatch path to rebuild object-bound config, then clear the diff --git a/ergon_core/ergon_core/core/application/experiments/service.py b/ergon_core/ergon_core/core/application/experiments/service.py index e43173baa..2697d0504 100644 --- a/ergon_core/ergon_core/core/application/experiments/service.py +++ b/ergon_core/ergon_core/core/application/experiments/service.py @@ -37,6 +37,7 @@ async def submit( policy_version: int | None, ) -> "ExperimentSubmitResult": ... + def persist_benchmark(benchmark: "Benchmark") -> DefinitionHandle: """Persist a configured object-bound Benchmark as an experiment definition.""" From 98fe37b26fbb972115263160ddbbf4446badc223 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 15:39:05 +0100 Subject: [PATCH 7/7] Resume environment cursors for candidate refill --- .../application/experiments/candidate_pool.py | 2 +- .../experiments/test_candidate_pool.py | 33 +++++++++++++++++++ 2 files changed, 34 insertions(+), 1 deletion(-) diff --git a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py index a4e7e953c..aea1b9285 100644 --- a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -45,7 +45,7 @@ def fill( } known_keys = self._repository.known_sample_keys_by_environment(handle.experiment_id) iterators = { - environment.name: iter(environment.iter_samples()) + environment.name: environment.iter_candidate_samples() for environment in experiment.environments } active_names = [environment.name for environment in experiment.environments] diff --git a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py index 61aa5b984..2b1f21ca9 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -165,6 +165,39 @@ def test_candidate_pool_reuses_unselected_entries_before_advancing_stream( assert counted_streaming_experiment.environments[0].pull_count == 10 +def test_candidate_pool_resumes_stateless_stream_cursor_after_retained_entries( + session: Session, + streamed_experiment: Experiment, +) -> None: + handle = persist_experiment(session=session, experiment=streamed_experiment) + pool = SampleCandidatePool(session) + + first = pool.fill( + experiment=streamed_experiment, + handle=handle, + candidate_pool_size=8, + ) + pool.mark_selected(first[:2], sampler_invocation_id=uuid4()) + session.commit() + + second = pool.fill( + experiment=streamed_experiment, + handle=handle, + candidate_pool_size=8, + ) + + assert [entry.sample_key for entry in second] == [ + "2", + "3", + "4", + "5", + "6", + "7", + "8", + "9", + ] + + def test_candidate_pool_round_robins_new_entries_across_environments( session: Session, ) -> None: