diff --git a/cueweaver/application/database.py b/cueweaver/application/database.py index 3c953e5..1769dfe 100644 --- a/cueweaver/application/database.py +++ b/cueweaver/application/database.py @@ -42,6 +42,7 @@ class JobRow(Base): stream_index: Mapped[int | None] = mapped_column(Integer) target_language_code: Mapped[str] = mapped_column(String, nullable=False) term_map_mode: Mapped[str] = mapped_column(String, nullable=False) + # These columns are immutable Job-owned snapshot metadata, not live-map fields. term_map_id: Mapped[str | None] = mapped_column(String) term_map_name: Mapped[str | None] = mapped_column(String) output_path: Mapped[str] = mapped_column(String, nullable=False) diff --git a/cueweaver/application/jobs/__init__.py b/cueweaver/application/jobs/__init__.py index 31b6ff7..ebc29af 100644 --- a/cueweaver/application/jobs/__init__.py +++ b/cueweaver/application/jobs/__init__.py @@ -206,9 +206,7 @@ def create(self, request: CreateJobRequest) -> dict[str, object]: } if subtitle is None: record["extraction"] = None - self._write_record(job_id, record) - with self._lock: - self._records[job_id] = record + self._write_record(record) self._pending.put(job_id) return self._record_with_queue_position(record) @@ -305,8 +303,7 @@ def retry(self, job_id: str) -> dict[str, object]: "message": error.message, **context, } - self._write_record(job_id, failed_record) - self._records[job_id] = failed_record + self._write_record(failed_record) raise safe_error from error with self._lock: retry_record = copy_job_record(record) @@ -331,8 +328,7 @@ def retry(self, job_id: str) -> dict[str, object]: self._base_output_path(retry_request).relative_to(self._media_root) ) retry_record["queue_sequence"] = next_queue_sequence - self._write_record(job_id, retry_record) - self._records[job_id] = retry_record + self._write_record(retry_record) self._next_queue_sequence = next_queue_sequence self._pending.put(job_id) return self._record_with_queue_position(retry_record) @@ -357,8 +353,7 @@ def cancel(self, job_id: str) -> dict[str, object]: ) cancelled_record["finished_at"] = cancelled_at cancelled_record["error"] = None - self._write_record(job_id, cancelled_record) - self._records[job_id] = cancelled_record + self._write_record(cancelled_record) return self._record_with_queue_position(cancelled_record) def delete(self, job_id: str) -> dict[str, object]: @@ -861,13 +856,14 @@ def _load_records(self) -> None: record = copy_job_record(loaded_record) if status in {"Extracting", "Translating"}: record = _interrupted_record(record) - self._write_record(job_id, record) + self._write_record(record) elif status == "Queued": self._recovered_queue_ids.append(job_id) + if status not in {"Extracting", "Translating"}: + self._records[job_id] = copy_job_record(record) self._next_queue_sequence = max( self._next_queue_sequence, queue_sequence(record) ) - self._records[job_id] = copy_job_record(record) def _run(self) -> None: while True: @@ -1045,10 +1041,9 @@ def _persist_embedded_progress( } transition_status(record, progress.phase, at=_timestamp()) try: - self._write_record(job_id, record) + self._write_record(record) except Exception as error: raise JobExecutionProgressPersistenceError from error - self._records[job_id] = record return True def _prepare_execution( @@ -1071,7 +1066,7 @@ def _prepare_execution( record["started_at"] = started_at output_path = self._execution_output_path(request) request["output_path"] = str(output_path.relative_to(self._media_root)) - self._write_record(job_id, record) + self._write_record(record) return request, embedded, self._jobs_root / job_id, record def _execution_output_path(self, request: dict[str, object]) -> Path: @@ -1126,13 +1121,13 @@ def _finish( transition_status(record, status, at=finished_at, terminal=True) record["finished_at"] = finished_at record["error"] = error - self._write_record(job_id, record) + self._write_record(record) def _finish_interrupted(self, job_id: str) -> None: with self._lock: interrupted = _interrupted_record(self._records[job_id]) try: - self._write_record(job_id, interrupted) + self._write_record(interrupted) except (OSError, ServiceError) as error: logger.warning( "Could not persist interrupted Job %s during shutdown: %s", @@ -1157,7 +1152,7 @@ def _mark_failed_after_worker_error(self, job_id: str, error: Exception) -> None "message": "Job execution could not be persisted", } try: - self._write_record(job_id, record) + self._write_record(record) except Exception as persistence_error: logger.error( "Could not persist worker failure for Job %s: %s", @@ -1165,9 +1160,11 @@ def _mark_failed_after_worker_error(self, job_id: str, error: Exception) -> None persistence_error, ) - def _write_record(self, job_id: str, record: dict[str, object]) -> None: + def _write_record(self, record: dict[str, object]) -> None: persisted = copy_job_record(record) self._record_store.write(persisted) + job_id = persisted["id"] + assert isinstance(job_id, str) self._records[job_id] = persisted def _check_jobs_root(self) -> None: diff --git a/cueweaver/application/jobs/store.py b/cueweaver/application/jobs/store.py index f8b3542..311a7f9 100644 --- a/cueweaver/application/jobs/store.py +++ b/cueweaver/application/jobs/store.py @@ -43,12 +43,39 @@ def __init__(self, database: SqliteDatabase) -> None: def load(self) -> list[JobRecord]: try: with self._database.read_session() as session: - rows = session.scalars( + job_rows = session.scalars( select(JobRow).order_by( JobRow.queue_sequence, JobRow.created_at, JobRow.id ) ).all() - return [_record_from_row(session, row) for row in rows] + history_rows = session.scalars( + select(JobStatusHistoryRow).order_by( + JobStatusHistoryRow.job_id, JobStatusHistoryRow.sequence + ) + ).all() + snapshot_rows = session.scalars( + select(JobTermMapSnapshotRow).order_by( + JobTermMapSnapshotRow.job_id, JobTermMapSnapshotRow.position + ) + ).all() + histories_by_job: dict[str, list[JobStatusHistoryRow]] = {} + for history_row in history_rows: + histories_by_job.setdefault(history_row.job_id, []).append( + history_row + ) + snapshots_by_job: dict[str, list[JobTermMapSnapshotRow]] = {} + for snapshot_row in snapshot_rows: + snapshots_by_job.setdefault(snapshot_row.job_id, []).append( + snapshot_row + ) + return [ + _record_from_row( + row, + histories_by_job.get(row.id, []), + snapshots_by_job.get(row.id, []), + ) + for row in job_rows + ] except (sqlite3.Error, SQLAlchemyError) as error: raise ServiceError( "database_unavailable", "Job records cannot be loaded" @@ -104,6 +131,7 @@ def _upsert_row(session: Session, record: JobRecord) -> None: extraction = record.get("extraction") extraction_values = extraction if isinstance(extraction, dict) else {} row = session.get(JobRow, record["id"]) + is_new_job = row is None if row is None: row = JobRow(id=str(record["id"])) session.add(row) @@ -122,6 +150,9 @@ def _upsert_row(session: Session, record: JobRecord) -> None: term_map = _set_request_fields(row, request) _set_extraction_fields(row, extraction_values) + if is_new_job: + _set_snapshot_fields(row, term_map, session) + session.execute( delete(JobStatusHistoryRow).where(JobStatusHistoryRow.job_id == row.id) ) @@ -140,36 +171,32 @@ def _upsert_row(session: Session, record: JobRecord) -> None: ) ) - # A Job owns this snapshot. Once populated, later Job writes cannot replace it. - has_snapshot = ( - session.scalar( - select(JobTermMapSnapshotRow.position) - .where(JobTermMapSnapshotRow.job_id == row.id) - .limit(1) + +def _set_snapshot_fields(row: JobRow, term_map: object, session: Session) -> None: + if not isinstance(term_map, dict): + return + row.term_map_id = _optional_str(term_map.get("id")) + row.term_map_name = _optional_str(term_map.get("name")) + content = term_map.get("content") + if not isinstance(content, dict): + return + for position, (source, target) in enumerate(content.items()): + session.add( + JobTermMapSnapshotRow( + job_id=row.id, + position=position, + source=str(source), + source_folded=str(source).casefold(), + target=str(target), + ) ) - is not None - ) - if not has_snapshot and isinstance(term_map, dict): - content = term_map.get("content") - if isinstance(content, dict): - for position, (source, target) in enumerate(content.items()): - session.add( - JobTermMapSnapshotRow( - job_id=row.id, - position=position, - source=str(source), - source_folded=str(source).casefold(), - target=str(target), - ) - ) -def _record_from_row(session: Session, row: JobRow) -> JobRecord: - snapshot_rows = session.scalars( - select(JobTermMapSnapshotRow) - .where(JobTermMapSnapshotRow.job_id == row.id) - .order_by(JobTermMapSnapshotRow.position) - ).all() +def _record_from_row( + row: JobRow, + history_rows: list[JobStatusHistoryRow], + snapshot_rows: list[JobTermMapSnapshotRow], +) -> JobRecord: content = {item.source: item.target for item in snapshot_rows} term_map: dict[str, object] | None = None if row.term_map_id is not None: @@ -214,11 +241,7 @@ def _record_from_row(session: Session, row: JobRow) -> JobRecord: "started_at": item.started_at, "finished_at": item.finished_at, } - for item in session.scalars( - select(JobStatusHistoryRow) - .where(JobStatusHistoryRow.job_id == row.id) - .order_by(JobStatusHistoryRow.sequence) - ).all() + for item in history_rows ], } if row.stream_index is not None: @@ -256,12 +279,6 @@ def _set_request_fields(row: JobRow, request: dict[str, object]) -> object: row.target_language_code = str(request["target_language_code"]) row.term_map_mode = str(request["term_map_mode"]) term_map = request.get("term_map") - row.term_map_id = ( - _optional_str(term_map.get("id")) if isinstance(term_map, dict) else None - ) - row.term_map_name = ( - _optional_str(term_map.get("name")) if isinstance(term_map, dict) else None - ) row.output_path = str(request["output_path"]) row.source_format = str(request["source_format"]) row.dynamic_terminology_enabled = bool(request["dynamic_terminology_enabled"]) diff --git a/tests/test_jobs.py b/tests/test_jobs.py index 3192a84..99393e6 100644 --- a/tests/test_jobs.py +++ b/tests/test_jobs.py @@ -9,6 +9,8 @@ import pytest from fastapi.testclient import TestClient +from sqlalchemy import event +from sqlalchemy.orm import Session from cueweaver.adapters.output import AtomicOutputPublisher from cueweaver.application.database import SqliteDatabase @@ -238,6 +240,116 @@ def test_sqlite_record_store_persists_records_and_uses_a_transactional_database( assert store.load() == [] +def test_sqlite_record_store_load_uses_three_selects_for_many_jobs(tmp_path: Path): + database = SqliteDatabase(tmp_path / "cueweaver.sqlite3") + store = SqliteJobRecordStore(database) + records = [] + for index in range(4): + record = persisted_job_record(f"sqlite-job-{index}") + record["queue_sequence"] = 4 - index + timestamp = record["created_at"] + assert isinstance(timestamp, str) + record["status_history"] = [ + { + "status": "Queued", + "attempt": 1, + "started_at": timestamp, + "finished_at": timestamp, + }, + { + "status": "Failed", + "attempt": 1, + "started_at": timestamp, + "finished_at": timestamp, + }, + ] + request = record["request"] + assert isinstance(request, dict) + request["term_map_mode"] = "selected" + request["term_map"] = { + "id": f"map-{index}", + "name": f"Map {index}", + "content": {"Captain": "队长", "Doctor": "医生"}, + } + records.append(record) + store.write(record) + + select_count = 0 + + def count_selects( + execute_state, + ) -> None: + nonlocal select_count + if execute_state.is_select: + select_count += 1 + + event.listen(Session, "do_orm_execute", count_selects) + try: + loaded = store.load() + finally: + event.remove(Session, "do_orm_execute", count_selects) + + assert select_count == 3 + assert [record["id"] for record in loaded] == [ + "sqlite-job-3", + "sqlite-job-2", + "sqlite-job-1", + "sqlite-job-0", + ] + assert loaded == sorted(records, key=lambda record: record["queue_sequence"]) + + +def test_sqlite_job_term_map_snapshot_is_immutable_after_creation(tmp_path: Path): + store = SqliteJobRecordStore(SqliteDatabase(tmp_path / "cueweaver.sqlite3")) + original = persisted_job_record("snapshot-job") + original_request = original["request"] + assert isinstance(original_request, dict) + original_request["term_map_mode"] = "selected" + original_request["term_map"] = { + "id": "map-original", + "name": "Original", + "content": {"Captain": "队长"}, + } + store.write(original) + + updated = copy_job_record(original) + updated_request = updated["request"] + assert isinstance(updated_request, dict) + updated_request["term_map"] = { + "id": "map-replaced", + "name": "Replaced", + "content": {"Captain": "舰长", "Doctor": "医生"}, + } + updated["error"] = {"code": "translation_failed", "message": "Failed"} + store.write(updated) + + loaded = store.load()[0] + assert loaded["request"]["term_map"] == { + "id": "map-original", + "name": "Original", + "content": {"Captain": "队长"}, + } + + +def test_existing_job_without_snapshot_cannot_acquire_one(tmp_path: Path): + store = SqliteJobRecordStore(SqliteDatabase(tmp_path / "cueweaver.sqlite3")) + original = persisted_job_record("no-snapshot-job") + store.write(original) + + updated = copy_job_record(original) + request = updated["request"] + assert isinstance(request, dict) + request["term_map_mode"] = "follow" + request["term_map"] = { + "id": "map-later", + "name": "Later", + "content": {"Captain": "队长"}, + } + store.write(updated) + + assert store.load()[0]["request"]["term_map"] is None + + def test_sqlite_job_write_rolls_back_related_rows_on_database_failure(tmp_path: Path): database_path = tmp_path / "cueweaver.sqlite3" database = SqliteDatabase(database_path) @@ -678,6 +790,7 @@ def create_term_map_job(client: TestClient): "target_language_code": "zh-Hans", "term_map_mode": "selected", "term_map_id": term_map["id"], + "output_conflict_policy": "append-number", }, ).json() return term_map, queued @@ -803,6 +916,24 @@ def test_job_snapshot_is_independent_and_job_children_cascade(tmp_path: Path): ).fetchone() == (0,) +def test_job_term_map_snapshot_survives_failure_and_retry(tmp_path: Path): + media_root, work_root, _media, _subtitle = make_roots(tmp_path) + translator = FakeTranslator(error=RuntimeError("boom")) + + with make_client(media_root, work_root, translator) as client: + _term_map, queued = create_term_map_job(client) + wait_for_status(client, queued["id"], "Failed") + before = sqlite_job_record(work_root, queued["id"]) + + translator.error = None + response = client.post(f"/api/jobs/{queued['id']}/retry") + assert response.status_code == 200 + wait_for_status(client, queued["id"], "Completed") + after = sqlite_job_record(work_root, queued["id"]) + + assert after["request"]["term_map"] == before["request"]["term_map"] + + def test_none_job_mode_is_accepted_without_a_term_map(tmp_path: Path): media_root, work_root, _media, _subtitle = make_roots(tmp_path) @@ -936,11 +1067,14 @@ def assert_failed_record_persisted(jobs: Jobs, work_root: Path, job_id: str) -> def persisted_external_job( - tmp_path: Path, + tmp_path: Path, *, with_term_map: bool = False ) -> tuple[Path, Path, dict[str, object], Path, dict[str, object]]: media_root, work_root, _media, _subtitle = make_roots(tmp_path) with make_client(media_root, work_root, FakeTranslator()) as client: - queued = create_job(client).json() + if with_term_map: + _term_map, queued = create_term_map_job(client) + else: + queued = create_job(client).json() wait_for_status(client, queued["id"], "Completed") record_path = work_root / "cueweaver.sqlite3" record = sqlite_job_record(work_root, str(queued["id"])) @@ -1803,12 +1937,10 @@ def test_cancel_persistence_failure_keeps_job_queued_and_persisted( job_id = str(queued["id"]) original_write = Jobs._write_record - def fail_cancel_write( - instance: Jobs, record_id: str, record: dict[str, object] - ) -> None: - if record_id == job_id and record["status"] == "Cancelled": + def fail_cancel_write(instance: Jobs, record: dict[str, object]) -> None: + if record["id"] == job_id and record["status"] == "Cancelled": raise OSError("record unavailable") - original_write(instance, record_id, record) + original_write(instance, record) monkeypatch.setattr(Jobs, "_write_record", fail_cancel_write) @@ -1829,10 +1961,10 @@ def test_retry_persistence_failure_keeps_job_terminal_and_unqueued( ) original_write = Jobs._write_record - def fail_retry_write(jobs: Jobs, job_id: str, record: dict[str, object]) -> None: + def fail_retry_write(jobs: Jobs, record: dict[str, object]) -> None: if record["status"] == "Queued": raise OSError("record unavailable") - original_write(jobs, job_id, record) + original_write(jobs, record) monkeypatch.setattr(Jobs, "_write_record", fail_retry_write) @@ -1968,14 +2100,12 @@ def test_embedded_phase_persistence_failure_marks_worker_failed( original_write = Jobs._write_record failed_once = False - def fail_translating_write( - current_jobs: Jobs, job_id: str, record: dict[str, object] - ) -> None: + def fail_translating_write(current_jobs: Jobs, record: dict[str, object]) -> None: nonlocal failed_once if not failed_once and record["status"] == "Translating": failed_once = True raise OSError("record unavailable") - original_write(current_jobs, job_id, record) + original_write(current_jobs, record) monkeypatch.setattr(Jobs, "_write_record", fail_translating_write) queued = jobs.create( @@ -2598,12 +2728,12 @@ def test_worker_survives_a_persistence_failure_and_processes_next_job( original_write = Jobs._write_record failed_once = False - def fail_translating_write(jobs, job_id, record): + def fail_translating_write(jobs, record): nonlocal failed_once if not failed_once and record["status"] == "Translating": failed_once = True raise OSError("record unavailable") - original_write(jobs, job_id, record) + original_write(jobs, record) monkeypatch.setattr(Jobs, "_write_record", fail_translating_write) first = create_job(client, "zh").json() @@ -3254,15 +3384,9 @@ def test_restart_recovers_queued_jobs_and_interrupts_running_jobs( tmp_path: Path, active_status: str ): media_root, work_root, queued, _record_path, record = persisted_external_job( - tmp_path + tmp_path, with_term_map=True ) set_record_status(record, active_status, finished_at=None) - record["request"]["term_map"] = { - "id": "map-1", - "name": "Characters", - "content": {"Captain": "队长"}, - } - record["request"]["term_map_mode"] = "selected" record["finished_at"] = None persist_job_record(work_root, record) work_directory = work_root / "jobs" / queued["id"]