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..aea1b9285 --- /dev/null +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -0,0 +1,130 @@ +"""Candidate-pool persistence for public experiment samples.""" + +from collections.abc import Sequence +from uuid import UUID, uuid4 + +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.application.experiments.repository import ( + ExperimentRepository, +) +from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow + + +class SampleCandidatePool: + def __init__( + self, + session: Session, + *, + max_duplicate_pulls_per_environment: int = 1_000, + repository: ExperimentRepository | None = None, + ) -> None: + self._repository = repository or ExperimentRepository(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._repository.pending_unselected_pool_entries(handle.experiment_id) + if len(entries) >= candidate_pool_size: + return entries[:candidate_pool_size] + + env_rows = { + environment.name: self._repository.experiment_environment_row( + experiment_id=handle.experiment_id, + environment_name=environment.name, + ) + for environment in experiment.environments + } + known_keys = self._repository.known_sample_keys_by_environment(handle.experiment_id) + iterators = { + environment.name: environment.iter_candidate_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._repository.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: + self._repository.mark_pool_entries_selected( + list(entries), + sampler_invocation_id=sampler_invocation_id, + ) + + +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/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 new file mode 100644 index 000000000..671fdc1ab --- /dev/null +++ b/ergon_core/ergon_core/core/application/experiments/repository.py @@ -0,0 +1,183 @@ +"""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.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 + + +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() + + 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=environment_ids, + 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, + policy_version: int | None = None, + 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, + policy_version=policy_version, + 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_ref_json=sample.sample_ref, + 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 persist_experiment(*, session: Session, experiment: Experiment) -> ExperimentRef: + return ExperimentRepository(session).persist_experiment(experiment) + + +def record_sampler_invocation( + *, + session: Session, + experiment_ref: ExperimentRef, + sampler_name: str, + requested_k: int, + candidate_pool_size: int, + selected_count: int = 0, + policy_version: int | None = None, + sampler_config: dict[str, JsonValue] | None = None, +) -> ExperimentSamplerInvocationRow: + 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, + policy_version=policy_version, + sampler_config=sampler_config, + ) diff --git a/ergon_core/ergon_core/core/application/experiments/service.py b/ergon_core/ergon_core/core/application/experiments/service.py index 4317b7d1a..2697d0504 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 @@ -39,24 +38,6 @@ async def submit( ) -> "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/__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..41d94243e --- /dev/null +++ b/ergon_core/ergon_core/core/persistence/experiments/models.py @@ -0,0 +1,82 @@ +"""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 + policy_version: int | None = Field(default=None, index=True) + 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_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, + 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..408b4f1af --- /dev/null +++ b/ergon_core/migrations/versions/00000001_add_experiment_persistence.py @@ -0,0 +1,197 @@ +"""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("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"]), + 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_policy_version"), + "experiment_sampler_invocations", + ["policy_version"], + ) + 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_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), + 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_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"), + 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..6f956da56 --- /dev/null +++ b/ergon_core/tests/integration/experiments/test_experiment_persistence_roundtrip.py @@ -0,0 +1,137 @@ +from collections.abc import Iterator +from importlib import import_module +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 + +from ergon_core.api import Environment, Experiment, Sample +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 + +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( + 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..a034ad370 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,59 @@ ) 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" / "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"), + ), + 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 +93,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..2b1f21ca9 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_candidate_pool.py @@ -0,0 +1,271 @@ +from collections.abc import Iterator +from importlib import import_module +from pathlib import Path +from typing import Literal +from uuid import uuid4 + +import pytest +from sqlmodel import Session, SQLModel, create_engine, select + +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.repository 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 + +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( + 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 + 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( + 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_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: + 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_ref_json == {"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 + + +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, +) -> 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..9ed4913ac --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_persistence.py @@ -0,0 +1,154 @@ +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 + +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, + 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 + +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( + 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, + 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())