diff --git a/docs/architecture/04_persistence.md b/docs/architecture/04_persistence.md index d9e9bca7c..54f5174e9 100644 --- a/docs/architecture/04_persistence.md +++ b/docs/architecture/04_persistence.md @@ -70,6 +70,23 @@ 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. + +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) The graph is append-only at the row level: new nodes and edges can enter 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/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/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/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 new file mode 100644 index 000000000..5ac45139e --- /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.persistence import CoreExperimentPersistencePort +from ergon_core.core.application.experiments.repository import ( + 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: + handle = await CoreExperimentPersistencePort(self._session).persist_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), + policy_version=policy_version, + 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..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, @@ -105,78 +104,188 @@ 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) - 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() - ready_ids = await get_initial_ready_tasks( + 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) + append_sample_status_changed( session, - command.sample_id, - command.definition_id, - graph_repo=self._graph_repo, - graph_lookup=graph_lookup, + sample_id=command.sample_id, + 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.""" @@ -212,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 56034dccb..d44d7cd11 100644 --- a/ergon_core/ergon_core/core/application/runtime/task_execution.py +++ b/ergon_core/ergon_core/core/application/runtime/task_execution.py @@ -3,6 +3,7 @@ import logging from uuid import UUID +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, @@ -155,18 +156,34 @@ 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, + ) = 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 + 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 +200,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 +322,13 @@ async def finalize_failure(self, command: FailTaskExecutionCommand) -> None: new_status=graph_status.FAILED, old_status=graph_status.RUNNING, ) + + +async def _resolve_sample_worker_config( + task_json: dict, + *, + task_id: UUID, + assigned_worker_slug: str | 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/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/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 new file mode 100644 index 000000000..998255ce4 --- /dev/null +++ b/ergon_core/ergon_core/core/application/samples/materialization.py @@ -0,0 +1,173 @@ +"""Materialize authored Samples into typed runtime WAL and graph projections.""" + +from uuid import UUID, uuid4 + +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 +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, + SampleTaskEventRow, + SampleWorkerEventRow, +) +from ergon_core.core.persistence.shared.enums import SampleStatus +from ergon_core.core.persistence.telemetry.models import SampleRecord + + +def materialize_sample( + session: Session, + *, + sample: Sample, + sample_row: SampleRecord, +) -> None: + wal = SampleRuntimeEventAppender(session) + 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] = {} + 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 = task.worker.model_dump(mode="json") + sandbox_snapshot = task.sandbox.model_dump(mode="json") + worker_slug = task.worker.type_slug + sandbox_type = _snapshot_type( + sandbox_snapshot, + fallback=component_type_path(task.sandbox), + ) + sandbox_slug = _component_display_slug(sandbox_type) + + 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, + ) + ) + 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=_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}, + 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=sandbox_type, + sandbox_snapshot_json=sandbox_snapshot, + payload_json={"task_id": str(task_id), "sandbox": sandbox_snapshot}, + actor="system:materialization", + ) + ) + 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.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", + ) + ) + + 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 _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/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/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/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..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,18 +1,23 @@ from collections.abc import Iterator +from importlib import import_module from typing import Literal from uuid import uuid4 import pytest -from sqlmodel import select +from sqlmodel import Session, SQLModel, create_engine, select 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", -) +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): @@ -34,6 +39,19 @@ 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:") + 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 @@ -42,7 +60,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/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..f2501c04e --- /dev/null +++ b/ergon_core/tests/integration/experiments/test_submit_starts_materialized_sample.py @@ -0,0 +1,255 @@ +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 + +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 + +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" + + 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/architecture/test_application_domain_boundaries.py b/ergon_core/tests/unit/architecture/test_application_domain_boundaries.py index 1de86a5b4..c6f9b8da7 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", @@ -51,7 +53,9 @@ "definition_writer.py", "handles.py", "launch.py", + "persistence.py", "repositories.py", + "submission.py", }, "ports": {"dashboard.py", "resources.py"}, "resources": {"publishing.py"}, @@ -78,7 +82,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 a034ad370..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"), @@ -66,10 +66,26 @@ / "test_experiment_persistence.py": ( re.compile(r"experiment_id == handle\.experiment_id"), re.compile(r"handle\.experiment_id"), + re.compile(r"ref\.experiment_id"), ), - ROOT / "ergon_core" / "tests" / "unit" / "api" / "test_sampler_contract.py": ( - re.compile(r"experiment_id=uuid4"), - re.compile(r"result\.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\)"), + ), + 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"),), + 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/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..466442b82 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_experiment_submit.py @@ -0,0 +1,369 @@ +from collections.abc import Iterator, Sequence +from importlib import import_module +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 + +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, + 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 +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: + 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_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, + 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() + 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_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 +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..debff0268 --- /dev/null +++ b/ergon_core/tests/unit/core/application/experiments/test_sample_materialization.py @@ -0,0 +1,158 @@ +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 + +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, + SampleSandboxEventRow, + 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 + +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" + + 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() + 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 == 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 == 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 == 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 + + +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 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")