Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
130 changes: 130 additions & 0 deletions ergon_core/ergon_core/core/application/experiments/candidate_pool.py
Original file line number Diff line number Diff line change
@@ -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
20 changes: 20 additions & 0 deletions ergon_core/ergon_core/core/application/experiments/persistence.py
Original file line number Diff line number Diff line change
@@ -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)
183 changes: 183 additions & 0 deletions ergon_core/ergon_core/core/application/experiments/repository.py
Original file line number Diff line number Diff line change
@@ -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,
)
19 changes: 0 additions & 19 deletions ergon_core/ergon_core/core/application/experiments/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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."""

Expand Down
15 changes: 15 additions & 0 deletions ergon_core/ergon_core/core/persistence/experiments/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
Loading
Loading