From d9553d6a350f2c17b8e96af0b6c3f7bf735857c7 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 16:10:07 +0100 Subject: [PATCH 1/9] Materialize experiment samples on submit --- docs/architecture/04_persistence.md | 12 + .../core/application/events/runtime.py | 22 +- .../application/experiments/candidate_pool.py | 7 +- .../application/experiments/submission.py | 185 +++++++++++ .../core/application/runtime/events.py | 4 +- .../core/application/runtime/lifecycle.py | 44 ++- .../core/application/runtime/orchestration.py | 14 +- .../application/runtime/sample_identity.py | 2 +- .../application/runtime/sample_lifecycle.py | 204 +++++++++--- .../application/runtime/task_execution.py | 76 ++++- .../application/runtime/task_management.py | 6 +- .../application/samples/materialization.py | 189 +++++++++++ .../resources/persist_outputs/contract.py | 2 +- .../core/jobs/sandbox/setup/contract.py | 2 +- .../core/jobs/task/cancel_orphans/job.py | 2 +- .../core/jobs/task/worker_execute/contract.py | 2 +- .../core/jobs/task/worker_execute/job.py | 4 +- .../core/jobs/workflow/start/job.py | 48 ++- .../core/persistence/telemetry/models.py | 9 +- ...000002_add_sample_experiment_provenance.py | 96 ++++++ .../test_public_experiment_submit_smoke.py | 20 +- .../test_submit_starts_materialized_sample.py | 250 ++++++++++++++ .../experiments/test_experiment_submit.py | 310 ++++++++++++++++++ .../test_sample_materialization.py | 140 ++++++++ 24 files changed, 1528 insertions(+), 122 deletions(-) create mode 100644 ergon_core/ergon_core/core/application/experiments/submission.py create mode 100644 ergon_core/ergon_core/core/application/samples/materialization.py create mode 100644 ergon_core/migrations/versions/00000002_add_sample_experiment_provenance.py create mode 100644 ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py create mode 100644 ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py create mode 100644 ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py diff --git a/docs/architecture/04_persistence.md b/docs/architecture/04_persistence.md index d9e9bca7c..2fb058635 100644 --- a/docs/architecture/04_persistence.md +++ b/docs/architecture/04_persistence.md @@ -70,6 +70,18 @@ rows and eagerly creates every statically-declared node and edge. Task payloads are attached as annotations in the core-reserved `"payload"` namespace. +PR05 adds the public experiment submit path alongside that bridge: +`Experiment.submit(...)` persists an experiment and candidate pool, records the +sampler invocation, marks selected candidate rows, and creates one +`SampleRecord` per selected sample. The materialization boundary is the authored +`Sample`'s concrete `Task` objects: each task is persisted as +`task.model_dump(mode="json")` into typed sample WAL rows and the +`sample_graph_nodes.task_json` projection. Environment and experiment objects +remain provenance rows; they are not serialized into the runtime replay +contract. `workflow/started` may now carry only `sample_id` when the sample +graph already exists, while the definition-backed initialization path remains as +a temporary bridge until the runtime consolidation PR removes it. + ### Manager-spawned subtasks (dynamic graph growth) The graph is append-only at the row level: new nodes and edges can enter diff --git a/ergon_core/ergon_core/core/application/events/runtime.py b/ergon_core/ergon_core/core/application/events/runtime.py index f39d4185b..beba967f7 100644 --- a/ergon_core/ergon_core/core/application/events/runtime.py +++ b/ergon_core/ergon_core/core/application/events/runtime.py @@ -7,7 +7,7 @@ from __future__ import annotations -from typing import ClassVar, Literal +from typing import Any, ClassVar, Literal from uuid import UUID from ergon_core.core.application.events.base import InngestEventContract @@ -31,7 +31,7 @@ class TaskReadyEvent(InngestEventContract): name: ClassVar[str] = "task/ready" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID @@ -39,7 +39,7 @@ class TaskStartedEvent(InngestEventContract): name: ClassVar[str] = "task/started" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID @@ -58,7 +58,7 @@ class TaskCancelledEvent(InngestEventContract): name: ClassVar[str] = "task/cancelled" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID | None cause: CancelCause @@ -70,7 +70,7 @@ class TaskCompletedEvent(InngestEventContract): name: ClassVar[str] = "task/completed" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID sandbox_id: str @@ -80,7 +80,7 @@ class TaskFailedEvent(InngestEventContract): name: ClassVar[str] = "task/failed" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID error: str @@ -91,19 +91,23 @@ class WorkflowStartedEvent(InngestEventContract): name: ClassVar[str] = "workflow/started" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None + + def model_dump(self, *args: Any, **kwargs: Any) -> dict[str, Any]: + kwargs.setdefault("exclude_none", True) + return super().model_dump(*args, **kwargs) class WorkflowCompletedEvent(InngestEventContract): name: ClassVar[str] = "workflow/completed" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None class WorkflowFailedEvent(InngestEventContract): name: ClassVar[str] = "workflow/failed" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None error: str 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 aea1b9285..a5d236563 100644 --- a/ergon_core/ergon_core/core/application/experiments/candidate_pool.py +++ b/ergon_core/ergon_core/core/application/experiments/candidate_pool.py @@ -5,9 +5,10 @@ from sqlmodel import Session +from ergon_core.api.benchmark import Task +from ergon_core.api.benchmark.task import EmptyTaskPayload 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, ) @@ -126,5 +127,9 @@ async def _task_from_candidate_snapshot(task_json: object) -> Task: # 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. + if isinstance(task.dependency_task_slugs, list): + task.dependency_task_slugs = tuple(task.dependency_task_slugs) + if isinstance(task.task_payload, dict) and not task.task_payload: + task.task_payload = EmptyTaskPayload() task._task_id = None return task diff --git a/ergon_core/ergon_core/core/application/experiments/submission.py b/ergon_core/ergon_core/core/application/experiments/submission.py new file mode 100644 index 000000000..cb0e8c4fe --- /dev/null +++ b/ergon_core/ergon_core/core/application/experiments/submission.py @@ -0,0 +1,185 @@ +"""Experiment submit orchestration for selected sample materialization.""" + +from collections.abc import Sequence +from typing import Protocol +from uuid import UUID + +from pydantic import JsonValue +from sqlmodel import Session + +from ergon_core.api.experiment.experiment import Experiment, ExperimentSubmitResult +from ergon_core.api.experiment.sample import Sample +from ergon_core.api.experiment.sampling import Sampler, SamplingContext +from ergon_core.core.application.events.runtime import WorkflowStartedEvent +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.application.samples.materialization import materialize_sample +from ergon_core.core.infrastructure.inngest.client import InngestEvent, inngest_client +from ergon_core.core.persistence.experiments.models import ( + ExperimentSamplePoolEntryRow, + ExperimentSamplerInvocationRow, +) +from ergon_core.core.persistence.shared.enums import SampleStatus +from ergon_core.core.persistence.telemetry.models import SampleRecord + + +class EventBus(Protocol): + async def publish(self, event: WorkflowStartedEvent) -> None: ... + + +class InngestWorkflowEventBus: + async def publish(self, event: WorkflowStartedEvent) -> None: + await inngest_client.send( + InngestEvent( + name=WorkflowStartedEvent.name, + data=event.model_dump(mode="json"), + ) + ) + + +class ExperimentSubmissionService: + def __init__(self, *, session: Session, event_bus: EventBus | None = None) -> None: + self._session = session + self._event_bus = event_bus or InngestWorkflowEventBus() + + @classmethod + def for_session( + cls, + session: Session, + *, + event_bus: EventBus | None = None, + ) -> "ExperimentSubmissionService": + return cls(session=session, event_bus=event_bus) + + async def submit( + self, + *, + experiment: Experiment, + k: int, + sampler: Sampler, + candidate_pool_size: int | None, + policy_version: int | None = None, + ) -> ExperimentSubmitResult: + del policy_version + handle = persist_experiment(session=self._session, experiment=experiment) + pool_size = candidate_pool_size or k + pool = SampleCandidatePool(self._session) + entries = pool.fill( + experiment=experiment, + handle=handle, + candidate_pool_size=pool_size, + ) + candidates = [await sample_from_pool_entry(entry) for entry in entries] + sampler_selected = list( + await sampler.select( + samples=candidates, + k=k, + context=SamplingContext( + experiment_id=handle.experiment_id, + candidate_pool_size=pool_size, + ), + ) + ) + selected = sampler_selected[:k] + selected_entries = _entries_for_selected_samples(entries=entries, selected=selected) + invocation = record_sampler_invocation( + session=self._session, + experiment_ref=handle, + sampler_name=sampler.name, + requested_k=k, + candidate_pool_size=pool_size, + selected_count=len(selected_entries), + sampler_config=sampler.config(), + ) + pool.mark_selected(selected_entries, sampler_invocation_id=invocation.id) + sample_ids = self._materialize_selected( + invocation=invocation, + selected=selected, + selected_entries=selected_entries, + ) + self._session.commit() + await self._start_samples(sample_ids) + return ExperimentSubmitResult( + experiment_id=handle.experiment_id, + sampler_invocation_id=invocation.id, + requested_k=k, + candidate_pool_size=pool_size, + selected_count=len(sample_ids), + sample_ids=sample_ids, + ) + + def _materialize_selected( + self, + *, + invocation: ExperimentSamplerInvocationRow, + selected: Sequence[Sample], + selected_entries: Sequence[ExperimentSamplePoolEntryRow], + ) -> list[UUID]: + sample_ids: list[UUID] = [] + for sample, entry in zip(selected, selected_entries, strict=True): + row = SampleRecord( + experiment_id=entry.experiment_id, + environment_id=entry.environment_id, + sampler_invocation_id=invocation.id, + pool_entry_id=entry.id, + sample_key=sample.sample_key, + sample_ref_json=dict(sample.sample_ref), + benchmark_type="experiment", + instance_key=sample.sample_key, + sample_id=sample.sample_key, + worker_team_json={}, + dependency_extras_json={}, + assignment_json=_assignment_json(sample), + experiment=str(entry.experiment_id), + status=SampleStatus.PENDING, + ) + self._session.add(row) + self._session.flush() + materialize_sample(session=self._session, sample=sample, sample_row=row) + sample_ids.append(row.id) + self._session.flush() + return sample_ids + + async def _start_samples(self, sample_ids: Sequence[UUID]) -> None: + for sample_id in sample_ids: + await self._event_bus.publish(WorkflowStartedEvent(sample_id=sample_id)) + + +def _entries_for_selected_samples( + *, + entries: Sequence[ExperimentSamplePoolEntryRow], + selected: Sequence[Sample], +) -> list[ExperimentSamplePoolEntryRow]: + remaining = list(entries) + selected_entries: list[ExperimentSamplePoolEntryRow] = [] + for sample in selected: + for index, entry in enumerate(remaining): + if ( + entry.sample_key == sample.sample_key + and entry.sample_json.get("environment_name") == sample.environment_name + ): + selected_entries.append(remaining.pop(index)) + break + else: + raise ValueError( + "Sampler returned a sample that was not present in the candidate pool: " + f"{sample.environment_name}/{sample.sample_key}" + ) + return selected_entries + + +def _assignment_json(sample: Sample) -> dict[str, JsonValue]: + return { + "sample_key": sample.sample_key, + "sample_name": sample.name, + "environment_name": sample.environment_name, + "sample_ref": dict(sample.sample_ref), + "source_metadata": dict(sample.source_metadata), + "metadata": dict(sample.metadata), + } diff --git a/ergon_core/ergon_core/core/application/runtime/events.py b/ergon_core/ergon_core/core/application/runtime/events.py index 5be7f7850..d14588f0f 100644 --- a/ergon_core/ergon_core/core/application/runtime/events.py +++ b/ergon_core/ergon_core/core/application/runtime/events.py @@ -18,7 +18,7 @@ logger = logging.getLogger(__name__) -TaskReadyDispatcher = Callable[[UUID, UUID, UUID], Awaitable[None]] +TaskReadyDispatcher = Callable[[UUID, UUID | None, UUID], Awaitable[None]] class RuntimeEventDispatcher: @@ -31,7 +31,7 @@ async def dispatch_task_ready( self, *, sample_id: UUID, - definition_id: UUID, + definition_id: UUID | None, task_id: UUID, ) -> None: """Emit the canonical ``task/ready`` event for a committed task state.""" diff --git a/ergon_core/ergon_core/core/application/runtime/lifecycle.py b/ergon_core/ergon_core/core/application/runtime/lifecycle.py index 468a2c5c0..a01acae66 100644 --- a/ergon_core/ergon_core/core/application/runtime/lifecycle.py +++ b/ergon_core/ergon_core/core/application/runtime/lifecycle.py @@ -58,7 +58,7 @@ async def mark_task_ready( session, sample_id, task_id, - graph_status.PENDING, + graph_status.READY, graph_repo=graph_repo, graph_lookup=graph_lookup, ) @@ -107,23 +107,38 @@ async def mark_task_failed( async def get_initial_ready_tasks( session: Session, sample_id: UUID, - definition_id: UUID, + definition_id: UUID | None = None, *, - graph_repo: RuntimeGraphRepository, - graph_lookup: GraphNodeLookup, + graph_repo: RuntimeGraphRepository | None = None, + graph_lookup: GraphNodeLookup | None = None, + commit: bool = True, ) -> list[UUID]: """Return task IDs that have zero dependencies.""" - all_tasks_stmt = select(ExperimentDefinitionTask.id).where( - ExperimentDefinitionTask.experiment_definition_id == definition_id, - ) + graph_repo = graph_repo or RuntimeGraphRepository() + graph_lookup = graph_lookup or GraphNodeLookup(session, sample_id) + if definition_id is None: + all_tasks_stmt = select(SampleGraphNode.task_id).where( + SampleGraphNode.sample_id == sample_id, + ) + tasks_with_deps_stmt = select(SampleGraphEdge.target_task_id).where( + SampleGraphEdge.sample_id == sample_id, + ) + else: + all_tasks_stmt = select(ExperimentDefinitionTask.id).where( + ExperimentDefinitionTask.experiment_definition_id == definition_id, + ) + tasks_with_deps_stmt = select(ExperimentDefinitionTaskDependency.task_id).where( + ExperimentDefinitionTaskDependency.experiment_definition_id == definition_id, + ) all_task_ids = set(session.exec(all_tasks_stmt).all()) - - tasks_with_deps_stmt = select(ExperimentDefinitionTaskDependency.task_id).where( - ExperimentDefinitionTaskDependency.experiment_definition_id == definition_id, - ) tasks_with_deps = set(session.exec(tasks_with_deps_stmt).all()) - ready_ids = list(all_task_ids - tasks_with_deps) + ready_ids = [] + for task_id in sorted(all_task_ids - tasks_with_deps): + node = session.get(SampleGraphNode, (sample_id, task_id)) + if node is None or node.status != graph_status.PENDING: + continue + ready_ids.append(task_id) for task_id in ready_ids: await mark_task_ready( @@ -134,7 +149,10 @@ async def get_initial_ready_tasks( graph_lookup=graph_lookup, ) - session.commit() + if commit: + session.commit() + else: + session.flush() return ready_ids diff --git a/ergon_core/ergon_core/core/application/runtime/orchestration.py b/ergon_core/ergon_core/core/application/runtime/orchestration.py index a4d24d631..47ff91d60 100644 --- a/ergon_core/ergon_core/core/application/runtime/orchestration.py +++ b/ergon_core/ergon_core/core/application/runtime/orchestration.py @@ -36,14 +36,14 @@ class InitializeWorkflowCommand(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None class InitializedWorkflow(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None benchmark_type: str total_tasks: int total_root_tasks: int @@ -55,7 +55,7 @@ class PrepareTaskExecutionCommand(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID @@ -69,7 +69,7 @@ class PreparedTaskExecution(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID task_slug: str task_description: str @@ -110,7 +110,7 @@ class PropagateTaskCompletionCommand(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID @@ -119,7 +119,7 @@ class PropagationResult(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None completed_task_id: UUID ready_tasks: list[TaskDescriptor] = Field(default_factory=list) workflow_terminal_state: WorkflowTerminalState = WorkflowTerminalState.NONE @@ -129,7 +129,7 @@ class FinalizeWorkflowCommand(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None class FinalizedWorkflowResult(BaseModel): diff --git a/ergon_core/ergon_core/core/application/runtime/sample_identity.py b/ergon_core/ergon_core/core/application/runtime/sample_identity.py index fa0df8a04..1e9358d1e 100644 --- a/ergon_core/ergon_core/core/application/runtime/sample_identity.py +++ b/ergon_core/ergon_core/core/application/runtime/sample_identity.py @@ -9,7 +9,7 @@ from sqlmodel import Session, select -def definition_id_for_run(session: Session, sample_id: UUID) -> UUID: +def definition_id_for_run(session: Session, sample_id: UUID) -> UUID | None: """Return the definition id for a run or fail the runtime invariant loudly.""" run = session.exec(select(SampleRecord).where(SampleRecord.id == sample_id)).first() if run is None: diff --git a/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py b/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py index 7cb0c36a8..ebee9962c 100644 --- a/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py +++ b/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py @@ -105,45 +105,88 @@ def __init__( self._task_execution_repo = TaskExecutionRepository() self._runtime_events = RuntimeEventDispatcher(task_ready_dispatcher) - async def initialize(self, command: InitializeWorkflowCommand) -> InitializedWorkflow: + async def initialize( + self, + command: InitializeWorkflowCommand, + *, + session: Session | None = None, + ) -> InitializedWorkflow: """Load a definition, seed graph state, and return initially ready tasks.""" + if session is not None: + return await self._initialize_in_session(session, command, commit=False) with get_session() as session: - definition = require_not_none( - session.get(ExperimentDefinition, command.definition_id), - f"Definition {command.definition_id} not found", - ) - all_tasks = list( - session.exec( - select(ExperimentDefinitionTask).where( - ExperimentDefinitionTask.experiment_definition_id == command.definition_id, - ) - ).all() - ) + return await self._initialize_in_session(session, command, commit=True) - self._graph_repo.initialize_from_definition( + async def _initialize_in_session( + self, + session: Session, + command: InitializeWorkflowCommand, + *, + commit: bool, + ) -> InitializedWorkflow: + materialized_nodes = list( + session.exec( + select(SampleGraphNode).where(SampleGraphNode.sample_id == command.sample_id) + ).all() + ) + if materialized_nodes: + return await self._initialize_materialized_sample( session, - command.sample_id, - command.definition_id, - initial_node_status=graph_status.PENDING, - initial_edge_status=graph_status.EDGE_PENDING, - meta=MutationMeta(actor="system:workflow_init"), + command, + nodes=materialized_nodes, + commit=commit, ) - session.commit() + if command.definition_id is None: + raise ValueError( + f"Sample {command.sample_id} has no materialized graph and no definition_id" + ) + return await self._initialize_definition_backed_sample(session, command, commit=commit) - task_descriptors = [ - TaskDescriptor( - task_id=t.id, - task_slug=t.task_slug, - parent_task_id=t.parent_task_id, + async def _initialize_definition_backed_sample( + self, + session: Session, + command: InitializeWorkflowCommand, + *, + commit: bool, + ) -> InitializedWorkflow: + definition = require_not_none( + session.get(ExperimentDefinition, command.definition_id), + f"Definition {command.definition_id} not found", + ) + all_tasks = list( + session.exec( + select(ExperimentDefinitionTask).where( + ExperimentDefinitionTask.experiment_definition_id == command.definition_id, ) - for t in all_tasks - ] - graph_lookup = GraphNodeLookup(session, command.sample_id) + ).all() + ) - run_record = require_not_none( - session.get(SampleRecord, command.sample_id), - f"SampleRecord {command.sample_id} not found", + self._graph_repo.initialize_from_definition( + session, + command.sample_id, + command.definition_id, + initial_node_status=graph_status.PENDING, + initial_edge_status=graph_status.EDGE_PENDING, + meta=MutationMeta(actor="system:workflow_init"), + ) + if commit: + session.commit() + + task_descriptors = [ + TaskDescriptor( + task_id=t.id, + task_slug=t.task_slug, + parent_task_id=t.parent_task_id, ) + for t in all_tasks + ] + graph_lookup = GraphNodeLookup(session, command.sample_id) + + run_record = require_not_none( + session.get(SampleRecord, command.sample_id), + f"SampleRecord {command.sample_id} not found", + ) + if run_record.status == SampleStatus.PENDING: run_record.status = SampleStatus.EXECUTING run_record.started_at = utcnow() session.add(run_record) @@ -156,27 +199,96 @@ async def initialize(self, command: InitializeWorkflowCommand) -> InitializedWor event_timestamp=run_record.started_at, ) ) + if commit: session.commit() - ready_ids = await get_initial_ready_tasks( - session, - command.sample_id, - command.definition_id, - graph_repo=self._graph_repo, - graph_lookup=graph_lookup, + ready_ids = await get_initial_ready_tasks( + session, + command.sample_id, + command.definition_id, + graph_repo=self._graph_repo, + graph_lookup=graph_lookup, + commit=commit, + ) + if not commit: + session.flush() + ready_id_set = set(ready_ids) + root_count = sum(1 for t in all_tasks if t.parent_task_id is None) + + return InitializedWorkflow( + sample_id=command.sample_id, + definition_id=command.definition_id, + benchmark_type=definition.benchmark_type, + total_tasks=len(all_tasks), + total_root_tasks=root_count, + pending_tasks=task_descriptors, + initial_ready_tasks=[td for td in task_descriptors if td.task_id in ready_id_set], + ) + + async def _initialize_materialized_sample( + self, + session: Session, + command: InitializeWorkflowCommand, + *, + nodes: list[SampleGraphNode], + commit: bool, + ) -> InitializedWorkflow: + graph_lookup = GraphNodeLookup(session, command.sample_id) + run_record = require_not_none( + session.get(SampleRecord, command.sample_id), + f"SampleRecord {command.sample_id} not found", + ) + if run_record.status == SampleStatus.PENDING: + run_record.status = SampleStatus.EXECUTING + run_record.started_at = utcnow() + session.add(run_record) + SampleRuntimeEventAppender(session).append_status_event( + SampleStatusEventRow( + sample_id=command.sample_id, + event_type="sample.status_changed", + status=SampleStatus.EXECUTING, + actor="system:workflow_init", + event_timestamp=run_record.started_at, + ) ) - ready_id_set = set(ready_ids) - root_count = sum(1 for t in all_tasks if t.parent_task_id is None) + ready_ids = await get_initial_ready_tasks( + session, + command.sample_id, + None, + graph_repo=self._graph_repo, + graph_lookup=graph_lookup, + commit=commit, + ) + if commit: + session.commit() + else: + session.flush() - return InitializedWorkflow( - sample_id=command.sample_id, - definition_id=command.definition_id, - benchmark_type=definition.benchmark_type, - total_tasks=len(all_tasks), - total_root_tasks=root_count, - pending_tasks=task_descriptors, - initial_ready_tasks=[td for td in task_descriptors if td.task_id in ready_id_set], + ready_id_set = set(ready_ids) + dependency_target_ids = set( + session.exec( + select(SampleGraphEdge.target_task_id).where( + SampleGraphEdge.sample_id == command.sample_id + ) + ).all() + ) + task_descriptors = [ + TaskDescriptor( + task_id=node.task_id, + task_slug=node.task_slug, + parent_task_id=node.parent_task_id, ) + for node in sorted(nodes, key=lambda node: (node.level, node.task_slug, str(node.task_id))) + ] + return InitializedWorkflow( + sample_id=command.sample_id, + definition_id=command.definition_id, + benchmark_type=run_record.benchmark_type, + total_tasks=len(nodes), + total_root_tasks=sum(1 for node in nodes if node.task_id not in dependency_target_ids), + pending_tasks=task_descriptors, + initial_ready_tasks=[td for td in task_descriptors if td.task_id in ready_id_set], + ) def finalize(self, command: FinalizeWorkflowCommand) -> FinalizedWorkflowResult: """Aggregate evaluations and close the run.""" diff --git a/ergon_core/ergon_core/core/application/runtime/task_execution.py b/ergon_core/ergon_core/core/application/runtime/task_execution.py index 56034dccb..13235a0c0 100644 --- a/ergon_core/ergon_core/core/application/runtime/task_execution.py +++ b/ergon_core/ergon_core/core/application/runtime/task_execution.py @@ -1,8 +1,10 @@ """Task execution lifecycle: prepare, finalize success, finalize failure.""" import logging +from collections.abc import Mapping from uuid import UUID +from pydantic import JsonValue from ergon_core.core.application.events.service import get_dashboard_event_publisher from ergon_core.core.persistence.definitions.models import ( ExperimentDefinition, @@ -155,18 +157,29 @@ async def _prepare_run_node( sample_id=command.sample_id, task_id=lookup_id, ) - definition = require_not_none( - session.get(ExperimentDefinition, command.definition_id), - f"Definition {command.definition_id} not found", - ) - assigned_worker_slug = node.assigned_worker_slug - worker_type, model_target, definition_worker_id = self._resolve_worker_config( - session, - definition_id=command.definition_id, - sample_id=command.sample_id, - assigned_worker_slug=assigned_worker_slug, + run_record = require_not_none( + session.get(SampleRecord, command.sample_id), + f"SampleRecord {command.sample_id} not found", ) + if command.definition_id is None: + worker_type, model_target, definition_worker_id = _resolve_sample_worker_config( + node.task_json, + assigned_worker_slug=assigned_worker_slug, + ) + benchmark_type = run_record.benchmark_type + else: + definition = require_not_none( + session.get(ExperimentDefinition, command.definition_id), + f"Definition {command.definition_id} not found", + ) + worker_type, model_target, definition_worker_id = self._resolve_worker_config( + session, + definition_id=command.definition_id, + sample_id=command.sample_id, + assigned_worker_slug=assigned_worker_slug, + ) + benchmark_type = definition.benchmark_type execution = SampleTaskAttempt( sample_id=command.sample_id, @@ -183,10 +196,7 @@ async def _prepare_run_node( # Snapshot ORM-derived scalars before commit. SQLAlchemy's # `expire_on_commit=True` default expires every loaded # instance on commit, and `with get_session() as session:` - # closes the session immediately after — so the post-commit - # reads of `definition.benchmark_type` / `execution.id` - # below would raise DetachedInstanceError. - benchmark_type = definition.benchmark_type + # closes the session immediately after. execution_id = execution.id await self._graph_repo.update_node_status( session, @@ -308,3 +318,41 @@ async def finalize_failure(self, command: FailTaskExecutionCommand) -> None: new_status=graph_status.FAILED, old_status=graph_status.RUNNING, ) + + +def _resolve_sample_worker_config( + task_json: Mapping[str, JsonValue], + *, + assigned_worker_slug: str | None, +) -> tuple[str | None, str | None, None]: + worker_snapshot = _component_snapshot(task_json.get("worker")) + if worker_snapshot is None: + return assigned_worker_slug, None, None + worker_slug = assigned_worker_slug or _component_slug(worker_snapshot, fallback="worker") + return ( + _component_type(worker_snapshot, fallback=worker_slug), + _component_model_target(worker_snapshot), + None, + ) + + +def _component_snapshot(value: JsonValue | None) -> Mapping[str, JsonValue] | None: + return value if isinstance(value, dict) else None + + +def _component_slug(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: + value = snapshot.get("type_slug") or snapshot.get("slug") or snapshot.get("name") + if isinstance(value, str) and value: + return value + component_type = _component_type(snapshot, fallback=fallback) + return component_type.rsplit(":", 1)[-1].rsplit(".", 1)[-1] + + +def _component_type(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: + value = snapshot.get("_type") or snapshot.get("type") + return value if isinstance(value, str) and value else fallback + + +def _component_model_target(snapshot: Mapping[str, JsonValue]) -> str | None: + value = snapshot.get("model") or snapshot.get("model_target") + return value if isinstance(value, str) else None diff --git a/ergon_core/ergon_core/core/application/runtime/task_management.py b/ergon_core/ergon_core/core/application/runtime/task_management.py index fe8572e52..de5a5a489 100644 --- a/ergon_core/ergon_core/core/application/runtime/task_management.py +++ b/ergon_core/ergon_core/core/application/runtime/task_management.py @@ -125,7 +125,7 @@ async def spawn_dynamic_task( the full Task snapshot lives in sample_graph_nodes.task_json with is_dynamic=True. """ - dispatch: tuple[UUID, UUID, UUID] | None = None + dispatch: tuple[UUID, UUID | None, UUID] | None = None with get_session() as session: parent = self._graph_repo.get_node(session, sample_id=sample_id, task_id=parent_task_id) node = await self._graph_repo.add_node( @@ -240,7 +240,7 @@ async def cancel_orphans( session: Session, *, sample_id: UUID, - definition_id: UUID, + definition_id: UUID | None, parent_task_id: UUID, cause: PropagationCancelCause, ) -> CancelOrphansResult: @@ -636,7 +636,7 @@ def _task_cancelled_event( session: Session, *, sample_id: UUID, - definition_id: UUID, + definition_id: UUID | None, task_id: UUID, cause: CancelCause, ) -> TaskCancelledEvent: diff --git a/ergon_core/ergon_core/core/application/samples/materialization.py b/ergon_core/ergon_core/core/application/samples/materialization.py new file mode 100644 index 000000000..fb1aeaf36 --- /dev/null +++ b/ergon_core/ergon_core/core/application/samples/materialization.py @@ -0,0 +1,189 @@ +"""Materialize authored Samples into typed runtime WAL and graph projections.""" + +from collections.abc import Mapping +from uuid import UUID, uuid4 + +from pydantic import JsonValue +from sqlmodel import Session + +from ergon_core.api.experiment.sample import Sample +from ergon_core.core.application.runtime import status as graph_status +from ergon_core.core.application.samples.events import SampleRuntimeEventAppender +from ergon_core.core.persistence.graph.models import SampleGraphEdge, SampleGraphNode +from ergon_core.core.persistence.samples.models import ( + SampleEdgeEventRow, + SampleEvaluatorEventRow, + SampleSandboxEventRow, + SampleStatusEventRow, + SampleTaskEventRow, + SampleWorkerEventRow, +) +from ergon_core.core.persistence.shared.enums import SampleStatus +from ergon_core.core.persistence.telemetry.models import SampleRecord + + +def component_slug(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: + value = snapshot.get("type_slug") or snapshot.get("slug") or snapshot.get("name") + if isinstance(value, str) and value: + return value + component_type = component_type_path(snapshot, fallback=fallback) + return component_type.rsplit(":", 1)[-1].rsplit(".", 1)[-1] + + +def component_type_path(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: + value = snapshot.get("_type") or snapshot.get("type") + return value if isinstance(value, str) and value else fallback + + +def component_model_target(snapshot: Mapping[str, JsonValue]) -> str | None: + value = snapshot.get("model") or snapshot.get("model_target") + return value if isinstance(value, str) else None + + +def materialize_sample( + session: Session, + *, + sample: Sample, + sample_row: SampleRecord, +) -> None: + wal = SampleRuntimeEventAppender(session) + wal.append_status_event( + SampleStatusEventRow( + sample_id=sample_row.id, + event_type="sample.status_changed", + status=SampleStatus.PENDING, + payload_json={"reason": "materialized"}, + actor="system:materialization", + ) + ) + + task_ids_by_key: dict[str, UUID] = {} + for task in sample.tasks: + task_id = uuid4() + task_ids_by_key[task.task_slug] = task_id + task_json = task.model_dump(mode="json") + worker_snapshot = _required_component_snapshot(task_json, "worker") + sandbox_snapshot = _required_component_snapshot(task_json, "sandbox") + worker_slug = component_slug(worker_snapshot, fallback="worker") + sandbox_slug = component_slug(sandbox_snapshot, fallback="sandbox") + + session.add( + SampleGraphNode( + sample_id=sample_row.id, + task_id=task_id, + instance_key=task.instance_key, + task_slug=task.task_slug, + description=task.description, + task_json=task_json, + is_dynamic=False, + status=graph_status.PENDING, + assigned_worker_slug=worker_slug, + ) + ) + task_event = wal.append_task_event( + SampleTaskEventRow( + sample_id=sample_row.id, + task_id=task_id, + event_type="task.added", + task_slug=task.task_slug, + status=graph_status.PENDING, + task_snapshot_json=task_json, + payload_json={ + "task_id": str(task_id), + "task_slug": task.task_slug, + "worker_slug": worker_slug, + "sandbox_slug": sandbox_slug, + }, + actor="system:materialization", + ) + ) + wal.append_worker_event( + SampleWorkerEventRow( + sample_id=sample_row.id, + task_id=task_id, + event_type="worker.added", + worker_slug=worker_slug, + worker_type=component_type_path(worker_snapshot, fallback=worker_slug), + model_target=component_model_target(worker_snapshot), + worker_snapshot_json=worker_snapshot, + payload_json={"task_id": str(task_id), "worker": worker_snapshot}, + actor="system:materialization", + ) + ) + wal.append_sandbox_event( + SampleSandboxEventRow( + sample_id=sample_row.id, + task_id=task_id, + event_type="sandbox.added", + sandbox_slug=sandbox_slug, + sandbox_type=component_type_path(sandbox_snapshot, fallback=sandbox_slug), + sandbox_snapshot_json=sandbox_snapshot, + payload_json={"task_id": str(task_id), "sandbox": sandbox_snapshot}, + actor="system:materialization", + ) + ) + for evaluator_snapshot in _evaluator_snapshots(task_json): + evaluator_slug = component_slug(evaluator_snapshot, fallback="default") + wal.append_evaluator_event( + SampleEvaluatorEventRow( + sample_id=sample_row.id, + task_id=task_id, + event_type="evaluator.added", + evaluator_slug=evaluator_slug, + evaluator_type=component_type_path(evaluator_snapshot, fallback=evaluator_slug), + evaluator_snapshot_json=evaluator_snapshot, + payload_json={"task_id": str(task_id), "evaluator": evaluator_snapshot}, + actor="system:materialization", + ) + ) + session.add(task_event) + + session.flush() + for task in sample.tasks: + target_task_id = task_ids_by_key[task.task_slug] + for dependency_key in task.dependency_task_slugs: + source_task_id = task_ids_by_key[dependency_key] + edge = SampleGraphEdge( + sample_id=sample_row.id, + source_task_id=source_task_id, + target_task_id=target_task_id, + status=graph_status.EDGE_PENDING, + ) + session.add(edge) + session.flush() + wal.append_edge_event( + SampleEdgeEventRow( + sample_id=sample_row.id, + edge_id=edge.id, + event_type="edge.added", + source_task_id=source_task_id, + target_task_id=target_task_id, + status=graph_status.EDGE_PENDING, + edge_snapshot_json={ + "source_task_slug": dependency_key, + "target_task_slug": task.task_slug, + }, + payload_json={ + "source_task_id": str(source_task_id), + "target_task_id": str(target_task_id), + "source_task_slug": dependency_key, + "target_task_slug": task.task_slug, + }, + actor="system:materialization", + ) + ) + session.flush() + + +def _required_component_snapshot(task_json: Mapping[str, JsonValue], key: str) -> dict: + value = task_json.get(key) + if not isinstance(value, dict): + raise ValueError(f"Task snapshot is missing object-bound {key}") + return value + + +def _evaluator_snapshots(task_json: Mapping[str, JsonValue]) -> list[dict]: + evaluators = task_json.get("evaluators", []) + if not isinstance(evaluators, list): + return [] + return [snapshot for snapshot in evaluators if isinstance(snapshot, dict)] diff --git a/ergon_core/ergon_core/core/jobs/resources/persist_outputs/contract.py b/ergon_core/ergon_core/core/jobs/resources/persist_outputs/contract.py index af0a37e7f..86020e367 100644 --- a/ergon_core/ergon_core/core/jobs/resources/persist_outputs/contract.py +++ b/ergon_core/ergon_core/core/jobs/resources/persist_outputs/contract.py @@ -10,7 +10,7 @@ class PersistOutputsRequest(InngestEventContract): name: ClassVar[str] = "task/persist-outputs" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID sandbox_id: str | None = None diff --git a/ergon_core/ergon_core/core/jobs/sandbox/setup/contract.py b/ergon_core/ergon_core/core/jobs/sandbox/setup/contract.py index 29775f336..33e779b5d 100644 --- a/ergon_core/ergon_core/core/jobs/sandbox/setup/contract.py +++ b/ergon_core/ergon_core/core/jobs/sandbox/setup/contract.py @@ -10,7 +10,7 @@ class SandboxSetupRequest(InngestEventContract): name: ClassVar[str] = "task/sandbox-setup" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID benchmark_type: str sandbox_slug: str | None = None diff --git a/ergon_core/ergon_core/core/jobs/task/cancel_orphans/job.py b/ergon_core/ergon_core/core/jobs/task/cancel_orphans/job.py index 15ab6b3d5..606e0d1aa 100644 --- a/ergon_core/ergon_core/core/jobs/task/cancel_orphans/job.py +++ b/ergon_core/ergon_core/core/jobs/task/cancel_orphans/job.py @@ -26,7 +26,7 @@ async def _cancel_orphans_for( ctx: Any, *, sample_id: UUID, - definition_id: UUID, + definition_id: UUID | None, parent_task_id: UUID, cause: PropagationCancelCause, ) -> int: diff --git a/ergon_core/ergon_core/core/jobs/task/worker_execute/contract.py b/ergon_core/ergon_core/core/jobs/task/worker_execute/contract.py index 5b6d80b60..f46131ece 100644 --- a/ergon_core/ergon_core/core/jobs/task/worker_execute/contract.py +++ b/ergon_core/ergon_core/core/jobs/task/worker_execute/contract.py @@ -11,7 +11,7 @@ class WorkerExecuteRequest(InngestEventContract): name: ClassVar[str] = "task/worker-execute" sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID execution_id: UUID sandbox_id: str diff --git a/ergon_core/ergon_core/core/jobs/task/worker_execute/job.py b/ergon_core/ergon_core/core/jobs/task/worker_execute/job.py index 3080bb88c..4ab5e88ae 100644 --- a/ergon_core/ergon_core/core/jobs/task/worker_execute/job.py +++ b/ergon_core/ergon_core/core/jobs/task/worker_execute/job.py @@ -208,7 +208,7 @@ class _ReadyDispatch(BaseModel): model_config = {"frozen": True} sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None task_id: UUID @@ -272,7 +272,7 @@ async def _run_spawn() -> _SpawnTaskStepResult: async def _collect_ready_dispatch( self, sample_id: UUID, - definition_id: UUID, + definition_id: UUID | None, task_id: UUID, ) -> None: if self._active_ready_dispatches is None: diff --git a/ergon_core/ergon_core/core/jobs/workflow/start/job.py b/ergon_core/ergon_core/core/jobs/workflow/start/job.py index 265bc571c..ba810b4da 100644 --- a/ergon_core/ergon_core/core/jobs/workflow/start/job.py +++ b/ergon_core/ergon_core/core/jobs/workflow/start/job.py @@ -9,6 +9,7 @@ from ergon_core.core.application.runtime.orchestration import InitializeWorkflowCommand from ergon_core.core.application.runtime.sample_lifecycle import WorkflowService from ergon_core.core.jobs._events import send_job_events +from sqlmodel import Session from ergon_core.core.infrastructure.tracing import ( CompletedSpan, get_trace_sink, @@ -49,21 +50,22 @@ async def run_start_workflow_job(payload: WorkflowStartedEvent) -> WorkflowStart await send_job_events(events) - snapshot = SampleSnapshotReadService().build_snapshot(payload.sample_id) - if snapshot is None: - raise RuntimeError(f"Run snapshot {payload.sample_id} not found after workflow start") + if payload.definition_id is not None: + snapshot = SampleSnapshotReadService().build_snapshot(payload.sample_id) + if snapshot is None: + raise RuntimeError(f"Run snapshot {payload.sample_id} not found after workflow start") - await get_dashboard_event_publisher().publish( - DashboardWorkflowStartedEvent( - sample_id=payload.sample_id, - definition_id=payload.definition_id, - workflow_name=initialized.benchmark_type, - snapshot=snapshot, - started_at=snapshot.started_at or utcnow(), - total_tasks=snapshot.total_tasks, - total_leaf_tasks=snapshot.total_leaf_tasks, + await get_dashboard_event_publisher().publish( + DashboardWorkflowStartedEvent( + sample_id=payload.sample_id, + definition_id=payload.definition_id, + workflow_name=initialized.benchmark_type, + snapshot=snapshot, + started_at=snapshot.started_at or utcnow(), + total_tasks=snapshot.total_tasks, + total_leaf_tasks=snapshot.total_leaf_tasks, + ) ) - ) result = WorkflowStartResult( sample_id=payload.sample_id, @@ -92,3 +94,23 @@ async def run_start_workflow_job(payload: WorkflowStartedEvent) -> WorkflowStart result.total_tasks, ) return result + + +async def run_workflow_start_job( + *, + session: Session, + event: WorkflowStartedEvent, +) -> WorkflowStartResult: + svc = WorkflowService() + initialized = await svc.initialize( + InitializeWorkflowCommand( + sample_id=event.sample_id, + definition_id=event.definition_id, + ), + session=session, + ) + return WorkflowStartResult( + sample_id=event.sample_id, + initial_ready_tasks=len(initialized.initial_ready_tasks), + total_tasks=initialized.total_tasks, + ) diff --git a/ergon_core/ergon_core/core/persistence/telemetry/models.py b/ergon_core/ergon_core/core/persistence/telemetry/models.py index 33c2f76c1..4176df6c2 100644 --- a/ergon_core/ergon_core/core/persistence/telemetry/models.py +++ b/ergon_core/ergon_core/core/persistence/telemetry/models.py @@ -34,11 +34,18 @@ class SampleRecord(SQLModel, table=True): __tablename__ = "samples" id: UUID = Field(default_factory=uuid4, primary_key=True) - definition_id: UUID = Field( + definition_id: UUID | None = Field( + default=None, foreign_key="experiment_definitions.id", index=True, description="Canonical runtime ExperimentDefinition id for this run.", ) + experiment_id: UUID | None = Field(default=None, index=True) + environment_id: UUID | None = Field(default=None, index=True) + sampler_invocation_id: UUID | None = Field(default=None, index=True) + pool_entry_id: UUID | None = Field(default=None, index=True) + sample_key: str | None = Field(default=None, index=True) + sample_ref_json: dict = Field(default_factory=dict, sa_column=Column(JSON)) benchmark_type: str = Field(index=True) instance_key: str = Field(index=True) sample_id: str | None = Field(default=None, index=True) diff --git a/ergon_core/migrations/versions/00000002_add_sample_experiment_provenance.py b/ergon_core/migrations/versions/00000002_add_sample_experiment_provenance.py new file mode 100644 index 000000000..c20df30bf --- /dev/null +++ b/ergon_core/migrations/versions/00000002_add_sample_experiment_provenance.py @@ -0,0 +1,96 @@ +"""Add experiment provenance to materialized samples. + +Revision ID: 00000002 +Revises: 00000001 +Create Date: 2026-05-26 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op +from sqlalchemy import inspect + + +revision = "00000002" +down_revision = "00000001" +branch_labels = None +depends_on = None + + +_NEW_COLUMNS = ( + sa.Column("experiment_id", sa.Uuid(), nullable=True), + sa.Column("environment_id", sa.Uuid(), nullable=True), + sa.Column("sampler_invocation_id", sa.Uuid(), nullable=True), + sa.Column("pool_entry_id", sa.Uuid(), nullable=True), + sa.Column("sample_key", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("sample_ref_json", sa.JSON(), nullable=True), +) + +_NEW_INDEXES = ( + ("ix_samples_experiment_id", ["experiment_id"]), + ("ix_samples_environment_id", ["environment_id"]), + ("ix_samples_sampler_invocation_id", ["sampler_invocation_id"]), + ("ix_samples_pool_entry_id", ["pool_entry_id"]), + ("ix_samples_sample_key", ["sample_key"]), +) + + +def upgrade() -> None: + context = op.get_context() + if context.as_sql: + _upgrade_unconditional() + return + + bind = op.get_bind() + existing_columns = {column["name"] for column in inspect(bind).get_columns("samples")} + if bind.dialect.name == "sqlite": + with op.batch_alter_table("samples") as batch_op: + batch_op.alter_column("definition_id", existing_type=sa.Uuid(), nullable=True) + else: + op.alter_column("samples", "definition_id", existing_type=sa.Uuid(), nullable=True) + for column in _NEW_COLUMNS: + if column.name not in existing_columns: + op.add_column("samples", column.copy()) + + existing_indexes = {index["name"] for index in inspect(bind).get_indexes("samples")} + for index_name, columns in _NEW_INDEXES: + if index_name not in existing_indexes: + op.create_index(index_name, "samples", columns) + + +def downgrade() -> None: + context = op.get_context() + if context.as_sql: + _downgrade_unconditional() + return + + bind = op.get_bind() + existing_columns = {column["name"] for column in inspect(bind).get_columns("samples")} + existing_indexes = {index["name"] for index in inspect(bind).get_indexes("samples")} + for index_name, _columns in reversed(_NEW_INDEXES): + if index_name in existing_indexes: + op.drop_index(index_name, table_name="samples") + for column in reversed(_NEW_COLUMNS): + if column.name in existing_columns: + op.drop_column("samples", column.name) + if bind.dialect.name == "sqlite": + with op.batch_alter_table("samples") as batch_op: + batch_op.alter_column("definition_id", existing_type=sa.Uuid(), nullable=False) + else: + op.alter_column("samples", "definition_id", existing_type=sa.Uuid(), nullable=False) + + +def _upgrade_unconditional() -> None: + op.alter_column("samples", "definition_id", existing_type=sa.Uuid(), nullable=True) + for column in _NEW_COLUMNS: + op.add_column("samples", column.copy()) + for index_name, columns in _NEW_INDEXES: + op.create_index(index_name, "samples", columns) + + +def _downgrade_unconditional() -> None: + for index_name, _columns in reversed(_NEW_INDEXES): + op.drop_index(index_name, table_name="samples") + for column in reversed(_NEW_COLUMNS): + op.drop_column("samples", column.name) + op.alter_column("samples", "definition_id", existing_type=sa.Uuid(), nullable=False) diff --git a/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py b/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py index 0bb0bd055..4fea078f5 100644 --- a/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py +++ b/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py @@ -3,17 +3,17 @@ from uuid import uuid4 import pytest -from sqlmodel import select +from sqlmodel import Session, SQLModel, create_engine, select +import ergon_core.core.persistence.definitions.models # noqa: F401 +import ergon_core.core.persistence.experiments.models # noqa: F401 +import ergon_core.core.persistence.graph.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, RandomSampler, Sample from ergon_core.core.persistence.samples.models import SampleTaskEventRow from ergon_core.test_support.task_factory import task_with_id -pytestmark = pytest.mark.xfail( - strict=True, - reason="PR05 implements ExperimentSubmissionService materialization and runtime start", -) - class MaterializedEnvironment(Environment): source_mode: Literal["materialized"] = "materialized" @@ -34,6 +34,14 @@ def iter_samples(self) -> Iterator[Sample]: ) +@pytest.fixture() +def session() -> Iterator[Session]: + engine = create_engine("sqlite:///:memory:") + SQLModel.metadata.create_all(engine) + with Session(engine) as session: + yield session + + @pytest.mark.asyncio async def test_public_experiment_submit_materializes_samples_and_typed_wal(session) -> None: from ergon_core.core.application.experiments.submission import ExperimentSubmissionService diff --git a/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py b/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py new file mode 100644 index 000000000..49d2452f6 --- /dev/null +++ b/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py @@ -0,0 +1,250 @@ +from collections.abc import Iterable, Iterator +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.experiments.models # noqa: F401 +import ergon_core.core.persistence.graph.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 Sample +from ergon_core.api.benchmark import Task +from ergon_core.api.criterion.outcome import CriterionOutcome +from ergon_core.api.rubric.evaluator import Evaluator +from ergon_core.api.rubric.results import TaskEvaluationResult +from ergon_core.core.application.events.runtime import TaskCancelledEvent, WorkflowStartedEvent +from ergon_core.core.application.runtime.lifecycle import get_initial_ready_tasks +from ergon_core.core.application.runtime.orchestration import ( + InitializeWorkflowCommand, + PrepareTaskExecutionCommand, +) +from ergon_core.core.application.runtime.sample_lifecycle import WorkflowService +from ergon_core.core.application.samples.materialization import materialize_sample +from ergon_core.core.jobs.resources.persist_outputs.contract import PersistOutputsRequest +from ergon_core.core.jobs.sandbox.setup.contract import SandboxSetupRequest +from ergon_core.core.jobs.task.worker_execute.contract import WorkerExecuteRequest +from ergon_core.core.jobs.workflow.start.job import run_workflow_start_job +from ergon_core.core.persistence.graph.models import SampleGraphNode +from ergon_core.core.persistence.samples.models import SampleStatusEventRow, SampleTaskEventRow +from ergon_core.core.persistence.shared.enums import SampleStatus +from ergon_core.core.persistence.telemetry.models import SampleRecord +from ergon_core.test_support.task_factory import task_with_id + + +class StartPathEvaluator(Evaluator): + type_slug = "start-test-evaluator" + + def criteria_for(self, task: Task) -> Iterable: + return () + + def aggregate_task( + self, + task: Task, + criterion_results: Iterable[CriterionOutcome], + ) -> TaskEvaluationResult: + return TaskEvaluationResult( + task_slug=task.task_slug, + score=1.0, + passed=True, + evaluator_name=self.name, + criterion_results=list(criterion_results), + ) + + +@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 materialized_sample(session: Session) -> SampleRecord: + root = task_with_id( + uuid4(), + task_slug="root", + instance_key="a", + description="Root task", + evaluators=(StartPathEvaluator(name="judge"),), + ) + child = task_with_id( + uuid4(), + task_slug="child", + instance_key="a", + description="Child task", + dependency_task_slugs=("root",), + ) + sample = Sample.from_tasks( + name="mini:a", + sample_key="a", + environment_name="mini", + tasks=[root, child], + ) + sample_row = SampleRecord( + benchmark_type="experiment", + instance_key="a", + status=SampleStatus.PENDING, + ) + session.add(sample_row) + session.flush() + materialize_sample(session=session, sample=sample, sample_row=sample_row) + session.commit() + return sample_row + + +@pytest.mark.asyncio +async def test_workflow_start_uses_materialized_sample_graph_without_definition( + session: Session, + materialized_sample: SampleRecord, +) -> None: + result = await run_workflow_start_job( + session=session, + event=WorkflowStartedEvent(sample_id=materialized_sample.id), + ) + assert result.initial_ready_tasks == 1 + + ready = await get_initial_ready_tasks( + session=session, + sample_id=materialized_sample.id, + definition_id=None, + ) + + nodes_by_id = { + node.task_id: node + for node in session.exec( + select(SampleGraphNode).where(SampleGraphNode.sample_id == materialized_sample.id) + ) + } + assert ready == [] + assert [node.task_slug for node in nodes_by_id.values() if node.status == "ready"] == ["root"] + + +def test_sample_only_task_execution_contracts_allow_missing_definition_id() -> None: + sample_id = uuid4() + task_id = uuid4() + execution_id = uuid4() + + assert PrepareTaskExecutionCommand(sample_id=sample_id, task_id=task_id).definition_id is None + assert ( + SandboxSetupRequest( + sample_id=sample_id, + task_id=task_id, + benchmark_type="experiment", + ).definition_id + is None + ) + assert ( + WorkerExecuteRequest( + sample_id=sample_id, + task_id=task_id, + execution_id=execution_id, + sandbox_id="sandbox", + task_slug="root", + task_description="Root task", + assigned_worker_slug="worker", + worker_type="worker.Type", + benchmark_type="experiment", + ).definition_id + is None + ) + assert ( + PersistOutputsRequest( + sample_id=sample_id, + task_id=task_id, + execution_id=execution_id, + benchmark_type="experiment", + ).definition_id + is None + ) + assert ( + TaskCancelledEvent( + sample_id=sample_id, + task_id=task_id, + execution_id=None, + cause="manager_decision", + ).definition_id + is None + ) + + +@pytest.mark.asyncio +async def test_materialized_workflow_start_is_idempotent( + session: Session, + materialized_sample: SampleRecord, +) -> None: + service = WorkflowService() + + first = await service.initialize( + InitializeWorkflowCommand(sample_id=materialized_sample.id), + session=session, + ) + first_ready_ids = [task.task_id for task in first.initial_ready_tasks] + assert first_ready_ids + root_id = first_ready_ids[0] + root = session.get(SampleGraphNode, (materialized_sample.id, root_id)) + assert root is not None + root.status = "running" + session.add(root) + session.commit() + + second = await service.initialize( + InitializeWorkflowCommand(sample_id=materialized_sample.id), + session=session, + ) + + assert second.initial_ready_tasks == [] + assert second.total_root_tasks == first.total_root_tasks == 1 + root_after_retry = session.get(SampleGraphNode, (materialized_sample.id, root_id)) + assert root_after_retry is not None + assert root_after_retry.status == "running" + assert ( + len( + session.exec( + select(SampleStatusEventRow) + .where(SampleStatusEventRow.sample_id == materialized_sample.id) + .where(SampleStatusEventRow.status == SampleStatus.EXECUTING) + ).all() + ) + == 1 + ) + task_events = session.exec( + select(SampleTaskEventRow) + .where(SampleTaskEventRow.sample_id == materialized_sample.id) + .where(SampleTaskEventRow.task_id == root_id) + .where(SampleTaskEventRow.event_type == "task.status_changed") + ).all() + assert len(task_events) == 1 + + +@pytest.mark.asyncio +async def test_materialized_workflow_start_retry_before_worker_claim_is_idempotent( + session: Session, + materialized_sample: SampleRecord, +) -> None: + service = WorkflowService() + + first = await service.initialize( + InitializeWorkflowCommand(sample_id=materialized_sample.id), + session=session, + ) + first_ready_ids = [task.task_id for task in first.initial_ready_tasks] + assert first_ready_ids + + second = await service.initialize( + InitializeWorkflowCommand(sample_id=materialized_sample.id), + session=session, + ) + + assert second.initial_ready_tasks == [] + assert second.total_root_tasks == first.total_root_tasks == 1 + task_events = session.exec( + select(SampleTaskEventRow) + .where(SampleTaskEventRow.sample_id == materialized_sample.id) + .where(SampleTaskEventRow.task_id == first_ready_ids[0]) + .where(SampleTaskEventRow.event_type == "task.status_changed") + ).all() + assert len(task_events) == 1 + assert task_events[0].status == "ready" diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py new file mode 100644 index 000000000..54f995371 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py @@ -0,0 +1,310 @@ +from collections.abc import Iterator, Sequence +from pathlib import Path +from types import SimpleNamespace +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.experiments.models # noqa: F401 +import ergon_core.core.persistence.graph.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, RandomSampler, Sample +from ergon_core.api.experiment.sampling import SamplingContext +from ergon_core.core.application.events.runtime import WorkflowStartedEvent +from ergon_core.core.application.experiments.submission import ( + ExperimentSubmissionService, + InngestWorkflowEventBus, +) +from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow +from ergon_core.core.persistence.graph.models import SampleGraphNode +from ergon_core.core.persistence.samples.models import SampleEdgeEventRow, SampleTaskEventRow +from ergon_core.core.persistence.telemetry.models import SampleRecord +from ergon_core.test_support.task_factory import task_with_id + + +class FakeEventBus: + def __init__(self) -> None: + self.events: list[SimpleNamespace] = [] + + async def publish(self, event: object) -> None: + name = getattr(event, "name") + payload = event.model_dump(mode="json") + self.events.append(SimpleNamespace(name=name, payload=payload)) + + +class CommittedStateEventBus: + def __init__(self, engine) -> None: + self._engine = engine + self.events: list[WorkflowStartedEvent] = [] + + async def publish(self, event: WorkflowStartedEvent) -> None: + with Session(self._engine) as observer: + assert observer.get(SampleRecord, event.sample_id) is not None + assert observer.exec( + select(SampleGraphNode).where(SampleGraphNode.sample_id == event.sample_id) + ).first() + assert observer.exec( + select(SampleTaskEventRow).where(SampleTaskEventRow.sample_id == event.sample_id) + ).first() + self.events.append(event) + + +class SequentialSampler: + name = "sequential" + + def config(self) -> dict: + return {} + + async def select( + self, + *, + samples: Sequence[Sample], + k: int, + context: SamplingContext, + ) -> Sequence[Sample]: + return list(samples)[:k] + + +def make_task(task_slug: str, key: str, dependencies: tuple[str, ...] = ()): + return task_with_id( + uuid4(), + task_slug=task_slug, + instance_key=key, + description=f"Solve {key}", + dependency_task_slugs=dependencies, + ) + + +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=[ + make_task("root", key), + make_task("child", key, ("root",)), + ], + 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 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 experiment() -> Experiment: + return Experiment( + name="submit smoke", + environments=[MaterializedEnvironment(name="mini-validation", keys=("a", "b", "c"))], + ) + + +@pytest.fixture() +def streaming_experiment() -> Experiment: + return Experiment( + name="streaming submit", + environments=[StreamingEnvironment(name="stream", total=20)], + ) + + +@pytest.fixture() +def two_env_experiment() -> Experiment: + return Experiment( + name="two env submit", + environments=[ + MaterializedEnvironment(name="mini-validation", keys=("a", "b")), + MaterializedEnvironment(name="swe-validation", keys=("1", "2")), + ], + ) + + +@pytest.mark.asyncio +async def test_submit_selects_and_materializes_k_samples( + session: Session, + experiment: Experiment, +) -> None: + service = ExperimentSubmissionService(session=session, event_bus=FakeEventBus()) + + result = await experiment.submit(service=service, k=3, sampler=SequentialSampler()) + + assert result.selected_count == 3 + assert len(result.sample_ids) == 3 + assert not hasattr(result, "run_ids") + assert session.exec(select(SampleRecord)).all() + assert session.exec( + select(SampleTaskEventRow).where(SampleTaskEventRow.event_type == "task.added") + ).all() + assert session.exec( + select(SampleEdgeEventRow).where(SampleEdgeEventRow.event_type == "edge.added") + ).all() + + +@pytest.mark.asyncio +async def test_submit_retains_unselected_candidate_pool_entries( + session: Session, + streaming_experiment: Experiment, +) -> None: + service = ExperimentSubmissionService(session=session, event_bus=FakeEventBus()) + + result = await streaming_experiment.submit( + service=service, + k=2, + sampler=SequentialSampler(), + candidate_pool_size=8, + ) + + rows = session.exec(select(ExperimentSamplePoolEntryRow)).all() + assert result.selected_count == 2 + assert len(rows) == 8 + assert sum(row.selected for row in rows) == 2 + + +@pytest.mark.asyncio +async def test_submit_caps_random_sampler_selection_to_requested_k( + session: Session, + streaming_experiment: Experiment, +) -> None: + service = ExperimentSubmissionService(session=session, event_bus=FakeEventBus()) + + result = await streaming_experiment.submit( + service=service, + k=2, + sampler=RandomSampler(seed=7), + candidate_pool_size=8, + ) + + pool_rows = session.exec(select(ExperimentSamplePoolEntryRow)).all() + sample_rows = session.exec(select(SampleRecord)).all() + assert result.selected_count == 2 + assert len(result.sample_ids) == 2 + assert len(pool_rows) == 8 + assert sum(row.selected for row in pool_rows) == 2 + assert len(sample_rows) == 2 + + +@pytest.mark.asyncio +async def test_submit_records_selected_sample_provenance( + session: Session, + two_env_experiment: Experiment, +) -> None: + service = ExperimentSubmissionService(session=session, event_bus=FakeEventBus()) + + result = await two_env_experiment.submit(service=service, k=2, sampler=SequentialSampler()) + + sample_rows = session.exec(select(SampleRecord).order_by(SampleRecord.created_at)).all() + assert [row.id for row in sample_rows] == list(result.sample_ids) + assert {row.experiment_id for row in sample_rows} == {result.experiment_id} + assert all(row.environment_id for row in sample_rows) + assert all(row.pool_entry_id for row in sample_rows) + assert all(row.sampler_invocation_id == result.sampler_invocation_id for row in sample_rows) + assert all(row.sample_key for row in sample_rows) + assert all(isinstance(row.sample_ref_json, dict) for row in sample_rows) + + +@pytest.mark.asyncio +async def test_submit_emits_sample_start_events_without_definition_or_run_ids( + session: Session, + experiment: Experiment, +) -> None: + event_bus = FakeEventBus() + service = ExperimentSubmissionService(session=session, event_bus=event_bus) + + result = await experiment.submit(service=service, k=1, sampler=SequentialSampler()) + + assert event_bus.events + assert event_bus.events[0].payload["sample_id"] == str(result.sample_ids[0]) + assert "definition_id" not in event_bus.events[0].payload + assert "run_id" not in event_bus.events[0].payload + + +@pytest.mark.asyncio +async def test_submit_commits_materialized_sample_before_start_event(tmp_path: Path) -> None: + engine = create_engine(f"sqlite:///{tmp_path / 'submit.db'}") + SQLModel.metadata.create_all(engine) + event_bus = CommittedStateEventBus(engine) + + with Session(engine) as session: + experiment = Experiment( + name="submit commit boundary", + environments=[MaterializedEnvironment(name="mini-validation", keys=("a",))], + ) + service = ExperimentSubmissionService(session=session, event_bus=event_bus) + + result = await experiment.submit(service=service, k=1, sampler=SequentialSampler()) + + assert [event.sample_id for event in event_bus.events] == list(result.sample_ids) + + +@pytest.mark.asyncio +async def test_default_event_bus_sends_workflow_started(monkeypatch: pytest.MonkeyPatch) -> None: + sent: list[object] = [] + + async def fake_send(event: object) -> None: + sent.append(event) + + monkeypatch.setattr( + "ergon_core.core.application.experiments.submission.inngest_client.send", + fake_send, + ) + event = WorkflowStartedEvent(sample_id=uuid4()) + + await InngestWorkflowEventBus().publish(event) + + assert len(sent) == 1 + assert getattr(sent[0], "name") == WorkflowStartedEvent.name + assert getattr(sent[0], "data") == event.model_dump(mode="json") + + +@pytest.mark.asyncio +async def test_submit_uses_inngest_event_bus_by_default( + session: Session, + experiment: Experiment, + monkeypatch: pytest.MonkeyPatch, +) -> None: + sent: list[object] = [] + + async def fake_send(event: object) -> None: + sent.append(event) + + monkeypatch.setattr( + "ergon_core.core.application.experiments.submission.inngest_client.send", + fake_send, + ) + service = ExperimentSubmissionService(session=session) + + result = await experiment.submit(service=service, k=1, sampler=SequentialSampler()) + + assert len(sent) == 1 + assert getattr(sent[0], "name") == WorkflowStartedEvent.name + assert getattr(sent[0], "data") == { + "sample_id": str(result.sample_ids[0]), + } diff --git a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py new file mode 100644 index 000000000..766460fb3 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py @@ -0,0 +1,140 @@ +from collections.abc import Iterable, Iterator +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.experiments.models # noqa: F401 +import ergon_core.core.persistence.graph.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 Sample +from ergon_core.api.benchmark import Task +from ergon_core.api.criterion.outcome import CriterionOutcome +from ergon_core.api.rubric.evaluator import Evaluator +from ergon_core.api.rubric.results import TaskEvaluationResult +from ergon_core.core.application.samples.materialization import materialize_sample +from ergon_core.core.persistence.graph.models import SampleGraphEdge, SampleGraphNode +from ergon_core.core.persistence.samples.models import ( + SampleEvaluatorEventRow, + SampleTaskEventRow, + SampleWorkerEventRow, +) +from ergon_core.core.persistence.shared.enums import SampleStatus +from ergon_core.core.persistence.telemetry.models import SampleRecord +from ergon_core.test_support.task_factory import task_with_id + + +class MaterializationEvaluator(Evaluator): + type_slug = "test-evaluator" + + def criteria_for(self, task: Task) -> Iterable: + return () + + def aggregate_task( + self, + task: Task, + criterion_results: Iterable[CriterionOutcome], + ) -> TaskEvaluationResult: + return TaskEvaluationResult( + task_slug=task.task_slug, + score=1.0, + passed=True, + evaluator_name=self.name, + criterion_results=list(criterion_results), + ) + + +@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 selected_sample() -> Sample: + root = task_with_id( + uuid4(), + task_slug="root", + instance_key="a", + description="Root task", + evaluators=(MaterializationEvaluator(name="judge"),), + ) + child = task_with_id( + uuid4(), + task_slug="child", + instance_key="a", + description="Child task", + dependency_task_slugs=("root",), + ) + return Sample.from_tasks( + name="mini:a", + sample_key="a", + environment_name="mini", + tasks=[root, child], + sample_ref={"key": "a"}, + ) + + +@pytest.fixture() +def sample_row(session: Session) -> SampleRecord: + row = SampleRecord( + benchmark_type="experiment", + instance_key="a", + status=SampleStatus.PENDING, + ) + session.add(row) + session.flush() + return row + + +def test_materialization_persists_task_json_not_environment_or_experiment( + session: Session, + selected_sample: Sample, + sample_row: SampleRecord, +) -> None: + materialize_sample(session=session, sample=selected_sample, sample_row=sample_row) + + node = session.exec( + select(SampleGraphNode).where(SampleGraphNode.sample_id == sample_row.id) + ).first() + task_event = session.exec( + select(SampleTaskEventRow).where(SampleTaskEventRow.sample_id == sample_row.id) + ).first() + worker_event = session.exec( + select(SampleWorkerEventRow).where(SampleWorkerEventRow.sample_id == sample_row.id) + ).first() + evaluator_event = session.exec( + select(SampleEvaluatorEventRow).where(SampleEvaluatorEventRow.sample_id == sample_row.id) + ).first() + + assert node is not None + assert task_event is not None + assert worker_event is not None + assert evaluator_event is not None + assert node.task_json["_type"] + assert node.task_json["worker"]["_type"] + assert node.task_json["sandbox"]["_type"] + assert node.task_json["evaluators"][0]["_type"] + assert task_event.task_snapshot_json == node.task_json + assert worker_event.worker_snapshot_json == node.task_json["worker"] + assert evaluator_event.evaluator_snapshot_json == node.task_json["evaluators"][0] + assert "environment" not in node.task_json + assert "experiment" not in node.task_json + + +def test_materialization_writes_graph_edges_from_task_dependencies( + session: Session, + selected_sample: Sample, + sample_row: SampleRecord, +) -> None: + materialize_sample(session=session, sample=selected_sample, sample_row=sample_row) + + nodes = session.exec(select(SampleGraphNode)).all() + edges = session.exec(select(SampleGraphEdge)).all() + + assert {node.task_slug for node in nodes} == {"root", "child"} + assert len(edges) == 1 From 61158ac4da87e290b039e4496bb0ebd8d3d25671 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 16:56:18 +0100 Subject: [PATCH 2/9] Support sample-only experiment submit views --- .../application/experiments/submission.py | 4 +- .../core/views/experiments/models.py | 2 +- .../core/views/experiments/service.py | 2 + .../ergon_core/core/views/samples/models.py | 4 +- .../ergon_core/core/views/samples/service.py | 47 ++++++++++++------- .../test_public_experiment_submit_smoke.py | 7 ++- .../test_definition_identity_naming.py | 20 +++++++- .../experiments/test_experiment_submit.py | 2 +- 8 files changed, 63 insertions(+), 25 deletions(-) diff --git a/ergon_core/ergon_core/core/application/experiments/submission.py b/ergon_core/ergon_core/core/application/experiments/submission.py index cb0e8c4fe..da6547d20 100644 --- a/ergon_core/ergon_core/core/application/experiments/submission.py +++ b/ergon_core/ergon_core/core/application/experiments/submission.py @@ -81,7 +81,7 @@ async def submit( samples=candidates, k=k, context=SamplingContext( - experiment_id=handle.experiment_id, + experiment_ref_id=handle.id, candidate_pool_size=pool_size, ), ) @@ -106,7 +106,7 @@ async def submit( self._session.commit() await self._start_samples(sample_ids) return ExperimentSubmitResult( - experiment_id=handle.experiment_id, + experiment_ref_id=handle.id, sampler_invocation_id=invocation.id, requested_k=k, candidate_pool_size=pool_size, diff --git a/ergon_core/ergon_core/core/views/experiments/models.py b/ergon_core/ergon_core/core/views/experiments/models.py index cf66927e3..3ec272149 100644 --- a/ergon_core/ergon_core/core/views/experiments/models.py +++ b/ergon_core/ergon_core/core/views/experiments/models.py @@ -61,7 +61,7 @@ class ExperimentRunMetricsDto(BaseModel): class ExperimentRunRowDto(BaseModel): sample_id: UUID - definition_id: UUID + definition_id: UUID | None = None benchmark_type: str instance_key: str status: str diff --git a/ergon_core/ergon_core/core/views/experiments/service.py b/ergon_core/ergon_core/core/views/experiments/service.py index e292d2daf..3c6101d06 100644 --- a/ergon_core/ergon_core/core/views/experiments/service.py +++ b/ergon_core/ergon_core/core/views/experiments/service.py @@ -76,6 +76,8 @@ def definitions_by_tag(self, tag: str) -> list[ExperimentTagDefinitionDto]: ) latest_by_definition: dict[UUID, SampleRecord] = {} for run in runs: + if run.definition_id is None: + continue latest_by_definition.setdefault(run.definition_id, run) rows: list[ExperimentTagDefinitionDto] = [] diff --git a/ergon_core/ergon_core/core/views/samples/models.py b/ergon_core/ergon_core/core/views/samples/models.py index 637b33a0d..8813d7746 100644 --- a/ergon_core/ergon_core/core/views/samples/models.py +++ b/ergon_core/ergon_core/core/views/samples/models.py @@ -195,7 +195,7 @@ class SampleSnapshotMetricsDto(CamelModel): class SampleSnapshotDto(CamelModel): id: str - definition_id: str + definition_id: str | None = None name: str status: str tasks: dict[str, SampleTaskDto] = Field(default_factory=dict) @@ -231,7 +231,7 @@ class SampleSummaryDto(BaseModel): completed_at: datetime | None = None latest_activity_at: datetime | None = None duration_seconds: float | None = None - definition_id: UUID + definition_id: UUID | None = None definition_name: str | None = None experiment: str | None = None benchmark_type: str diff --git a/ergon_core/ergon_core/core/views/samples/service.py b/ergon_core/ergon_core/core/views/samples/service.py index 03f3016a4..4d011e389 100644 --- a/ergon_core/ergon_core/core/views/samples/service.py +++ b/ergon_core/ergon_core/core/views/samples/service.py @@ -83,11 +83,12 @@ def list_samples( stmt = stmt.where(SampleRecord.experiment == experiment) stmt = stmt.offset(offset).limit(limit) rows = list(session.exec(stmt).all()) + definition_ids = [row.definition_id for row in rows if row.definition_id is not None] definition_names = { definition.id: definition.name for definition in session.exec( select(ExperimentDefinition).where( - col(ExperimentDefinition.id).in_([row.definition_id for row in rows]) + col(ExperimentDefinition.id).in_(definition_ids) ) ).all() } @@ -95,7 +96,11 @@ def list_samples( return [ _run_summary( row, - definition_name=definition_names.get(row.definition_id), + definition_name=( + definition_names.get(row.definition_id) + if row.definition_id is not None + else None + ), task_counts=task_counts.get(row.id), ) for row in rows @@ -112,11 +117,11 @@ def build_snapshot(self, sample_id: UUID) -> SampleSnapshotDto | None: if run is None: return None - definition = session.get(ExperimentDefinition, run.definition_id) - if definition is None: - return None - - def_id = run.definition_id + definition = ( + session.get(ExperimentDefinition, run.definition_id) + if run.definition_id is not None + else None + ) nodes = list( session.exec( select(SampleGraphNode).where(SampleGraphNode.sample_id == sample_id) @@ -127,12 +132,16 @@ def build_snapshot(self, sample_id: UUID) -> SampleSnapshotDto | None: select(SampleGraphEdge).where(SampleGraphEdge.sample_id == sample_id) ).all() ) - def_workers = list( - session.exec( - select(ExperimentDefinitionWorker).where( - ExperimentDefinitionWorker.experiment_definition_id == def_id - ) - ).all() + def_workers = ( + list( + session.exec( + select(ExperimentDefinitionWorker).where( + ExperimentDefinitionWorker.experiment_definition_id == run.definition_id + ) + ).all() + ) + if run.definition_id is not None + else [] ) executions = list( session.exec( @@ -198,12 +207,18 @@ def build_snapshot(self, sample_id: UUID) -> SampleSnapshotDto | None: sample_id_str = str(run.id) run_summary = run.parsed_summary() aggregated_metrics = aggregate_run_metrics(context_events, summary=run_summary) - meta = definition.parsed_metadata() - run_name = str(meta.get("name", definition.benchmark_type)) + assignment = run.parsed_assignment() + meta = definition.parsed_metadata() if definition is not None else assignment + run_name = str( + run_summary.get("name") + or assignment.get("sample_name") + or meta.get("name") + or run.benchmark_type + ) return SampleSnapshotDto( id=sample_id_str, - definition_id=str(run.definition_id), + definition_id=str(run.definition_id) if run.definition_id is not None else None, name=run_name, status=run.status, tasks=task_map, diff --git a/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py b/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py index 4fea078f5..a683b3ebe 100644 --- a/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py +++ b/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py @@ -34,6 +34,11 @@ def iter_samples(self) -> Iterator[Sample]: ) +class FakeEventBus: + async def publish(self, event) -> None: + pass + + @pytest.fixture() def session() -> Iterator[Session]: engine = create_engine("sqlite:///:memory:") @@ -50,7 +55,7 @@ async def test_public_experiment_submit_materializes_samples_and_typed_wal(sessi experiment = Experiment(name="mini-smoke", environments=[env]) result = await experiment.submit( - service=ExperimentSubmissionService.for_session(session), + service=ExperimentSubmissionService.for_session(session, event_bus=FakeEventBus()), k=1, sampler=RandomSampler(seed=0), ) 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 a034ad370..b4fd217a3 100644 --- a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py +++ b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py @@ -53,6 +53,7 @@ / "tests" / "integration" / "experiments" +<<<<<<< HEAD / "test_experiment_persistence_roundtrip.py": ( re.compile(r"experiment_id == handle\.experiment_id"), ), @@ -68,9 +69,24 @@ 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"), + re.compile(r"experiment_ref_id=uuid4"), + re.compile(r"result\.experiment_ref_id"), ), + ROOT / "ergon_core" / "ergon_core" / "core" / "application" / "experiments" / "submission.py": ( + re.compile(r"experiment_id=entry\.experiment_id"), + re.compile(r"experiment=str\(entry\.experiment_id\)"), + ), + ROOT / "ergon_core" / "ergon_core" / "core" / "persistence" / "telemetry" / "models.py": ( + re.compile(r"experiment_id"), + ), + ROOT + / "ergon_core" + / "tests" + / "unit" + / "core" + / "application" + / "experiments" + / "test_experiment_submit.py": (re.compile(r"row\.experiment_id"),), } diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py index 54f995371..6d2adc0cf 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py @@ -222,7 +222,7 @@ async def test_submit_records_selected_sample_provenance( sample_rows = session.exec(select(SampleRecord).order_by(SampleRecord.created_at)).all() assert [row.id for row in sample_rows] == list(result.sample_ids) - assert {row.experiment_id for row in sample_rows} == {result.experiment_id} + assert {row.experiment_id for row in sample_rows} == {result.experiment_ref_id} assert all(row.environment_id for row in sample_rows) assert all(row.pool_entry_id for row in sample_rows) assert all(row.sampler_invocation_id == result.sampler_invocation_id for row in sample_rows) From d4d603a90a90b832567e3984032273f2d455a0b2 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 17:03:35 +0100 Subject: [PATCH 3/9] Reduce materialization test suppressions --- .../test_public_experiment_submit_smoke.py | 15 ++++++++++----- .../test_submit_starts_materialized_sample.py | 15 ++++++++++----- .../experiments/test_experiment_submit.py | 15 ++++++++++----- .../experiments/test_sample_materialization.py | 15 ++++++++++----- 4 files changed, 40 insertions(+), 20 deletions(-) diff --git a/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py b/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py index a683b3ebe..f4d5d97cb 100644 --- a/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py +++ b/ergon_core/tests/integration/experiments/test_public_experiment_submit_smoke.py @@ -1,19 +1,24 @@ 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.experiments.models # noqa: F401 -import ergon_core.core.persistence.graph.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, RandomSampler, Sample from ergon_core.core.persistence.samples.models import SampleTaskEventRow 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.experiments.models", + "ergon_core.core.persistence.graph.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + class MaterializedEnvironment(Environment): source_mode: Literal["materialized"] = "materialized" diff --git a/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py b/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py index 49d2452f6..f2501c04e 100644 --- a/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py +++ b/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py @@ -1,14 +1,10 @@ from collections.abc import Iterable, Iterator +from importlib import import_module 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.experiments.models # noqa: F401 -import ergon_core.core.persistence.graph.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 Sample from ergon_core.api.benchmark import Task from ergon_core.api.criterion.outcome import CriterionOutcome @@ -32,6 +28,15 @@ 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.experiments.models", + "ergon_core.core.persistence.graph.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + class StartPathEvaluator(Evaluator): type_slug = "start-test-evaluator" diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py index 6d2adc0cf..067e98d3a 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py @@ -1,4 +1,5 @@ from collections.abc import Iterator, Sequence +from importlib import import_module from pathlib import Path from types import SimpleNamespace from typing import Literal @@ -7,11 +8,6 @@ import pytest from sqlmodel import Session, SQLModel, create_engine, select -import ergon_core.core.persistence.definitions.models # noqa: F401 -import ergon_core.core.persistence.experiments.models # noqa: F401 -import ergon_core.core.persistence.graph.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, RandomSampler, Sample from ergon_core.api.experiment.sampling import SamplingContext from ergon_core.core.application.events.runtime import WorkflowStartedEvent @@ -25,6 +21,15 @@ 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.experiments.models", + "ergon_core.core.persistence.graph.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + class FakeEventBus: def __init__(self) -> None: diff --git a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py index 766460fb3..40929a6d7 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py +++ b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py @@ -1,14 +1,10 @@ from collections.abc import Iterable, Iterator +from importlib import import_module 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.experiments.models # noqa: F401 -import ergon_core.core.persistence.graph.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 Sample from ergon_core.api.benchmark import Task from ergon_core.api.criterion.outcome import CriterionOutcome @@ -25,6 +21,15 @@ 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.experiments.models", + "ergon_core.core.persistence.graph.models", + "ergon_core.core.persistence.samples.models", + "ergon_core.core.persistence.telemetry.models", +): + import_module(module_name) + class MaterializationEvaluator(Evaluator): type_slug = "test-evaluator" From d6117e622e85a8ab5a00702b3d5cff7acce66354 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 17:07:46 +0100 Subject: [PATCH 4/9] Refresh sample-only dashboard contracts --- .../DashboardWorkflowStartedEvent.schema.json | 13 ++++++++++--- .../core/application/runtime/sample_lifecycle.py | 4 +++- 2 files changed, 13 insertions(+), 4 deletions(-) diff --git a/ergon-dashboard/src/generated/events/schemas/DashboardWorkflowStartedEvent.schema.json b/ergon-dashboard/src/generated/events/schemas/DashboardWorkflowStartedEvent.schema.json index 26a59671a..f2f0bd429 100644 --- a/ergon-dashboard/src/generated/events/schemas/DashboardWorkflowStartedEvent.schema.json +++ b/ergon-dashboard/src/generated/events/schemas/DashboardWorkflowStartedEvent.schema.json @@ -1092,8 +1092,16 @@ "type": "string" }, "definitionId": { - "title": "Definitionid", - "type": "string" + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Definitionid" }, "name": { "title": "Name", @@ -1272,7 +1280,6 @@ }, "required": [ "id", - "definitionId", "name", "status" ], diff --git a/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py b/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py index ebee9962c..1e0180951 100644 --- a/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py +++ b/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py @@ -278,7 +278,9 @@ async def _initialize_materialized_sample( task_slug=node.task_slug, parent_task_id=node.parent_task_id, ) - for node in sorted(nodes, key=lambda node: (node.level, node.task_slug, str(node.task_id))) + for node in sorted( + nodes, key=lambda node: (node.level, node.task_slug, str(node.task_id)) + ) ] return InitializedWorkflow( sample_id=command.sample_id, From 28d58053bedbb7861c99a618397918b53e47e079 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 17:13:16 +0100 Subject: [PATCH 5/9] Update architecture guards for sample provenance --- .../architecture/test_application_domain_boundaries.py | 7 +++++-- .../unit/architecture/test_definition_identity_naming.py | 3 +++ .../tests/unit/architecture/test_persistence_boundaries.py | 4 ++-- .../tests/unit/architecture/test_single_alembic_head.py | 1 + ergon_core/tests/unit/state/test_type_invariants.py | 2 +- 5 files changed, 12 insertions(+), 5 deletions(-) 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 1de86a5b4..eb5849952 100644 --- a/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py +++ b/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py @@ -19,15 +19,17 @@ # views, read-only task inspection, and orchestration command/result DTOs each # have distinct collaborators. PUBLIC_CROSS_DOMAIN_MODULES_BY_DOMAIN = { + "events": {"runtime"}, "runtime": { "orchestration", "resources", "sample_lifecycle", + "status", "task_execution", "task_inspection", "task_management", }, - "samples": {"events", "state"}, + "samples": {"events", "materialization", "state"}, } APPROVED_DOMAIN_FILES = { "__init__.py", @@ -52,6 +54,7 @@ "handles.py", "launch.py", "repositories.py", + "submission.py", }, "ports": {"dashboard.py", "resources.py"}, "resources": {"publishing.py"}, @@ -78,7 +81,7 @@ "workflow_errors.py", "workflow_models.py", }, - "samples": {"events.py", "state.py"}, + "samples": {"events.py", "materialization.py", "state.py"}, "testing": {"suppression_budget.py", "test_harness_service.py"}, } LAYOUT_DIR_EXCEPTIONS: dict[str, set[str]] = {} 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 b4fd217a3..9ee6a1b1f 100644 --- a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py +++ b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py @@ -87,6 +87,9 @@ / "application" / "experiments" / "test_experiment_submit.py": (re.compile(r"row\.experiment_id"),), + ROOT / "ergon_core" / "tests" / "unit" / "state" / "test_type_invariants.py": ( + re.compile(r"run\.experiment_id is None"), + ), } diff --git a/ergon_core/tests/unit/architecture/test_persistence_boundaries.py b/ergon_core/tests/unit/architecture/test_persistence_boundaries.py index 8ac936483..ef9045474 100644 --- a/ergon_core/tests/unit/architecture/test_persistence_boundaries.py +++ b/ergon_core/tests/unit/architecture/test_persistence_boundaries.py @@ -102,7 +102,7 @@ def test_run_record_uses_definition_id_as_single_runtime_definition_identity() - assert ("workflow" + "_definition_id") not in SampleRecord.model_fields -def test_run_record_does_not_expose_legacy_definition_group_identity() -> None: +def test_run_record_exposes_sample_experiment_provenance() -> None: from ergon_core.core.persistence.telemetry.models import SampleRecord - assert ("experiment" + "_id") not in SampleRecord.model_fields + assert ("experiment" + "_id") in SampleRecord.model_fields 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 39316985c..5ecfec712 100644 --- a/ergon_core/tests/unit/architecture/test_single_alembic_head.py +++ b/ergon_core/tests/unit/architecture/test_single_alembic_head.py @@ -12,6 +12,7 @@ def test_v2_has_one_initial_migration() -> None: assert migrations == [ "00000000_initial_v2.py", "00000001_add_experiment_persistence.py", + "00000002_add_sample_experiment_provenance.py", ] diff --git a/ergon_core/tests/unit/state/test_type_invariants.py b/ergon_core/tests/unit/state/test_type_invariants.py index 54b076133..d1cada9bc 100644 --- a/ergon_core/tests/unit/state/test_type_invariants.py +++ b/ergon_core/tests/unit/state/test_type_invariants.py @@ -136,7 +136,7 @@ def test_run_record_uses_definition_identity(): assert run.instance_key == "sample-1" assert run.parsed_worker_team() == {"primary": "test-worker"} assert not hasattr(run, "workflow" + "_definition_id") - assert not hasattr(run, "experiment" + "_id") + assert run.experiment_id is None assert not hasattr(run, "cohort_id") From 2fd68965fae740a82f4dd23290c12ad75f24c944 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Tue, 26 May 2026 22:47:38 +0100 Subject: [PATCH 6/9] Lean sample materialization on authored task objects --- docs/architecture/04_persistence.md | 7 +- .../application/runtime/sample_lifecycle.py | 45 +++++------ .../application/runtime/task_execution.py | 50 +++--------- .../core/application/samples/events.py | 22 +++++ .../application/samples/materialization.py | 81 ++++++------------- .../test_sample_materialization.py | 13 +++ 6 files changed, 99 insertions(+), 119 deletions(-) diff --git a/docs/architecture/04_persistence.md b/docs/architecture/04_persistence.md index 2fb058635..54f5174e9 100644 --- a/docs/architecture/04_persistence.md +++ b/docs/architecture/04_persistence.md @@ -80,7 +80,12 @@ sampler invocation, marks selected candidate rows, and creates one remain provenance rows; they are not serialized into the runtime replay contract. `workflow/started` may now carry only `sample_id` when the sample graph already exists, while the definition-backed initialization path remains as -a temporary bridge until the runtime consolidation PR removes it. +a temporary bridge. + +TODO(PR09): remove the definition-backed initialization bridge when runtime +launch no longer reads `experiment_definitions`, `experiment_definition_tasks`, +or definition worker/evaluator bindings. The remaining path should initialize +from already-materialized sample graph/WAL rows only. ### Manager-spawned subtasks (dynamic graph growth) diff --git a/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py b/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py index 1e0180951..3d6f5b0b1 100644 --- a/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py +++ b/ergon_core/ergon_core/core/application/runtime/sample_lifecycle.py @@ -9,7 +9,6 @@ ) from ergon_core.core.application.runtime import status as graph_status from ergon_core.core.persistence.graph.models import SampleGraphEdge, SampleGraphNode -from ergon_core.core.persistence.samples.models import SampleStatusEventRow from ergon_core.core.persistence.shared.db import get_session from ergon_core.core.persistence.shared.enums import ( SampleResourceKind, @@ -41,7 +40,7 @@ ) from ergon_core.core.application.runtime.graph_traversal import descendant_ids from ergon_core.core.application.runtime.models import GraphEdgeDto, GraphNodeDto, MutationMeta -from ergon_core.core.application.samples.events import SampleRuntimeEventAppender +from ergon_core.core.application.samples.events import append_sample_status_changed from ergon_core.core.application.runtime.graph_repository import RuntimeGraphRepository from ergon_core.core.application.runtime.orchestration import ( FinalizedWorkflowResult, @@ -190,14 +189,12 @@ async def _initialize_definition_backed_sample( run_record.status = SampleStatus.EXECUTING run_record.started_at = utcnow() session.add(run_record) - SampleRuntimeEventAppender(session).append_status_event( - SampleStatusEventRow( - sample_id=command.sample_id, - event_type="sample.status_changed", - status=SampleStatus.EXECUTING, - actor="system:workflow_init", - event_timestamp=run_record.started_at, - ) + append_sample_status_changed( + session, + sample_id=command.sample_id, + status=SampleStatus.EXECUTING, + actor="system:workflow_init", + event_timestamp=run_record.started_at, ) if commit: session.commit() @@ -242,14 +239,12 @@ async def _initialize_materialized_sample( run_record.status = SampleStatus.EXECUTING run_record.started_at = utcnow() session.add(run_record) - SampleRuntimeEventAppender(session).append_status_event( - SampleStatusEventRow( - sample_id=command.sample_id, - event_type="sample.status_changed", - status=SampleStatus.EXECUTING, - actor="system:workflow_init", - event_timestamp=run_record.started_at, - ) + append_sample_status_changed( + session, + sample_id=command.sample_id, + status=SampleStatus.EXECUTING, + actor="system:workflow_init", + event_timestamp=run_record.started_at, ) ready_ids = await get_initial_ready_tasks( session, @@ -326,14 +321,12 @@ def finalize(self, command: FinalizeWorkflowCommand) -> FinalizedWorkflowResult: "cost_observed": completion.cost_observed, } session.add(run_record) - SampleRuntimeEventAppender(session).append_status_event( - SampleStatusEventRow( - sample_id=command.sample_id, - event_type="sample.status_changed", - status=SampleStatus.COMPLETED, - actor="system:workflow_finalize", - event_timestamp=completion.completed_at, - ) + append_sample_status_changed( + session, + sample_id=command.sample_id, + status=SampleStatus.COMPLETED, + actor="system:workflow_finalize", + event_timestamp=completion.completed_at, ) session.commit() diff --git a/ergon_core/ergon_core/core/application/runtime/task_execution.py b/ergon_core/ergon_core/core/application/runtime/task_execution.py index 13235a0c0..d44d7cd11 100644 --- a/ergon_core/ergon_core/core/application/runtime/task_execution.py +++ b/ergon_core/ergon_core/core/application/runtime/task_execution.py @@ -1,10 +1,9 @@ """Task execution lifecycle: prepare, finalize success, finalize failure.""" import logging -from collections.abc import Mapping from uuid import UUID -from pydantic import JsonValue +from ergon_core.api.benchmark import Task from ergon_core.core.application.events.service import get_dashboard_event_publisher from ergon_core.core.persistence.definitions.models import ( ExperimentDefinition, @@ -163,8 +162,13 @@ async def _prepare_run_node( f"SampleRecord {command.sample_id} not found", ) if command.definition_id is None: - worker_type, model_target, definition_worker_id = _resolve_sample_worker_config( + ( + worker_type, + model_target, + definition_worker_id, + ) = await _resolve_sample_worker_config( node.task_json, + task_id=view.task_id, assigned_worker_slug=assigned_worker_slug, ) benchmark_type = run_record.benchmark_type @@ -320,39 +324,11 @@ async def finalize_failure(self, command: FailTaskExecutionCommand) -> None: ) -def _resolve_sample_worker_config( - task_json: Mapping[str, JsonValue], +async def _resolve_sample_worker_config( + task_json: dict, *, + task_id: UUID, assigned_worker_slug: str | None, -) -> tuple[str | None, str | None, None]: - worker_snapshot = _component_snapshot(task_json.get("worker")) - if worker_snapshot is None: - return assigned_worker_slug, None, None - worker_slug = assigned_worker_slug or _component_slug(worker_snapshot, fallback="worker") - return ( - _component_type(worker_snapshot, fallback=worker_slug), - _component_model_target(worker_snapshot), - None, - ) - - -def _component_snapshot(value: JsonValue | None) -> Mapping[str, JsonValue] | None: - return value if isinstance(value, dict) else None - - -def _component_slug(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: - value = snapshot.get("type_slug") or snapshot.get("slug") or snapshot.get("name") - if isinstance(value, str) and value: - return value - component_type = _component_type(snapshot, fallback=fallback) - return component_type.rsplit(":", 1)[-1].rsplit(".", 1)[-1] - - -def _component_type(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: - value = snapshot.get("_type") or snapshot.get("type") - return value if isinstance(value, str) and value else fallback - - -def _component_model_target(snapshot: Mapping[str, JsonValue]) -> str | None: - value = snapshot.get("model") or snapshot.get("model_target") - return value if isinstance(value, str) else None +) -> tuple[str, str, None]: + task = await Task.from_definition(task_json, task_id=task_id) + return assigned_worker_slug or task.worker.type_slug, task.worker.model, None diff --git a/ergon_core/ergon_core/core/application/samples/events.py b/ergon_core/ergon_core/core/application/samples/events.py index 8762fffc2..8f8c6ff9c 100644 --- a/ergon_core/ergon_core/core/application/samples/events.py +++ b/ergon_core/ergon_core/core/application/samples/events.py @@ -14,6 +14,7 @@ SampleWorkerEventRow, ) from ergon_core.core.shared.json_types import JsonObject +from ergon_core.core.shared.utils import utcnow from pydantic import BaseModel, Field from sqlmodel import Session, select @@ -110,6 +111,27 @@ def append_annotation_event(self, row: SampleAnnotationEventRow) -> SampleAnnota return row +def append_sample_status_changed( + session: Session, + *, + sample_id: UUID, + status: str, + actor: str, + event_timestamp: datetime | None = None, + payload: JsonObject | None = None, +) -> SampleStatusEventRow: + return SampleRuntimeEventAppender(session).append_status_event( + SampleStatusEventRow( + sample_id=sample_id, + event_type="sample.status_changed", + status=status, + actor=actor, + event_timestamp=event_timestamp or utcnow(), + payload_json=dict(payload or {}), + ) + ) + + class SampleRuntimeEventReadService: def list_events(self, session: Session, sample_id: UUID) -> list[SampleRuntimeEventView]: rows: list[SampleRuntimeEventRow] = [] diff --git a/ergon_core/ergon_core/core/application/samples/materialization.py b/ergon_core/ergon_core/core/application/samples/materialization.py index fb1aeaf36..49002410f 100644 --- a/ergon_core/ergon_core/core/application/samples/materialization.py +++ b/ergon_core/ergon_core/core/application/samples/materialization.py @@ -1,20 +1,21 @@ """Materialize authored Samples into typed runtime WAL and graph projections.""" -from collections.abc import Mapping from uuid import UUID, uuid4 -from pydantic import JsonValue from sqlmodel import Session +from ergon_core.api._serialization import component_type_path from ergon_core.api.experiment.sample import Sample from ergon_core.core.application.runtime import status as graph_status -from ergon_core.core.application.samples.events import SampleRuntimeEventAppender +from ergon_core.core.application.samples.events import ( + SampleRuntimeEventAppender, + append_sample_status_changed, +) from ergon_core.core.persistence.graph.models import SampleGraphEdge, SampleGraphNode from ergon_core.core.persistence.samples.models import ( SampleEdgeEventRow, SampleEvaluatorEventRow, SampleSandboxEventRow, - SampleStatusEventRow, SampleTaskEventRow, SampleWorkerEventRow, ) @@ -22,24 +23,6 @@ from ergon_core.core.persistence.telemetry.models import SampleRecord -def component_slug(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: - value = snapshot.get("type_slug") or snapshot.get("slug") or snapshot.get("name") - if isinstance(value, str) and value: - return value - component_type = component_type_path(snapshot, fallback=fallback) - return component_type.rsplit(":", 1)[-1].rsplit(".", 1)[-1] - - -def component_type_path(snapshot: Mapping[str, JsonValue], *, fallback: str) -> str: - value = snapshot.get("_type") or snapshot.get("type") - return value if isinstance(value, str) and value else fallback - - -def component_model_target(snapshot: Mapping[str, JsonValue]) -> str | None: - value = snapshot.get("model") or snapshot.get("model_target") - return value if isinstance(value, str) else None - - def materialize_sample( session: Session, *, @@ -47,14 +30,12 @@ def materialize_sample( sample_row: SampleRecord, ) -> None: wal = SampleRuntimeEventAppender(session) - wal.append_status_event( - SampleStatusEventRow( - sample_id=sample_row.id, - event_type="sample.status_changed", - status=SampleStatus.PENDING, - payload_json={"reason": "materialized"}, - actor="system:materialization", - ) + append_sample_status_changed( + session, + sample_id=sample_row.id, + status=SampleStatus.PENDING, + payload={"reason": "materialized"}, + actor="system:materialization", ) task_ids_by_key: dict[str, UUID] = {} @@ -62,10 +43,11 @@ def materialize_sample( task_id = uuid4() task_ids_by_key[task.task_slug] = task_id task_json = task.model_dump(mode="json") - worker_snapshot = _required_component_snapshot(task_json, "worker") - sandbox_snapshot = _required_component_snapshot(task_json, "sandbox") - worker_slug = component_slug(worker_snapshot, fallback="worker") - sandbox_slug = component_slug(sandbox_snapshot, fallback="sandbox") + worker_snapshot = task.worker.model_dump(mode="json") + sandbox_snapshot = task.sandbox.model_dump(mode="json") + worker_slug = task.worker.type_slug + sandbox_type = component_type_path(task.sandbox) + sandbox_slug = _component_display_slug(sandbox_type) session.add( SampleGraphNode( @@ -80,7 +62,7 @@ def materialize_sample( assigned_worker_slug=worker_slug, ) ) - task_event = wal.append_task_event( + wal.append_task_event( SampleTaskEventRow( sample_id=sample_row.id, task_id=task_id, @@ -103,8 +85,8 @@ def materialize_sample( task_id=task_id, event_type="worker.added", worker_slug=worker_slug, - worker_type=component_type_path(worker_snapshot, fallback=worker_slug), - model_target=component_model_target(worker_snapshot), + worker_type=task.worker.type_slug, + model_target=task.worker.model, worker_snapshot_json=worker_snapshot, payload_json={"task_id": str(task_id), "worker": worker_snapshot}, actor="system:materialization", @@ -116,27 +98,26 @@ def materialize_sample( task_id=task_id, event_type="sandbox.added", sandbox_slug=sandbox_slug, - sandbox_type=component_type_path(sandbox_snapshot, fallback=sandbox_slug), + sandbox_type=sandbox_type, sandbox_snapshot_json=sandbox_snapshot, payload_json={"task_id": str(task_id), "sandbox": sandbox_snapshot}, actor="system:materialization", ) ) - for evaluator_snapshot in _evaluator_snapshots(task_json): - evaluator_slug = component_slug(evaluator_snapshot, fallback="default") + for evaluator in task.evaluators: + evaluator_snapshot = evaluator.model_dump(mode="json") wal.append_evaluator_event( SampleEvaluatorEventRow( sample_id=sample_row.id, task_id=task_id, event_type="evaluator.added", - evaluator_slug=evaluator_slug, - evaluator_type=component_type_path(evaluator_snapshot, fallback=evaluator_slug), + evaluator_slug=evaluator.type_slug, + evaluator_type=evaluator.type_slug, evaluator_snapshot_json=evaluator_snapshot, payload_json={"task_id": str(task_id), "evaluator": evaluator_snapshot}, actor="system:materialization", ) ) - session.add(task_event) session.flush() for task in sample.tasks: @@ -175,15 +156,5 @@ def materialize_sample( session.flush() -def _required_component_snapshot(task_json: Mapping[str, JsonValue], key: str) -> dict: - value = task_json.get(key) - if not isinstance(value, dict): - raise ValueError(f"Task snapshot is missing object-bound {key}") - return value - - -def _evaluator_snapshots(task_json: Mapping[str, JsonValue]) -> list[dict]: - evaluators = task_json.get("evaluators", []) - if not isinstance(evaluators, list): - return [] - return [snapshot for snapshot in evaluators if isinstance(snapshot, dict)] +def _component_display_slug(type_path: str) -> str: + return type_path.rsplit(":", 1)[-1].rsplit(".", 1)[-1] diff --git a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py index 40929a6d7..41d4ba644 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py +++ b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py @@ -14,6 +14,7 @@ from ergon_core.core.persistence.graph.models import SampleGraphEdge, SampleGraphNode from ergon_core.core.persistence.samples.models import ( SampleEvaluatorEventRow, + SampleSandboxEventRow, SampleTaskEventRow, SampleWorkerEventRow, ) @@ -115,18 +116,30 @@ def test_materialization_persists_task_json_not_environment_or_experiment( evaluator_event = session.exec( select(SampleEvaluatorEventRow).where(SampleEvaluatorEventRow.sample_id == sample_row.id) ).first() + sandbox_event = session.exec( + select(SampleSandboxEventRow).where(SampleSandboxEventRow.sample_id == sample_row.id) + ).first() assert node is not None assert task_event is not None assert worker_event is not None assert evaluator_event is not None + assert sandbox_event is not None assert node.task_json["_type"] assert node.task_json["worker"]["_type"] assert node.task_json["sandbox"]["_type"] assert node.task_json["evaluators"][0]["_type"] assert task_event.task_snapshot_json == node.task_json + assert worker_event.worker_slug == "test-worker" + assert worker_event.worker_type == "test-worker" + assert worker_event.model_target == "test:none" assert worker_event.worker_snapshot_json == node.task_json["worker"] + assert evaluator_event.evaluator_slug == "test-evaluator" + assert evaluator_event.evaluator_type == "test-evaluator" assert evaluator_event.evaluator_snapshot_json == node.task_json["evaluators"][0] + assert sandbox_event.sandbox_slug == "TestSandbox" + assert sandbox_event.sandbox_type.endswith(":TestSandbox") + assert sandbox_event.sandbox_snapshot_json == node.task_json["sandbox"] assert "environment" not in node.task_json assert "experiment" not in node.task_json From afff3f28ad74c67017bd3732d1e2c4a9990466be Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 00:10:37 +0100 Subject: [PATCH 7/9] Update submission repository import --- .../ergon_core/core/application/experiments/submission.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ergon_core/ergon_core/core/application/experiments/submission.py b/ergon_core/ergon_core/core/application/experiments/submission.py index da6547d20..a41d44d2a 100644 --- a/ergon_core/ergon_core/core/application/experiments/submission.py +++ b/ergon_core/ergon_core/core/application/experiments/submission.py @@ -15,7 +15,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, ) From 663a7f28f5f98acb32b01be829ec521066b1a328 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 13:23:10 +0100 Subject: [PATCH 8/9] Use bound experiment persistence during submit --- .../application/experiments/repository.py | 33 ++++--------------- .../application/experiments/submission.py | 10 +++--- .../application/samples/materialization.py | 19 +++++++++-- .../test_application_domain_boundaries.py | 1 + .../test_definition_identity_naming.py | 9 ++--- .../experiments/test_experiment_submit.py | 30 +++++++++++++++-- .../test_sample_materialization.py | 6 ++-- 7 files changed, 62 insertions(+), 46 deletions(-) diff --git a/ergon_core/ergon_core/core/application/experiments/repository.py b/ergon_core/ergon_core/core/application/experiments/repository.py index 671fdc1ab..6fed5860b 100644 --- a/ergon_core/ergon_core/core/application/experiments/repository.py +++ b/ergon_core/ergon_core/core/application/experiments/repository.py @@ -61,30 +61,6 @@ def persist_experiment(self, experiment: Experiment) -> ExperimentRef: 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, @@ -172,12 +148,15 @@ def record_sampler_invocation( policy_version: int | None = None, 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, policy_version=policy_version, - sampler_config=sampler_config, + sampler_config_json=dict(sampler_config or {}), ) + session.add(row) + session.flush() + return row diff --git a/ergon_core/ergon_core/core/application/experiments/submission.py b/ergon_core/ergon_core/core/application/experiments/submission.py index a41d44d2a..5ac45139e 100644 --- a/ergon_core/ergon_core/core/application/experiments/submission.py +++ b/ergon_core/ergon_core/core/application/experiments/submission.py @@ -15,8 +15,8 @@ SampleCandidatePool, sample_from_pool_entry, ) +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.application.samples.materialization import materialize_sample @@ -66,8 +66,7 @@ async def submit( candidate_pool_size: int | None, policy_version: int | None = None, ) -> ExperimentSubmitResult: - del policy_version - handle = persist_experiment(session=self._session, experiment=experiment) + handle = await CoreExperimentPersistencePort(self._session).persist_experiment(experiment) pool_size = candidate_pool_size or k pool = SampleCandidatePool(self._session) entries = pool.fill( @@ -81,7 +80,7 @@ async def submit( samples=candidates, k=k, context=SamplingContext( - experiment_ref_id=handle.id, + experiment_id=handle.experiment_id, candidate_pool_size=pool_size, ), ) @@ -95,6 +94,7 @@ async def submit( requested_k=k, candidate_pool_size=pool_size, selected_count=len(selected_entries), + policy_version=policy_version, sampler_config=sampler.config(), ) pool.mark_selected(selected_entries, sampler_invocation_id=invocation.id) @@ -106,7 +106,7 @@ async def submit( self._session.commit() await self._start_samples(sample_ids) return ExperimentSubmitResult( - experiment_ref_id=handle.id, + experiment_id=handle.experiment_id, sampler_invocation_id=invocation.id, requested_k=k, candidate_pool_size=pool_size, diff --git a/ergon_core/ergon_core/core/application/samples/materialization.py b/ergon_core/ergon_core/core/application/samples/materialization.py index 49002410f..998255ce4 100644 --- a/ergon_core/ergon_core/core/application/samples/materialization.py +++ b/ergon_core/ergon_core/core/application/samples/materialization.py @@ -4,6 +4,8 @@ from sqlmodel import Session +from pydantic import JsonValue + from ergon_core.api._serialization import component_type_path from ergon_core.api.experiment.sample import Sample from ergon_core.core.application.runtime import status as graph_status @@ -46,7 +48,10 @@ def materialize_sample( worker_snapshot = task.worker.model_dump(mode="json") sandbox_snapshot = task.sandbox.model_dump(mode="json") worker_slug = task.worker.type_slug - sandbox_type = component_type_path(task.sandbox) + sandbox_type = _snapshot_type( + sandbox_snapshot, + fallback=component_type_path(task.sandbox), + ) sandbox_slug = _component_display_slug(sandbox_type) session.add( @@ -85,7 +90,7 @@ def materialize_sample( task_id=task_id, event_type="worker.added", worker_slug=worker_slug, - worker_type=task.worker.type_slug, + worker_type=_snapshot_type(worker_snapshot, fallback=task.worker.type_slug), model_target=task.worker.model, worker_snapshot_json=worker_snapshot, payload_json={"task_id": str(task_id), "worker": worker_snapshot}, @@ -112,7 +117,10 @@ def materialize_sample( task_id=task_id, event_type="evaluator.added", evaluator_slug=evaluator.type_slug, - evaluator_type=evaluator.type_slug, + evaluator_type=_snapshot_type( + evaluator_snapshot, + fallback=evaluator.type_slug, + ), evaluator_snapshot_json=evaluator_snapshot, payload_json={"task_id": str(task_id), "evaluator": evaluator_snapshot}, actor="system:materialization", @@ -158,3 +166,8 @@ def materialize_sample( def _component_display_slug(type_path: str) -> str: return type_path.rsplit(":", 1)[-1].rsplit(".", 1)[-1] + + +def _snapshot_type(snapshot: dict[str, JsonValue], *, fallback: str) -> str: + value = snapshot.get("_type") + return str(value) if value else fallback 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 eb5849952..c6f9b8da7 100644 --- a/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py +++ b/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py @@ -53,6 +53,7 @@ "definition_writer.py", "handles.py", "launch.py", + "persistence.py", "repositories.py", "submission.py", }, 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 9ee6a1b1f..14467245a 100644 --- a/ergon_core/tests/unit/architecture/test_definition_identity_naming.py +++ b/ergon_core/tests/unit/architecture/test_definition_identity_naming.py @@ -40,7 +40,7 @@ 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_id=handle\.experiment_id"), re.compile(r"experiment_ref\.experiment_id"), re.compile(r"experiment_id: UUID"), re.compile(r"experiment_id == experiment_id"), @@ -53,7 +53,6 @@ / "tests" / "integration" / "experiments" -<<<<<<< HEAD / "test_experiment_persistence_roundtrip.py": ( re.compile(r"experiment_id == handle\.experiment_id"), ), @@ -67,12 +66,10 @@ / "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_ref_id=uuid4"), - re.compile(r"result\.experiment_ref_id"), + re.compile(r"ref\.experiment_id"), ), ROOT / "ergon_core" / "ergon_core" / "core" / "application" / "experiments" / "submission.py": ( + re.compile(r"experiment_id=handle\.experiment_id"), re.compile(r"experiment_id=entry\.experiment_id"), re.compile(r"experiment=str\(entry\.experiment_id\)"), ), diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py index 067e98d3a..953ab87b8 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py @@ -15,7 +15,10 @@ ExperimentSubmissionService, InngestWorkflowEventBus, ) -from ergon_core.core.persistence.experiments.models import ExperimentSamplePoolEntryRow +from ergon_core.core.persistence.experiments.models import ( + ExperimentSamplePoolEntryRow, + ExperimentSamplerInvocationRow, +) from ergon_core.core.persistence.graph.models import SampleGraphNode from ergon_core.core.persistence.samples.models import SampleEdgeEventRow, SampleTaskEventRow from ergon_core.core.persistence.telemetry.models import SampleRecord @@ -226,13 +229,36 @@ async def test_submit_records_selected_sample_provenance( result = await two_env_experiment.submit(service=service, k=2, sampler=SequentialSampler()) sample_rows = session.exec(select(SampleRecord).order_by(SampleRecord.created_at)).all() + pool_rows = session.exec(select(ExperimentSamplePoolEntryRow)).all() assert [row.id for row in sample_rows] == list(result.sample_ids) - assert {row.experiment_id for row in sample_rows} == {result.experiment_ref_id} + assert {row.experiment_id for row in sample_rows} == {result.experiment_id} assert all(row.environment_id for row in sample_rows) assert all(row.pool_entry_id for row in sample_rows) assert all(row.sampler_invocation_id == result.sampler_invocation_id for row in sample_rows) assert all(row.sample_key for row in sample_rows) assert all(isinstance(row.sample_ref_json, dict) for row in sample_rows) + assert {(row.sample_key, row.sample_ref_json["key"]) for row in sample_rows} == { + (row.sample_key, row.sample_ref_json["key"]) for row in pool_rows if row.selected + } + + +@pytest.mark.asyncio +async def test_submit_records_sampler_policy_version( + session: Session, + experiment: Experiment, +) -> None: + service = ExperimentSubmissionService(session=session, event_bus=FakeEventBus()) + + result = await experiment.submit( + service=service, + k=1, + sampler=SequentialSampler(), + policy_version=7, + ) + + invocation = session.get(ExperimentSamplerInvocationRow, result.sampler_invocation_id) + assert invocation is not None + assert invocation.policy_version == 7 @pytest.mark.asyncio diff --git a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py index 41d4ba644..debff0268 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py +++ b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py @@ -131,14 +131,14 @@ def test_materialization_persists_task_json_not_environment_or_experiment( assert node.task_json["evaluators"][0]["_type"] assert task_event.task_snapshot_json == node.task_json assert worker_event.worker_slug == "test-worker" - assert worker_event.worker_type == "test-worker" + assert worker_event.worker_type == node.task_json["worker"]["_type"] assert worker_event.model_target == "test:none" assert worker_event.worker_snapshot_json == node.task_json["worker"] assert evaluator_event.evaluator_slug == "test-evaluator" - assert evaluator_event.evaluator_type == "test-evaluator" + assert evaluator_event.evaluator_type == node.task_json["evaluators"][0]["_type"] assert evaluator_event.evaluator_snapshot_json == node.task_json["evaluators"][0] assert sandbox_event.sandbox_slug == "TestSandbox" - assert sandbox_event.sandbox_type.endswith(":TestSandbox") + assert sandbox_event.sandbox_type == node.task_json["sandbox"]["_type"] assert sandbox_event.sandbox_snapshot_json == node.task_json["sandbox"] assert "environment" not in node.task_json assert "experiment" not in node.task_json From faa0aaee4f0c3c020926c428a4e3636c82d51d69 Mon Sep 17 00:00:00 2001 From: Charlie Masters <69640669+cm2435@users.noreply.github.com> Date: Wed, 27 May 2026 15:39:53 +0100 Subject: [PATCH 9/9] Reuse persisted experiments on repeated submit --- .../application/experiments/persistence.py | 7 ++++- .../experiments/test_experiment_submit.py | 28 +++++++++++++++++++ 2 files changed, 34 insertions(+), 1 deletion(-) diff --git a/ergon_core/ergon_core/core/application/experiments/persistence.py b/ergon_core/ergon_core/core/application/experiments/persistence.py index 0f0aff231..c8ebf269e 100644 --- a/ergon_core/ergon_core/core/application/experiments/persistence.py +++ b/ergon_core/ergon_core/core/application/experiments/persistence.py @@ -17,4 +17,9 @@ 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) + ref = experiment.persisted_ref() + if ref is not None: + return ref + ref = persist_row_graph(session=self._session, experiment=experiment) + experiment.mark_persisted(ref) + return ref diff --git a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py index 953ab87b8..466442b82 100644 --- a/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py @@ -196,6 +196,34 @@ async def test_submit_retains_unselected_candidate_pool_entries( assert sum(row.selected for row in rows) == 2 +@pytest.mark.asyncio +async def test_submit_reuses_persisted_experiment_and_retained_candidates( + session: Session, + streaming_experiment: Experiment, +) -> None: + service = ExperimentSubmissionService(session=session, event_bus=FakeEventBus()) + + first = await streaming_experiment.submit( + service=service, + k=2, + sampler=SequentialSampler(), + candidate_pool_size=8, + ) + second = await streaming_experiment.submit( + service=service, + k=2, + sampler=SequentialSampler(), + candidate_pool_size=8, + ) + + rows = session.exec( + select(ExperimentSamplePoolEntryRow).order_by(ExperimentSamplePoolEntryRow.created_at) + ).all() + assert second.experiment_id == first.experiment_id + assert [row.sample_key for row in rows] == [str(index) for index in range(10)] + assert sum(row.selected for row in rows) == 4 + + @pytest.mark.asyncio async def test_submit_caps_random_sampler_selection_to_requested_k( session: Session,