tunarag-python 0.2.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- tunarag/__init__.py +166 -0
- tunarag/cache.py +326 -0
- tunarag/config.py +257 -0
- tunarag/contracts.py +88 -0
- tunarag/dataset.py +481 -0
- tunarag/domain.py +83 -0
- tunarag/engine.py +979 -0
- tunarag/errors.py +179 -0
- tunarag/evaluators.py +186 -0
- tunarag/integrations/__init__.py +25 -0
- tunarag/integrations/mlflow.py +155 -0
- tunarag/integrations/runnables.py +270 -0
- tunarag/objective.py +75 -0
- tunarag/py.typed +1 -0
- tunarag/result.py +309 -0
- tunarag/retry.py +70 -0
- tunarag/search.py +286 -0
- tunarag/serialization.py +78 -0
- tunarag/stopping.py +193 -0
- tunarag/store.py +941 -0
- tunarag/synthetic.py +257 -0
- tunarag_python-0.2.1.dist-info/METADATA +1164 -0
- tunarag_python-0.2.1.dist-info/RECORD +24 -0
- tunarag_python-0.2.1.dist-info/WHEEL +4 -0
tunarag/store.py
ADDED
|
@@ -0,0 +1,941 @@
|
|
|
1
|
+
"""SQLite-backed study, trial, attempt, observation, and event persistence."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import math
|
|
5
|
+
import sqlite3
|
|
6
|
+
import time
|
|
7
|
+
import uuid
|
|
8
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
from enum import Enum
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
from typing import Any, cast
|
|
13
|
+
|
|
14
|
+
from .domain import Candidate, MetricValue, UsageRecord, ValueStatus
|
|
15
|
+
from .errors import StoreError, redact
|
|
16
|
+
from .serialization import canonical_json, content_hash, storage_json
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class StudyStatus(str, Enum):
|
|
20
|
+
"""Durable lifecycle states for a study."""
|
|
21
|
+
|
|
22
|
+
CREATED = "created"
|
|
23
|
+
RUNNING = "running"
|
|
24
|
+
COMPLETED = "completed"
|
|
25
|
+
FAILED = "failed"
|
|
26
|
+
CANCELLED = "cancelled"
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class TrialStatus(str, Enum):
|
|
30
|
+
"""Durable lifecycle states for a trial."""
|
|
31
|
+
|
|
32
|
+
PENDING = "pending"
|
|
33
|
+
RUNNING = "running"
|
|
34
|
+
SUCCEEDED = "succeeded"
|
|
35
|
+
FAILED = "failed"
|
|
36
|
+
ABANDONED = "abandoned"
|
|
37
|
+
CANCELLED = "cancelled"
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class AttemptStatus(str, Enum):
|
|
41
|
+
"""Durable lifecycle states for a trial attempt."""
|
|
42
|
+
|
|
43
|
+
RUNNING = "running"
|
|
44
|
+
SUCCEEDED = "succeeded"
|
|
45
|
+
FAILED = "failed"
|
|
46
|
+
ABANDONED = "abandoned"
|
|
47
|
+
CANCELLED = "cancelled"
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
_STUDY_TRANSITIONS = {
|
|
51
|
+
StudyStatus.CREATED: {StudyStatus.RUNNING, StudyStatus.FAILED, StudyStatus.CANCELLED},
|
|
52
|
+
StudyStatus.RUNNING: {StudyStatus.COMPLETED, StudyStatus.FAILED, StudyStatus.CANCELLED},
|
|
53
|
+
StudyStatus.COMPLETED: set(),
|
|
54
|
+
StudyStatus.FAILED: set(),
|
|
55
|
+
StudyStatus.CANCELLED: set(),
|
|
56
|
+
}
|
|
57
|
+
_TRIAL_TRANSITIONS = {
|
|
58
|
+
TrialStatus.PENDING: {
|
|
59
|
+
TrialStatus.RUNNING,
|
|
60
|
+
TrialStatus.ABANDONED,
|
|
61
|
+
TrialStatus.CANCELLED,
|
|
62
|
+
},
|
|
63
|
+
TrialStatus.RUNNING: {
|
|
64
|
+
TrialStatus.SUCCEEDED,
|
|
65
|
+
TrialStatus.FAILED,
|
|
66
|
+
TrialStatus.ABANDONED,
|
|
67
|
+
TrialStatus.CANCELLED,
|
|
68
|
+
},
|
|
69
|
+
TrialStatus.SUCCEEDED: set(),
|
|
70
|
+
TrialStatus.FAILED: set(),
|
|
71
|
+
TrialStatus.ABANDONED: set(),
|
|
72
|
+
TrialStatus.CANCELLED: set(),
|
|
73
|
+
}
|
|
74
|
+
_TERMINAL_TRIAL_STATUSES = {
|
|
75
|
+
TrialStatus.SUCCEEDED,
|
|
76
|
+
TrialStatus.FAILED,
|
|
77
|
+
TrialStatus.ABANDONED,
|
|
78
|
+
TrialStatus.CANCELLED,
|
|
79
|
+
}
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
@dataclass(frozen=True, slots=True)
|
|
83
|
+
class StudyRecord:
|
|
84
|
+
"""Persisted study identity and lifecycle state."""
|
|
85
|
+
|
|
86
|
+
id: str
|
|
87
|
+
name: str
|
|
88
|
+
config_hash: str
|
|
89
|
+
status: StudyStatus
|
|
90
|
+
created_at: str
|
|
91
|
+
updated_at: str
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
@dataclass(frozen=True, slots=True)
|
|
95
|
+
class TrialRecord:
|
|
96
|
+
"""Persisted trial identity and lifecycle state."""
|
|
97
|
+
|
|
98
|
+
id: str
|
|
99
|
+
study_id: str
|
|
100
|
+
sequence: int
|
|
101
|
+
candidate_hash: str
|
|
102
|
+
candidate_json: str
|
|
103
|
+
status: TrialStatus
|
|
104
|
+
error: str | None
|
|
105
|
+
created_at: str
|
|
106
|
+
updated_at: str
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@dataclass(frozen=True, slots=True)
|
|
110
|
+
class AttemptRecord:
|
|
111
|
+
"""Persisted execution attempt with a recovery lease."""
|
|
112
|
+
|
|
113
|
+
id: str
|
|
114
|
+
trial_id: str
|
|
115
|
+
number: int
|
|
116
|
+
status: AttemptStatus
|
|
117
|
+
started_at: float
|
|
118
|
+
lease_expires_at: float
|
|
119
|
+
finished_at: float | None
|
|
120
|
+
error: str | None
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
@dataclass(frozen=True, slots=True)
|
|
124
|
+
class PersistedUsageRecord:
|
|
125
|
+
"""Append-only usage observation associated with a trial and optional attempt."""
|
|
126
|
+
|
|
127
|
+
id: str
|
|
128
|
+
trial_id: str
|
|
129
|
+
attempt_id: str | None
|
|
130
|
+
usage: UsageRecord
|
|
131
|
+
created_at: str
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
@dataclass(frozen=True, slots=True)
|
|
135
|
+
class StudyEvent:
|
|
136
|
+
"""Ordered durable event for replay and observability."""
|
|
137
|
+
|
|
138
|
+
study_id: str
|
|
139
|
+
sequence: int
|
|
140
|
+
type: str
|
|
141
|
+
payload: Mapping[str, Any]
|
|
142
|
+
created_at: str
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
class SQLiteStore:
|
|
146
|
+
"""Transactional SQLite authority for optimization lifecycle state."""
|
|
147
|
+
|
|
148
|
+
schema_version = 2
|
|
149
|
+
|
|
150
|
+
def __init__(
|
|
151
|
+
self,
|
|
152
|
+
path: str | Path = ":memory:",
|
|
153
|
+
*,
|
|
154
|
+
clock: Callable[[], float] = time.time,
|
|
155
|
+
) -> None:
|
|
156
|
+
self._connection = sqlite3.connect(str(path), timeout=30.0)
|
|
157
|
+
self._connection.row_factory = sqlite3.Row
|
|
158
|
+
self._clock = clock
|
|
159
|
+
self.initialize()
|
|
160
|
+
|
|
161
|
+
def initialize(self) -> None:
|
|
162
|
+
"""Create or migrate the experiment schema to the supported version."""
|
|
163
|
+
|
|
164
|
+
self._connection.execute("PRAGMA foreign_keys = ON")
|
|
165
|
+
self._connection.execute(
|
|
166
|
+
"CREATE TABLE IF NOT EXISTS schema_version (version INTEGER PRIMARY KEY)"
|
|
167
|
+
)
|
|
168
|
+
self._connection.commit()
|
|
169
|
+
versions = self._schema_versions()
|
|
170
|
+
if not versions:
|
|
171
|
+
self._create_schema_version_1()
|
|
172
|
+
versions = [1]
|
|
173
|
+
if versions == [1]:
|
|
174
|
+
self._migrate_version_1_to_2()
|
|
175
|
+
versions = [1, 2]
|
|
176
|
+
if versions != [1, 2]:
|
|
177
|
+
self._connection.close()
|
|
178
|
+
raise StoreError(
|
|
179
|
+
"unsupported experiment schema versions",
|
|
180
|
+
stage="store.initialize",
|
|
181
|
+
details={"versions": versions, "supported": self.schema_version},
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
def create_study(self, name: str, configuration: object) -> StudyRecord:
|
|
185
|
+
"""Create a study in the created state."""
|
|
186
|
+
|
|
187
|
+
if not name.strip():
|
|
188
|
+
raise ValueError("study name must not be empty")
|
|
189
|
+
study_id = str(uuid.uuid4())
|
|
190
|
+
config_hash = content_hash(configuration, namespace="study:v1")
|
|
191
|
+
with self._connection:
|
|
192
|
+
self._connection.execute(
|
|
193
|
+
"""INSERT INTO studies(id, name, config_hash, status, updated_at)
|
|
194
|
+
VALUES (?, ?, ?, 'created', CURRENT_TIMESTAMP)""",
|
|
195
|
+
(study_id, name, config_hash),
|
|
196
|
+
)
|
|
197
|
+
self._append_event_locked(study_id, "study.created", {"name": name})
|
|
198
|
+
return self.get_study(study_id)
|
|
199
|
+
|
|
200
|
+
def get_study(self, study_id: str) -> StudyRecord:
|
|
201
|
+
"""Load one study or raise `KeyError`."""
|
|
202
|
+
|
|
203
|
+
row = self._connection.execute(
|
|
204
|
+
"""SELECT id, name, config_hash, status, created_at, updated_at
|
|
205
|
+
FROM studies WHERE id = ?""",
|
|
206
|
+
(study_id,),
|
|
207
|
+
).fetchone()
|
|
208
|
+
if row is None:
|
|
209
|
+
raise KeyError(f"unknown study: {study_id}")
|
|
210
|
+
return _study_from_row(cast(sqlite3.Row, row))
|
|
211
|
+
|
|
212
|
+
def transition_study(self, study_id: str, status: StudyStatus) -> StudyRecord:
|
|
213
|
+
"""Apply a guarded study lifecycle transition."""
|
|
214
|
+
|
|
215
|
+
self._begin_immediate()
|
|
216
|
+
try:
|
|
217
|
+
current = self._get_study_status_locked(study_id)
|
|
218
|
+
if status not in _STUDY_TRANSITIONS[current]:
|
|
219
|
+
raise ValueError(f"invalid study transition: {current.value} -> {status.value}")
|
|
220
|
+
if status is StudyStatus.COMPLETED:
|
|
221
|
+
active = self._connection.execute(
|
|
222
|
+
"""SELECT 1 FROM trials
|
|
223
|
+
WHERE study_id = ? AND status IN ('pending', 'running') LIMIT 1""",
|
|
224
|
+
(study_id,),
|
|
225
|
+
).fetchone()
|
|
226
|
+
if active is not None:
|
|
227
|
+
raise ValueError("study cannot complete while trials are active")
|
|
228
|
+
self._connection.execute(
|
|
229
|
+
"""UPDATE studies SET status = ?, updated_at = CURRENT_TIMESTAMP
|
|
230
|
+
WHERE id = ?""",
|
|
231
|
+
(status.value, study_id),
|
|
232
|
+
)
|
|
233
|
+
self._append_event_locked(
|
|
234
|
+
study_id, f"study.{status.value}", {"previous_status": current.value}
|
|
235
|
+
)
|
|
236
|
+
self._connection.commit()
|
|
237
|
+
except BaseException:
|
|
238
|
+
self._connection.rollback()
|
|
239
|
+
raise
|
|
240
|
+
return self.get_study(study_id)
|
|
241
|
+
|
|
242
|
+
def create_trial(self, study_id: str, sequence: int, candidate: Candidate) -> TrialRecord:
|
|
243
|
+
"""Create a trial with an explicit nonnegative sequence."""
|
|
244
|
+
|
|
245
|
+
if sequence < 0:
|
|
246
|
+
raise ValueError("trial sequence must not be negative")
|
|
247
|
+
self._begin_immediate()
|
|
248
|
+
try:
|
|
249
|
+
self._require_running_study_locked(study_id)
|
|
250
|
+
record = self._insert_trial_locked(study_id, sequence, candidate)
|
|
251
|
+
self._connection.commit()
|
|
252
|
+
except BaseException:
|
|
253
|
+
self._connection.rollback()
|
|
254
|
+
raise
|
|
255
|
+
return record
|
|
256
|
+
|
|
257
|
+
def reserve_trial(self, study_id: str, candidate: Candidate) -> TrialRecord:
|
|
258
|
+
"""Atomically allocate the next sequence and reserve a unique candidate."""
|
|
259
|
+
|
|
260
|
+
self._begin_immediate()
|
|
261
|
+
try:
|
|
262
|
+
self._require_running_study_locked(study_id)
|
|
263
|
+
row = self._connection.execute(
|
|
264
|
+
"SELECT COALESCE(MAX(sequence), -1) + 1 FROM trials WHERE study_id = ?",
|
|
265
|
+
(study_id,),
|
|
266
|
+
).fetchone()
|
|
267
|
+
sequence = int(cast(sqlite3.Row, row)[0])
|
|
268
|
+
record = self._insert_trial_locked(study_id, sequence, candidate)
|
|
269
|
+
self._connection.commit()
|
|
270
|
+
except BaseException:
|
|
271
|
+
self._connection.rollback()
|
|
272
|
+
raise
|
|
273
|
+
return record
|
|
274
|
+
|
|
275
|
+
def get_trial(self, trial_id: str) -> TrialRecord:
|
|
276
|
+
"""Load one trial or raise `KeyError`."""
|
|
277
|
+
|
|
278
|
+
row = self._connection.execute(
|
|
279
|
+
"""SELECT id, study_id, sequence, candidate_hash, candidate_json, status,
|
|
280
|
+
error, created_at, updated_at FROM trials WHERE id = ?""",
|
|
281
|
+
(trial_id,),
|
|
282
|
+
).fetchone()
|
|
283
|
+
if row is None:
|
|
284
|
+
raise KeyError(f"unknown trial: {trial_id}")
|
|
285
|
+
return _trial_from_row(cast(sqlite3.Row, row))
|
|
286
|
+
|
|
287
|
+
def list_trials(self, study_id: str) -> list[TrialRecord]:
|
|
288
|
+
"""Return trials in durable sequence order."""
|
|
289
|
+
|
|
290
|
+
rows = self._connection.execute(
|
|
291
|
+
"""SELECT id, study_id, sequence, candidate_hash, candidate_json, status,
|
|
292
|
+
error, created_at, updated_at FROM trials
|
|
293
|
+
WHERE study_id = ? ORDER BY sequence""",
|
|
294
|
+
(study_id,),
|
|
295
|
+
).fetchall()
|
|
296
|
+
return [_trial_from_row(row) for row in rows]
|
|
297
|
+
|
|
298
|
+
def transition_trial(
|
|
299
|
+
self, trial_id: str, status: TrialStatus, *, error: str | None = None
|
|
300
|
+
) -> TrialRecord:
|
|
301
|
+
"""Apply a guarded trial lifecycle transition."""
|
|
302
|
+
|
|
303
|
+
self._begin_immediate()
|
|
304
|
+
try:
|
|
305
|
+
row = self._get_trial_row_locked(trial_id)
|
|
306
|
+
current = TrialStatus(row["status"])
|
|
307
|
+
if status not in _TRIAL_TRANSITIONS[current]:
|
|
308
|
+
raise ValueError(f"invalid trial transition: {current.value} -> {status.value}")
|
|
309
|
+
if status in _TERMINAL_TRIAL_STATUSES:
|
|
310
|
+
active = self._connection.execute(
|
|
311
|
+
"SELECT 1 FROM trial_attempts WHERE trial_id = ? AND status = 'running'",
|
|
312
|
+
(trial_id,),
|
|
313
|
+
).fetchone()
|
|
314
|
+
if active is not None:
|
|
315
|
+
raise ValueError("active attempt must be completed through commit_trial_result")
|
|
316
|
+
self._connection.execute(
|
|
317
|
+
"""UPDATE trials SET status = ?, error = ?, updated_at = CURRENT_TIMESTAMP
|
|
318
|
+
WHERE id = ?""",
|
|
319
|
+
(status.value, error, trial_id),
|
|
320
|
+
)
|
|
321
|
+
self._append_event_locked(
|
|
322
|
+
row["study_id"],
|
|
323
|
+
f"trial.{status.value}",
|
|
324
|
+
{"trial_id": trial_id, "previous_status": current.value},
|
|
325
|
+
)
|
|
326
|
+
self._connection.commit()
|
|
327
|
+
except BaseException:
|
|
328
|
+
self._connection.rollback()
|
|
329
|
+
raise
|
|
330
|
+
return self.get_trial(trial_id)
|
|
331
|
+
|
|
332
|
+
def start_attempt(
|
|
333
|
+
self,
|
|
334
|
+
trial_id: str,
|
|
335
|
+
*,
|
|
336
|
+
lease_seconds: float,
|
|
337
|
+
now: float | None = None,
|
|
338
|
+
) -> AttemptRecord:
|
|
339
|
+
"""Start the next attempt and acquire its recovery lease."""
|
|
340
|
+
|
|
341
|
+
if not math.isfinite(lease_seconds) or lease_seconds <= 0:
|
|
342
|
+
raise ValueError("attempt lease must be a positive finite number")
|
|
343
|
+
started_at = self._timestamp(now)
|
|
344
|
+
self._begin_immediate()
|
|
345
|
+
try:
|
|
346
|
+
trial = self._get_trial_row_locked(trial_id)
|
|
347
|
+
trial_status = TrialStatus(trial["status"])
|
|
348
|
+
if trial_status not in {TrialStatus.PENDING, TrialStatus.RUNNING}:
|
|
349
|
+
raise ValueError(f"cannot start attempt for {trial_status.value} trial")
|
|
350
|
+
active = self._connection.execute(
|
|
351
|
+
"SELECT 1 FROM trial_attempts WHERE trial_id = ? AND status = 'running'",
|
|
352
|
+
(trial_id,),
|
|
353
|
+
).fetchone()
|
|
354
|
+
if active is not None:
|
|
355
|
+
raise ValueError("trial already has a running attempt")
|
|
356
|
+
row = self._connection.execute(
|
|
357
|
+
"""SELECT COALESCE(MAX(attempt_number), -1) + 1
|
|
358
|
+
FROM trial_attempts WHERE trial_id = ?""",
|
|
359
|
+
(trial_id,),
|
|
360
|
+
).fetchone()
|
|
361
|
+
number = int(cast(sqlite3.Row, row)[0])
|
|
362
|
+
attempt_id = str(uuid.uuid4())
|
|
363
|
+
lease_expires_at = started_at + lease_seconds
|
|
364
|
+
self._connection.execute(
|
|
365
|
+
"""INSERT INTO trial_attempts(
|
|
366
|
+
id, trial_id, attempt_number, status, started_at, lease_expires_at
|
|
367
|
+
) VALUES (?, ?, ?, 'running', ?, ?)""",
|
|
368
|
+
(attempt_id, trial_id, number, started_at, lease_expires_at),
|
|
369
|
+
)
|
|
370
|
+
if trial_status is TrialStatus.PENDING:
|
|
371
|
+
self._connection.execute(
|
|
372
|
+
"""UPDATE trials SET status = 'running', updated_at = CURRENT_TIMESTAMP
|
|
373
|
+
WHERE id = ?""",
|
|
374
|
+
(trial_id,),
|
|
375
|
+
)
|
|
376
|
+
self._append_event_locked(
|
|
377
|
+
trial["study_id"],
|
|
378
|
+
"trial.attempt.started",
|
|
379
|
+
{"trial_id": trial_id, "attempt_id": attempt_id, "attempt": number},
|
|
380
|
+
)
|
|
381
|
+
self._connection.commit()
|
|
382
|
+
except BaseException:
|
|
383
|
+
self._connection.rollback()
|
|
384
|
+
raise
|
|
385
|
+
return self.get_attempt(attempt_id)
|
|
386
|
+
|
|
387
|
+
def get_attempt(self, attempt_id: str) -> AttemptRecord:
|
|
388
|
+
"""Load one attempt or raise `KeyError`."""
|
|
389
|
+
|
|
390
|
+
row = self._connection.execute(
|
|
391
|
+
"""SELECT id, trial_id, attempt_number, status, started_at,
|
|
392
|
+
lease_expires_at, finished_at, error FROM trial_attempts WHERE id = ?""",
|
|
393
|
+
(attempt_id,),
|
|
394
|
+
).fetchone()
|
|
395
|
+
if row is None:
|
|
396
|
+
raise KeyError(f"unknown attempt: {attempt_id}")
|
|
397
|
+
return _attempt_from_row(cast(sqlite3.Row, row))
|
|
398
|
+
|
|
399
|
+
def list_attempts(self, trial_id: str) -> list[AttemptRecord]:
|
|
400
|
+
"""Return attempts in monotonic attempt-number order."""
|
|
401
|
+
|
|
402
|
+
rows = self._connection.execute(
|
|
403
|
+
"""SELECT id, trial_id, attempt_number, status, started_at,
|
|
404
|
+
lease_expires_at, finished_at, error FROM trial_attempts
|
|
405
|
+
WHERE trial_id = ? ORDER BY attempt_number""",
|
|
406
|
+
(trial_id,),
|
|
407
|
+
).fetchall()
|
|
408
|
+
return [_attempt_from_row(row) for row in rows]
|
|
409
|
+
|
|
410
|
+
def finish_attempt(
|
|
411
|
+
self,
|
|
412
|
+
attempt_id: str,
|
|
413
|
+
status: AttemptStatus,
|
|
414
|
+
*,
|
|
415
|
+
usage: Sequence[UsageRecord] = (),
|
|
416
|
+
error: str | None = None,
|
|
417
|
+
now: float | None = None,
|
|
418
|
+
) -> AttemptRecord:
|
|
419
|
+
"""Finish a running attempt while leaving retryable trial state intact."""
|
|
420
|
+
|
|
421
|
+
if status is AttemptStatus.RUNNING:
|
|
422
|
+
raise ValueError("attempt terminal status must not be running")
|
|
423
|
+
finished_at = self._timestamp(now)
|
|
424
|
+
self._begin_immediate()
|
|
425
|
+
try:
|
|
426
|
+
row = self._get_attempt_row_locked(attempt_id)
|
|
427
|
+
if AttemptStatus(row["status"]) is not AttemptStatus.RUNNING:
|
|
428
|
+
raise ValueError("only a running attempt can be finished")
|
|
429
|
+
self._connection.execute(
|
|
430
|
+
"""UPDATE trial_attempts SET status = ?, finished_at = ?, error = ?
|
|
431
|
+
WHERE id = ?""",
|
|
432
|
+
(status.value, finished_at, error, attempt_id),
|
|
433
|
+
)
|
|
434
|
+
self._insert_usage_locked(row["trial_id"], attempt_id, usage)
|
|
435
|
+
trial = self._get_trial_row_locked(row["trial_id"])
|
|
436
|
+
self._append_event_locked(
|
|
437
|
+
trial["study_id"],
|
|
438
|
+
f"trial.attempt.{status.value}",
|
|
439
|
+
{"trial_id": row["trial_id"], "attempt_id": attempt_id},
|
|
440
|
+
)
|
|
441
|
+
self._connection.commit()
|
|
442
|
+
except BaseException:
|
|
443
|
+
self._connection.rollback()
|
|
444
|
+
raise
|
|
445
|
+
return self.get_attempt(attempt_id)
|
|
446
|
+
|
|
447
|
+
def abandon_expired_attempts(
|
|
448
|
+
self, *, study_id: str | None = None, now: float | None = None
|
|
449
|
+
) -> list[AttemptRecord]:
|
|
450
|
+
"""Persist abandonment for expired attempts, optionally scoped to one study."""
|
|
451
|
+
|
|
452
|
+
timestamp = self._timestamp(now)
|
|
453
|
+
abandoned_ids: list[str] = []
|
|
454
|
+
self._begin_immediate()
|
|
455
|
+
try:
|
|
456
|
+
if study_id is None:
|
|
457
|
+
rows = self._connection.execute(
|
|
458
|
+
"""SELECT id, trial_id FROM trial_attempts
|
|
459
|
+
WHERE status = 'running' AND lease_expires_at <= ?
|
|
460
|
+
ORDER BY started_at, id""",
|
|
461
|
+
(timestamp,),
|
|
462
|
+
).fetchall()
|
|
463
|
+
else:
|
|
464
|
+
self._get_study_status_locked(study_id)
|
|
465
|
+
rows = self._connection.execute(
|
|
466
|
+
"""SELECT a.id, a.trial_id FROM trial_attempts AS a
|
|
467
|
+
JOIN trials AS t ON t.id = a.trial_id
|
|
468
|
+
WHERE t.study_id = ? AND a.status = 'running'
|
|
469
|
+
AND a.lease_expires_at <= ?
|
|
470
|
+
ORDER BY a.started_at, a.id""",
|
|
471
|
+
(study_id, timestamp),
|
|
472
|
+
).fetchall()
|
|
473
|
+
for row in rows:
|
|
474
|
+
self._connection.execute(
|
|
475
|
+
"""UPDATE trial_attempts
|
|
476
|
+
SET status = 'abandoned', finished_at = ?, error = 'lease expired'
|
|
477
|
+
WHERE id = ?""",
|
|
478
|
+
(timestamp, row["id"]),
|
|
479
|
+
)
|
|
480
|
+
trial = self._get_trial_row_locked(row["trial_id"])
|
|
481
|
+
self._append_event_locked(
|
|
482
|
+
trial["study_id"],
|
|
483
|
+
"trial.attempt.abandoned",
|
|
484
|
+
{"trial_id": row["trial_id"], "attempt_id": row["id"]},
|
|
485
|
+
)
|
|
486
|
+
abandoned_ids.append(row["id"])
|
|
487
|
+
self._connection.commit()
|
|
488
|
+
except BaseException:
|
|
489
|
+
self._connection.rollback()
|
|
490
|
+
raise
|
|
491
|
+
return [self.get_attempt(attempt_id) for attempt_id in abandoned_ids]
|
|
492
|
+
|
|
493
|
+
def list_recoverable_trials(self, study_id: str) -> list[TrialRecord]:
|
|
494
|
+
"""Return running trials that have no active attempt."""
|
|
495
|
+
|
|
496
|
+
rows = self._connection.execute(
|
|
497
|
+
"""SELECT t.id, t.study_id, t.sequence, t.candidate_hash, t.candidate_json,
|
|
498
|
+
t.status, t.error, t.created_at, t.updated_at
|
|
499
|
+
FROM trials AS t
|
|
500
|
+
WHERE t.study_id = ? AND t.status = 'running'
|
|
501
|
+
AND NOT EXISTS (
|
|
502
|
+
SELECT 1 FROM trial_attempts AS a
|
|
503
|
+
WHERE a.trial_id = t.id AND a.status = 'running'
|
|
504
|
+
)
|
|
505
|
+
ORDER BY t.sequence""",
|
|
506
|
+
(study_id,),
|
|
507
|
+
).fetchall()
|
|
508
|
+
return [_trial_from_row(row) for row in rows]
|
|
509
|
+
|
|
510
|
+
def commit_trial_result(
|
|
511
|
+
self,
|
|
512
|
+
trial_id: str,
|
|
513
|
+
*,
|
|
514
|
+
status: TrialStatus,
|
|
515
|
+
metrics: Sequence[MetricValue] = (),
|
|
516
|
+
usage: Sequence[UsageRecord] = (),
|
|
517
|
+
attempt_id: str | None = None,
|
|
518
|
+
error: str | None = None,
|
|
519
|
+
now: float | None = None,
|
|
520
|
+
) -> TrialRecord:
|
|
521
|
+
"""Atomically persist observations, attempt completion, and terminal trial state."""
|
|
522
|
+
|
|
523
|
+
if status not in _TERMINAL_TRIAL_STATUSES:
|
|
524
|
+
raise ValueError("trial result status must be terminal")
|
|
525
|
+
metric_names = [metric.name for metric in metrics]
|
|
526
|
+
if len(metric_names) != len(set(metric_names)):
|
|
527
|
+
raise ValueError("trial result contains duplicate metric names")
|
|
528
|
+
finished_at = self._timestamp(now)
|
|
529
|
+
self._begin_immediate()
|
|
530
|
+
try:
|
|
531
|
+
trial = self._get_trial_row_locked(trial_id)
|
|
532
|
+
if TrialStatus(trial["status"]) is not TrialStatus.RUNNING:
|
|
533
|
+
raise ValueError("only a running trial can commit a result")
|
|
534
|
+
if attempt_id is not None:
|
|
535
|
+
attempt = self._get_attempt_row_locked(attempt_id)
|
|
536
|
+
if attempt["trial_id"] != trial_id:
|
|
537
|
+
raise ValueError("attempt does not belong to trial")
|
|
538
|
+
if AttemptStatus(attempt["status"]) is not AttemptStatus.RUNNING:
|
|
539
|
+
raise ValueError("only a running attempt can commit a trial result")
|
|
540
|
+
attempt_status = _attempt_status_for_trial(status)
|
|
541
|
+
self._connection.execute(
|
|
542
|
+
"""UPDATE trial_attempts SET status = ?, finished_at = ?, error = ?
|
|
543
|
+
WHERE id = ?""",
|
|
544
|
+
(attempt_status.value, finished_at, error, attempt_id),
|
|
545
|
+
)
|
|
546
|
+
for metric in metrics:
|
|
547
|
+
self._connection.execute(
|
|
548
|
+
"""INSERT INTO metric_values(
|
|
549
|
+
trial_id, name, value, status, coverage
|
|
550
|
+
) VALUES (?, ?, ?, ?, ?)""",
|
|
551
|
+
(trial_id, metric.name, metric.value, metric.status.value, metric.coverage),
|
|
552
|
+
)
|
|
553
|
+
self._insert_usage_locked(trial_id, attempt_id, usage)
|
|
554
|
+
self._connection.execute(
|
|
555
|
+
"""UPDATE trials SET status = ?, error = ?, updated_at = CURRENT_TIMESTAMP
|
|
556
|
+
WHERE id = ?""",
|
|
557
|
+
(status.value, error, trial_id),
|
|
558
|
+
)
|
|
559
|
+
self._append_event_locked(
|
|
560
|
+
trial["study_id"],
|
|
561
|
+
f"trial.{status.value}",
|
|
562
|
+
{
|
|
563
|
+
"trial_id": trial_id,
|
|
564
|
+
"attempt_id": attempt_id,
|
|
565
|
+
"metric_count": len(metrics),
|
|
566
|
+
"usage_count": len(usage),
|
|
567
|
+
},
|
|
568
|
+
)
|
|
569
|
+
self._connection.commit()
|
|
570
|
+
except BaseException:
|
|
571
|
+
self._connection.rollback()
|
|
572
|
+
raise
|
|
573
|
+
return self.get_trial(trial_id)
|
|
574
|
+
|
|
575
|
+
def list_metrics(self, trial_id: str) -> list[MetricValue]:
|
|
576
|
+
"""Return committed metrics ordered by name."""
|
|
577
|
+
|
|
578
|
+
rows = self._connection.execute(
|
|
579
|
+
"""SELECT name, value, status, coverage FROM metric_values
|
|
580
|
+
WHERE trial_id = ? ORDER BY name""",
|
|
581
|
+
(trial_id,),
|
|
582
|
+
).fetchall()
|
|
583
|
+
return [
|
|
584
|
+
MetricValue(row["name"], row["value"], ValueStatus(row["status"]), row["coverage"])
|
|
585
|
+
for row in rows
|
|
586
|
+
]
|
|
587
|
+
|
|
588
|
+
def list_usage(self, trial_id: str) -> list[PersistedUsageRecord]:
|
|
589
|
+
"""Return append-only usage in insertion order."""
|
|
590
|
+
|
|
591
|
+
rows = self._connection.execute(
|
|
592
|
+
"""SELECT id, trial_id, attempt_id, component, input_tokens, output_tokens,
|
|
593
|
+
cost_text, latency_seconds, status, pricing_version, created_at
|
|
594
|
+
FROM usage_records WHERE trial_id = ? ORDER BY rowid""",
|
|
595
|
+
(trial_id,),
|
|
596
|
+
).fetchall()
|
|
597
|
+
records: list[PersistedUsageRecord] = []
|
|
598
|
+
for row in rows:
|
|
599
|
+
cost = None if row["cost_text"] is None else float(row["cost_text"])
|
|
600
|
+
records.append(
|
|
601
|
+
PersistedUsageRecord(
|
|
602
|
+
id=row["id"],
|
|
603
|
+
trial_id=row["trial_id"],
|
|
604
|
+
attempt_id=row["attempt_id"],
|
|
605
|
+
usage=UsageRecord(
|
|
606
|
+
component=row["component"],
|
|
607
|
+
input_tokens=row["input_tokens"],
|
|
608
|
+
output_tokens=row["output_tokens"],
|
|
609
|
+
cost=cost,
|
|
610
|
+
latency_seconds=row["latency_seconds"],
|
|
611
|
+
status=ValueStatus(row["status"]),
|
|
612
|
+
pricing_version=row["pricing_version"],
|
|
613
|
+
),
|
|
614
|
+
created_at=row["created_at"],
|
|
615
|
+
)
|
|
616
|
+
)
|
|
617
|
+
return records
|
|
618
|
+
|
|
619
|
+
def append_event(
|
|
620
|
+
self, study_id: str, event_type: str, payload: Mapping[str, Any]
|
|
621
|
+
) -> StudyEvent:
|
|
622
|
+
"""Append a canonical event with an atomic per-study sequence."""
|
|
623
|
+
|
|
624
|
+
self._validate_event_type(event_type)
|
|
625
|
+
self._begin_immediate()
|
|
626
|
+
try:
|
|
627
|
+
self._get_study_status_locked(study_id)
|
|
628
|
+
event = self._append_event_locked(study_id, event_type, payload)
|
|
629
|
+
self._connection.commit()
|
|
630
|
+
except BaseException:
|
|
631
|
+
self._connection.rollback()
|
|
632
|
+
raise
|
|
633
|
+
return event
|
|
634
|
+
|
|
635
|
+
def list_events(self, study_id: str, *, after_sequence: int = -1) -> list[StudyEvent]:
|
|
636
|
+
"""Return durable events after a sequence in replay order."""
|
|
637
|
+
|
|
638
|
+
rows = self._connection.execute(
|
|
639
|
+
"""SELECT study_id, sequence, type, payload_json, created_at
|
|
640
|
+
FROM study_events WHERE study_id = ? AND sequence > ? ORDER BY sequence""",
|
|
641
|
+
(study_id, after_sequence),
|
|
642
|
+
).fetchall()
|
|
643
|
+
return [_event_from_row(row) for row in rows]
|
|
644
|
+
|
|
645
|
+
def close(self) -> None:
|
|
646
|
+
"""Close the underlying SQLite connection."""
|
|
647
|
+
|
|
648
|
+
self._connection.close()
|
|
649
|
+
|
|
650
|
+
def _insert_usage_locked(
|
|
651
|
+
self,
|
|
652
|
+
trial_id: str,
|
|
653
|
+
attempt_id: str | None,
|
|
654
|
+
usage: Sequence[UsageRecord],
|
|
655
|
+
) -> None:
|
|
656
|
+
for item in usage:
|
|
657
|
+
self._connection.execute(
|
|
658
|
+
"""INSERT INTO usage_records(
|
|
659
|
+
id, trial_id, attempt_id, component, input_tokens, output_tokens,
|
|
660
|
+
cost_text, latency_seconds, status, pricing_version
|
|
661
|
+
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
|
662
|
+
(
|
|
663
|
+
str(uuid.uuid4()),
|
|
664
|
+
trial_id,
|
|
665
|
+
attempt_id,
|
|
666
|
+
item.component,
|
|
667
|
+
item.input_tokens,
|
|
668
|
+
item.output_tokens,
|
|
669
|
+
None if item.cost is None else format(item.cost, ".17g"),
|
|
670
|
+
item.latency_seconds,
|
|
671
|
+
item.status.value,
|
|
672
|
+
item.pricing_version,
|
|
673
|
+
),
|
|
674
|
+
)
|
|
675
|
+
|
|
676
|
+
def _insert_trial_locked(
|
|
677
|
+
self, study_id: str, sequence: int, candidate: Candidate
|
|
678
|
+
) -> TrialRecord:
|
|
679
|
+
candidate_bytes = storage_json(candidate.parameters)
|
|
680
|
+
trial_id = str(uuid.uuid4())
|
|
681
|
+
self._connection.execute(
|
|
682
|
+
"""INSERT INTO trials(
|
|
683
|
+
id, study_id, sequence, candidate_hash, candidate_json, status, updated_at
|
|
684
|
+
) VALUES (?, ?, ?, ?, ?, 'pending', CURRENT_TIMESTAMP)""",
|
|
685
|
+
(
|
|
686
|
+
trial_id,
|
|
687
|
+
study_id,
|
|
688
|
+
sequence,
|
|
689
|
+
content_hash(candidate.parameters, namespace="candidate:v1"),
|
|
690
|
+
candidate_bytes.decode("utf-8"),
|
|
691
|
+
),
|
|
692
|
+
)
|
|
693
|
+
self._append_event_locked(
|
|
694
|
+
study_id, "trial.reserved", {"trial_id": trial_id, "sequence": sequence}
|
|
695
|
+
)
|
|
696
|
+
return self.get_trial(trial_id)
|
|
697
|
+
|
|
698
|
+
def _append_event_locked(
|
|
699
|
+
self, study_id: str, event_type: str, payload: Mapping[str, Any]
|
|
700
|
+
) -> StudyEvent:
|
|
701
|
+
self._validate_event_type(event_type)
|
|
702
|
+
row = self._connection.execute(
|
|
703
|
+
"SELECT COALESCE(MAX(sequence), -1) + 1 FROM study_events WHERE study_id = ?",
|
|
704
|
+
(study_id,),
|
|
705
|
+
).fetchone()
|
|
706
|
+
sequence = int(cast(sqlite3.Row, row)[0])
|
|
707
|
+
safe_payload = cast(dict[str, Any], redact(payload))
|
|
708
|
+
payload_json = canonical_json(safe_payload).decode("utf-8")
|
|
709
|
+
self._connection.execute(
|
|
710
|
+
"""INSERT INTO study_events(study_id, sequence, type, payload_json)
|
|
711
|
+
VALUES (?, ?, ?, ?)""",
|
|
712
|
+
(study_id, sequence, event_type, payload_json),
|
|
713
|
+
)
|
|
714
|
+
row = self._connection.execute(
|
|
715
|
+
"""SELECT study_id, sequence, type, payload_json, created_at
|
|
716
|
+
FROM study_events WHERE study_id = ? AND sequence = ?""",
|
|
717
|
+
(study_id, sequence),
|
|
718
|
+
).fetchone()
|
|
719
|
+
return _event_from_row(cast(sqlite3.Row, row))
|
|
720
|
+
|
|
721
|
+
def _get_study_status_locked(self, study_id: str) -> StudyStatus:
|
|
722
|
+
row = self._connection.execute(
|
|
723
|
+
"SELECT status FROM studies WHERE id = ?", (study_id,)
|
|
724
|
+
).fetchone()
|
|
725
|
+
if row is None:
|
|
726
|
+
raise KeyError(f"unknown study: {study_id}")
|
|
727
|
+
return StudyStatus(row["status"])
|
|
728
|
+
|
|
729
|
+
def _require_running_study_locked(self, study_id: str) -> None:
|
|
730
|
+
status = self._get_study_status_locked(study_id)
|
|
731
|
+
if status is not StudyStatus.RUNNING:
|
|
732
|
+
raise ValueError(f"study must be running to reserve trials, not {status.value}")
|
|
733
|
+
|
|
734
|
+
def _get_trial_row_locked(self, trial_id: str) -> sqlite3.Row:
|
|
735
|
+
row = self._connection.execute(
|
|
736
|
+
"""SELECT id, study_id, sequence, candidate_hash, candidate_json, status,
|
|
737
|
+
error, created_at, updated_at FROM trials WHERE id = ?""",
|
|
738
|
+
(trial_id,),
|
|
739
|
+
).fetchone()
|
|
740
|
+
if row is None:
|
|
741
|
+
raise KeyError(f"unknown trial: {trial_id}")
|
|
742
|
+
return cast(sqlite3.Row, row)
|
|
743
|
+
|
|
744
|
+
def _get_attempt_row_locked(self, attempt_id: str) -> sqlite3.Row:
|
|
745
|
+
row = self._connection.execute(
|
|
746
|
+
"""SELECT id, trial_id, attempt_number, status, started_at,
|
|
747
|
+
lease_expires_at, finished_at, error FROM trial_attempts WHERE id = ?""",
|
|
748
|
+
(attempt_id,),
|
|
749
|
+
).fetchone()
|
|
750
|
+
if row is None:
|
|
751
|
+
raise KeyError(f"unknown attempt: {attempt_id}")
|
|
752
|
+
return cast(sqlite3.Row, row)
|
|
753
|
+
|
|
754
|
+
def _schema_versions(self) -> list[int]:
|
|
755
|
+
return [
|
|
756
|
+
int(row[0])
|
|
757
|
+
for row in self._connection.execute(
|
|
758
|
+
"SELECT version FROM schema_version ORDER BY version"
|
|
759
|
+
)
|
|
760
|
+
]
|
|
761
|
+
|
|
762
|
+
def _create_schema_version_1(self) -> None:
|
|
763
|
+
try:
|
|
764
|
+
self._connection.executescript(
|
|
765
|
+
"""
|
|
766
|
+
BEGIN IMMEDIATE;
|
|
767
|
+
CREATE TABLE studies (
|
|
768
|
+
id TEXT PRIMARY KEY,
|
|
769
|
+
name TEXT NOT NULL,
|
|
770
|
+
config_hash TEXT NOT NULL,
|
|
771
|
+
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
|
|
772
|
+
);
|
|
773
|
+
CREATE TABLE trials (
|
|
774
|
+
id TEXT PRIMARY KEY,
|
|
775
|
+
study_id TEXT NOT NULL REFERENCES studies(id),
|
|
776
|
+
sequence INTEGER NOT NULL,
|
|
777
|
+
candidate_hash TEXT NOT NULL,
|
|
778
|
+
candidate_json TEXT NOT NULL,
|
|
779
|
+
status TEXT NOT NULL,
|
|
780
|
+
error TEXT,
|
|
781
|
+
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
782
|
+
UNIQUE(study_id, sequence)
|
|
783
|
+
);
|
|
784
|
+
INSERT INTO schema_version(version) VALUES (1);
|
|
785
|
+
COMMIT;
|
|
786
|
+
"""
|
|
787
|
+
)
|
|
788
|
+
except sqlite3.Error as error:
|
|
789
|
+
self._connection.rollback()
|
|
790
|
+
self._connection.close()
|
|
791
|
+
raise StoreError("could not create experiment schema", stage="store.migrate") from error
|
|
792
|
+
|
|
793
|
+
def _migrate_version_1_to_2(self) -> None:
|
|
794
|
+
try:
|
|
795
|
+
self._connection.executescript(
|
|
796
|
+
"""
|
|
797
|
+
BEGIN IMMEDIATE;
|
|
798
|
+
ALTER TABLE studies ADD COLUMN status TEXT NOT NULL DEFAULT 'created';
|
|
799
|
+
ALTER TABLE studies ADD COLUMN updated_at TEXT;
|
|
800
|
+
UPDATE studies SET updated_at = created_at WHERE updated_at IS NULL;
|
|
801
|
+
ALTER TABLE trials ADD COLUMN updated_at TEXT;
|
|
802
|
+
UPDATE trials SET updated_at = created_at WHERE updated_at IS NULL;
|
|
803
|
+
CREATE UNIQUE INDEX uq_trials_study_candidate
|
|
804
|
+
ON trials(study_id, candidate_hash);
|
|
805
|
+
CREATE TABLE trial_attempts (
|
|
806
|
+
id TEXT PRIMARY KEY,
|
|
807
|
+
trial_id TEXT NOT NULL REFERENCES trials(id),
|
|
808
|
+
attempt_number INTEGER NOT NULL,
|
|
809
|
+
status TEXT NOT NULL,
|
|
810
|
+
started_at REAL NOT NULL,
|
|
811
|
+
lease_expires_at REAL NOT NULL,
|
|
812
|
+
finished_at REAL,
|
|
813
|
+
error TEXT,
|
|
814
|
+
UNIQUE(trial_id, attempt_number),
|
|
815
|
+
CHECK(attempt_number >= 0),
|
|
816
|
+
CHECK(status IN ('running', 'succeeded', 'failed', 'abandoned', 'cancelled'))
|
|
817
|
+
);
|
|
818
|
+
CREATE UNIQUE INDEX uq_trial_attempts_running
|
|
819
|
+
ON trial_attempts(trial_id) WHERE status = 'running';
|
|
820
|
+
CREATE INDEX idx_trial_attempts_recovery
|
|
821
|
+
ON trial_attempts(status, lease_expires_at);
|
|
822
|
+
CREATE TABLE metric_values (
|
|
823
|
+
trial_id TEXT NOT NULL REFERENCES trials(id),
|
|
824
|
+
name TEXT NOT NULL,
|
|
825
|
+
value REAL,
|
|
826
|
+
status TEXT NOT NULL,
|
|
827
|
+
coverage REAL NOT NULL,
|
|
828
|
+
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
829
|
+
PRIMARY KEY(trial_id, name),
|
|
830
|
+
CHECK(status IN ('exact', 'estimated', 'unavailable')),
|
|
831
|
+
CHECK(coverage >= 0.0 AND coverage <= 1.0)
|
|
832
|
+
);
|
|
833
|
+
CREATE TABLE usage_records (
|
|
834
|
+
id TEXT PRIMARY KEY,
|
|
835
|
+
trial_id TEXT NOT NULL REFERENCES trials(id),
|
|
836
|
+
attempt_id TEXT REFERENCES trial_attempts(id),
|
|
837
|
+
component TEXT NOT NULL,
|
|
838
|
+
input_tokens INTEGER,
|
|
839
|
+
output_tokens INTEGER,
|
|
840
|
+
cost_text TEXT,
|
|
841
|
+
latency_seconds REAL,
|
|
842
|
+
status TEXT NOT NULL,
|
|
843
|
+
pricing_version TEXT,
|
|
844
|
+
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
845
|
+
CHECK(status IN ('exact', 'estimated', 'unavailable'))
|
|
846
|
+
);
|
|
847
|
+
CREATE INDEX idx_usage_records_trial ON usage_records(trial_id);
|
|
848
|
+
CREATE TABLE study_events (
|
|
849
|
+
study_id TEXT NOT NULL REFERENCES studies(id),
|
|
850
|
+
sequence INTEGER NOT NULL,
|
|
851
|
+
type TEXT NOT NULL,
|
|
852
|
+
payload_json TEXT NOT NULL,
|
|
853
|
+
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
854
|
+
PRIMARY KEY(study_id, sequence)
|
|
855
|
+
);
|
|
856
|
+
INSERT INTO schema_version(version) VALUES (2);
|
|
857
|
+
COMMIT;
|
|
858
|
+
"""
|
|
859
|
+
)
|
|
860
|
+
except sqlite3.Error as error:
|
|
861
|
+
self._connection.rollback()
|
|
862
|
+
self._connection.close()
|
|
863
|
+
raise StoreError(
|
|
864
|
+
"could not migrate experiment schema to version 2",
|
|
865
|
+
stage="store.migrate",
|
|
866
|
+
details={"target_version": 2, "cause": type(error).__name__},
|
|
867
|
+
) from error
|
|
868
|
+
|
|
869
|
+
def _begin_immediate(self) -> None:
|
|
870
|
+
self._connection.execute("BEGIN IMMEDIATE")
|
|
871
|
+
|
|
872
|
+
def _timestamp(self, value: float | None) -> float:
|
|
873
|
+
timestamp = float(self._clock() if value is None else value)
|
|
874
|
+
if not math.isfinite(timestamp):
|
|
875
|
+
raise ValueError("timestamp must be finite")
|
|
876
|
+
return timestamp
|
|
877
|
+
|
|
878
|
+
@staticmethod
|
|
879
|
+
def _validate_event_type(event_type: str) -> None:
|
|
880
|
+
if not event_type.strip():
|
|
881
|
+
raise ValueError("event type must not be empty")
|
|
882
|
+
|
|
883
|
+
|
|
884
|
+
def _study_from_row(row: sqlite3.Row) -> StudyRecord:
|
|
885
|
+
return StudyRecord(
|
|
886
|
+
id=row["id"],
|
|
887
|
+
name=row["name"],
|
|
888
|
+
config_hash=row["config_hash"],
|
|
889
|
+
status=StudyStatus(row["status"]),
|
|
890
|
+
created_at=row["created_at"],
|
|
891
|
+
updated_at=row["updated_at"],
|
|
892
|
+
)
|
|
893
|
+
|
|
894
|
+
|
|
895
|
+
def _trial_from_row(row: sqlite3.Row) -> TrialRecord:
|
|
896
|
+
return TrialRecord(
|
|
897
|
+
id=row["id"],
|
|
898
|
+
study_id=row["study_id"],
|
|
899
|
+
sequence=row["sequence"],
|
|
900
|
+
candidate_hash=row["candidate_hash"],
|
|
901
|
+
candidate_json=row["candidate_json"],
|
|
902
|
+
status=TrialStatus(row["status"]),
|
|
903
|
+
error=row["error"],
|
|
904
|
+
created_at=row["created_at"],
|
|
905
|
+
updated_at=row["updated_at"],
|
|
906
|
+
)
|
|
907
|
+
|
|
908
|
+
|
|
909
|
+
def _attempt_from_row(row: sqlite3.Row) -> AttemptRecord:
|
|
910
|
+
return AttemptRecord(
|
|
911
|
+
id=row["id"],
|
|
912
|
+
trial_id=row["trial_id"],
|
|
913
|
+
number=row["attempt_number"],
|
|
914
|
+
status=AttemptStatus(row["status"]),
|
|
915
|
+
started_at=row["started_at"],
|
|
916
|
+
lease_expires_at=row["lease_expires_at"],
|
|
917
|
+
finished_at=row["finished_at"],
|
|
918
|
+
error=row["error"],
|
|
919
|
+
)
|
|
920
|
+
|
|
921
|
+
|
|
922
|
+
def _event_from_row(row: sqlite3.Row) -> StudyEvent:
|
|
923
|
+
payload = json.loads(row["payload_json"])
|
|
924
|
+
if not isinstance(payload, dict):
|
|
925
|
+
raise StoreError("event payload is not an object", stage="store.read")
|
|
926
|
+
return StudyEvent(
|
|
927
|
+
study_id=row["study_id"],
|
|
928
|
+
sequence=row["sequence"],
|
|
929
|
+
type=row["type"],
|
|
930
|
+
payload=cast(dict[str, Any], payload),
|
|
931
|
+
created_at=row["created_at"],
|
|
932
|
+
)
|
|
933
|
+
|
|
934
|
+
|
|
935
|
+
def _attempt_status_for_trial(status: TrialStatus) -> AttemptStatus:
|
|
936
|
+
return {
|
|
937
|
+
TrialStatus.SUCCEEDED: AttemptStatus.SUCCEEDED,
|
|
938
|
+
TrialStatus.FAILED: AttemptStatus.FAILED,
|
|
939
|
+
TrialStatus.ABANDONED: AttemptStatus.ABANDONED,
|
|
940
|
+
TrialStatus.CANCELLED: AttemptStatus.CANCELLED,
|
|
941
|
+
}[status]
|