from __future__ import annotations import inspect import json from collections.abc import Iterator from contextlib import contextmanager from decimal import Decimal from typing import Any import pytest from app.modules.data_process.store import ( ConflictError, DataProcessStore, DataProcessStoreError, InvalidStateError, _decode_row, _preview_config_changed, _reasoning_output_is_valid, _source_storage_descriptor, repeat_task_id, ) def test_preview_replace_sql_never_uses_untyped_null_placeholders() -> None: source = inspect.getsource(DataProcessStore.replace_preview_items) assert "%s IS NULL" not in source assert "is_direct_build = preview_run_id is None" in source assert "workflow_step=CASE WHEN %s THEN 'preview'" in source class _Result: def __init__(self, *, row: dict[str, Any] | None = None, rows: list[dict[str, Any]] | None = None): self.row = row self.rows = rows or [] def fetchone(self) -> dict[str, Any] | None: return self.row def fetchall(self) -> list[dict[str, Any]]: return self.rows class _PublishConnection: def __init__( self, results: list[dict[str, Any]], datasets: list[dict[str, Any]] | None = None, ): self.results = results self.datasets: list[dict[str, Any]] = datasets or [] self.files: list[dict[str, Any]] = [] self.records: list[dict[str, Any]] = [] def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if params is not None: placeholder_count = normalized.count("%s") assert placeholder_count == len(params), ( f"SQL placeholder count {placeholder_count} does not match " f"parameter count {len(params)}" ) if normalized.startswith("SELECT * FROM data_process_results"): return _Result(rows=self.results) if normalized.startswith("SELECT * FROM datasets WHERE source='task'"): assert "source_task_id=%s" in normalized assert "source_task_id IS NULL AND task_id=%s" in normalized assert "deleted_at IS NULL" in normalized source_task_id, legacy_task_id = params return _Result( rows=[ item for item in self.datasets if item.get("source") == "task" and item.get("deleted_at") is None and ( item.get("source_task_id") == source_task_id or ( item.get("source_task_id") is None and item.get("task_id") == legacy_task_id ) ) ] ) if normalized.startswith("INSERT INTO datasets"): dataset = { "id": params[0], "name": params[1], "type": params[2], "source": "task", "task_id": params[4], "source_task_id": params[5], "count": params[8], "record_count": params[9], "metadata": params[11], "deleted_at": None, } self.datasets.append(dataset) return _Result(row=dataset) if normalized.startswith("UPDATE datasets SET name="): dataset = next(item for item in self.datasets if item["id"] == params[10]) dataset.update( { "name": params[0], "type": params[1], "count": params[5], "record_count": params[6], "metadata": params[8], } ) return _Result(row=dataset) if normalized.startswith("DELETE FROM dataset_records WHERE dataset_id"): self.records = [item for item in self.records if item["dataset_id"] != params[0]] if normalized.startswith("DELETE FROM dataset_files WHERE dataset_id"): self.files = [item for item in self.files if item["dataset_id"] != params[0]] if normalized.startswith("INSERT INTO dataset_files"): self.files.append( { "id": params[0], "dataset_id": params[1], "name": params[2], "record_count": params[11], } ) if normalized.startswith("INSERT INTO dataset_records"): self.records.append( { "dataset_id": params[1], "line_no": params[4], "split": params[5], "output": params[8], "raw": json.loads(params[9]), } ) return _Result() class _PublishStore(DataProcessStore): def __init__( self, conn: _PublishConnection, task_config: dict[str, Any] | None = None, ): self._conn = conn self._task_config = task_config or {} @contextmanager def connect(self) -> Iterator[_PublishConnection]: yield self._conn def _task_in_connection(self, conn: Any, task_id: str, *, for_update: bool = False) -> dict[str, Any]: train_dataset = next( (item for item in self._conn.datasets if item["type"] == "train"), None ) return { "id": task_id, "status": "completed", "results_confirmed": True, "description": "", "config": self._task_config, "output_dataset_id": train_dataset and train_dataset["id"], } @staticmethod def _source_ids(conn: Any, task_id: str) -> list[dict[str, Any]]: return [] class _RegenerationConnection: def __init__( self, task: dict[str, Any], datasets: list[dict[str, Any]] | None = None, ) -> None: self.task = task self.datasets = datasets or [ { "id": "dataset_train", "source": "task", "task_id": task["id"], "source_task_id": task["id"], "deleted_at": None, }, { "id": "dataset_validation", "source": "task", "task_id": task["id"], "source_task_id": task["id"], "deleted_at": None, }, { "id": "dataset_test", "source": "task", "task_id": task["id"], "source_task_id": task["id"], "deleted_at": None, }, ] self.sources = [{"id": "source_1"}] self.previews = [{"id": "preview_1"}] self.results = [{"id": "result_1"}] def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if params is not None: assert normalized.count("%s") == len(params) if normalized.startswith("UPDATE datasets SET source_task_id="): source_task_id, _, legacy_task_id = params for dataset in self.datasets: if ( dataset.get("source") == "task" and dataset.get("source_task_id") is None and dataset.get("task_id") == legacy_task_id and dataset.get("deleted_at") is None ): dataset["source_task_id"] = source_task_id return _Result() if normalized.startswith("SELECT EXISTS("): assert "source_task_id=%s" in normalized assert "source_task_id IS NULL AND task_id=%s" in normalized assert "deleted_at IS NULL" in normalized source_task_id, legacy_task_id = params exists = any( dataset.get("source") == "task" and dataset.get("deleted_at") is None and ( dataset.get("source_task_id") == source_task_id or ( dataset.get("source_task_id") is None and dataset.get("task_id") == legacy_task_id ) ) for dataset in self.datasets ) return _Result(row={"exists": exists}) if normalized.startswith("DELETE FROM data_process_results"): self.results.clear() return _Result() if normalized.startswith("DELETE FROM data_process_preview_items"): self.previews.clear() return _Result() if normalized.startswith("SELECT COUNT(*) AS count FROM data_process_preview_items"): return _Result(row={"count": len(self.previews)}) if normalized.startswith("UPDATE data_process_tasks SET name="): self.task.update( { "name": params[0], "description": params[1], "config": params[2], "updated_at": params[3], } ) return _Result(row=dict(self.task)) raise AssertionError(f"unexpected SQL: {normalized}") class _RegenerationStore(DataProcessStore): def __init__(self, conn: _RegenerationConnection): self._conn = conn @contextmanager def connect(self) -> Iterator[_RegenerationConnection]: yield self._conn def _task_in_connection( self, conn: Any, task_id: str, *, for_update: bool = False ) -> dict[str, Any]: assert for_update is True assert task_id == self._conn.task["id"] return dict(self._conn.task) class _TaskDetailConnection: def __init__( self, task: dict[str, Any], datasets: list[dict[str, Any]], ) -> None: self.task = task self.datasets = datasets def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) assert normalized.startswith("SELECT task.*") assert "dataset.source_task_id=task.id" in normalized assert "dataset.source_task_id IS NULL AND dataset.task_id=task.id" in normalized assert "dataset.deleted_at IS NULL" in normalized assert "dataset.metadata::jsonb" not in normalized assert "WHEN 'train' THEN 'train'" in normalized assert "WHEN 'val' THEN 'validation'" in normalized assert "WHEN 'test' THEN 'test'" in normalized task_id = params[0] visible = [ dataset for dataset in self.datasets if dataset.get("source") == "task" and dataset.get("deleted_at") is None and ( dataset.get("source_task_id") == task_id or ( dataset.get("source_task_id") is None and dataset.get("task_id") == task_id ) ) ] return _Result( row={ **self.task, "output_datasets": json.dumps( [ { "id": item["id"], "name": item["name"], "type": item["type"], "count": item["count"], "dataset_split": item["dataset_split"], } for item in visible ] ), } ) class _TaskDetailStore(DataProcessStore): def __init__(self, conn: _TaskDetailConnection) -> None: self._conn = conn @contextmanager def connect(self) -> Iterator[_TaskDetailConnection]: yield self._conn class _RepeatConnection: def __init__(self) -> None: self.source_files = [ { "id": "source-old", "name": "source.jsonl", "size_bytes": 12, "record_count": 1, "file_format": "jsonl", "checksum_sha256": "a" * 64, "content": '{"id":1}\n', "content_preview": '{"id":1}', "metadata": {"storage_backend": "local"}, "created_by": "user-1", } ] self.source_previews = [ { "id": "preview-old", "source_file_id": "source-old", "original_content": '{"id":1}', "edited_content": '{"id":1,"checked":true}', "source_start": 0, "source_end": 8, "source_start_line": 1, "source_end_line": 1, "token_count": 5, "status": "modified", "quality_score": {"overall": 90}, } ] self.created_task: dict[str, Any] | None = None self.created_files: list[dict[str, Any]] = [] self.created_previews: list[dict[str, Any]] = [] def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if params is not None: assert normalized.count("%s") == len(params) if normalized.startswith("SELECT * FROM data_process_tasks WHERE id="): return _Result(row=None) if normalized.startswith("SELECT * FROM data_process_source_files"): return _Result(rows=[dict(item) for item in self.source_files]) if normalized.startswith("SELECT * FROM data_process_preview_items"): return _Result(rows=[dict(item) for item in self.source_previews]) if normalized.startswith("INSERT INTO data_process_tasks"): self.created_task = { "id": params[0], "name": params[1], "description": params[2], "status": "pending", "process_type": params[3], "source_dataset_id": params[4], "config": params[5], "progress": 20, "input_count": params[6], "results_confirmed": False, "workflow_step": "preview", "preview_status": "completed", "preview_progress": 100, "preview_total_files": params[7], "preview_completed_files": params[8], "created_at": params[15], "updated_at": params[16], } return _Result(row=dict(self.created_task)) if normalized.startswith("INSERT INTO data_process_source_files"): self.created_files.append( { "id": params[0], "task_id": params[1], "storage_object_id": params[2], "content": params[8], } ) return _Result() if normalized.startswith("INSERT INTO data_process_preview_items"): self.created_previews.append( { "id": params[0], "task_id": params[1], "source_file_id": params[2], "edited_content": params[4], } ) return _Result() if normalized.startswith("SELECT (SELECT COUNT(*) FROM data_process_source_files"): return _Result( row={ "source_file_count": len(self.created_files), "preview_count": len(self.created_previews), } ) raise AssertionError(f"unexpected SQL: {normalized}") class _RepeatStore(DataProcessStore): def __init__(self, conn: _RepeatConnection) -> None: self._conn = conn @contextmanager def connect(self) -> Iterator[_RepeatConnection]: yield self._conn def _task_in_connection( self, conn: Any, task_id: str, *, for_update: bool = False, ) -> dict[str, Any]: assert task_id == "task-source" assert for_update is True return { "id": task_id, "name": "原任务", "description": "原描述", "status": "completed", "process_type": "structured", "source_dataset_id": None, "config": { "temperature": 0.3, "_regeneration_prepared": {"prepared": True}, }, "results_confirmed": True, "preview_status": "completed", "tenant_id": "tenant-1", "project_id": "project-1", "owner_id": "owner-1", "created_by": "user-1", "updated_at": "2026-07-28T12:00:00Z", } class _TaskListConnection: def __init__(self) -> None: self.task = { **_regeneration_task( status="pending", output_dataset_id=None, output_count=17, ), "created_at": "2026-07-25T18:00:00Z", "deleted_at": None, } self.sources = [ {"task_id": "task-1", "deleted_at": None}, {"task_id": "task-1", "deleted_at": None}, {"task_id": "task-1", "deleted_at": "2026-07-26T00:00:00Z"}, ] def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if normalized.startswith("SELECT COUNT(*) AS count FROM data_process_tasks task"): assert "task.deleted_at IS NULL" in normalized return _Result(row={"count": 1}) assert normalized.startswith("SELECT task.*") assert "source_file.task_id=task.id" in normalized assert "source_file.deleted_at IS NULL" in normalized source_file_count = sum( item["task_id"] == self.task["id"] and item["deleted_at"] is None for item in self.sources ) return _Result( rows=[ { **self.task, "source_file_count": source_file_count, } ] ) class _TaskListStore(DataProcessStore): def __init__(self, conn: _TaskListConnection) -> None: self._conn = conn @contextmanager def connect(self) -> Iterator[_TaskListConnection]: yield self._conn class _LegacyRecoveryConnection: def __init__(self) -> None: self.task = { **_regeneration_task( status="pending", progress=20, output_dataset_id=None, output_count=0, generation_run_id=None, started_at=None, completed_at=None, ), } self.datasets = [ {"id": "dataset_train", "type": "train", "count": 1, "created_at": "2026-07-25T18:12:00Z"}, {"id": "dataset_test", "type": "test", "count": 1, "created_at": "2026-07-25T18:12:00Z"}, ] self.records = [ { "id": "record_train", "dataset_id": "dataset_train", "line_no": 1, "split": "train", "instruction": "训练问题", "input": "", "output": "训练答案", "raw": json.dumps( { "source_result_id": "result_train", "preview_item_id": "preview_1", } ), "status": "valid", "source_result_id": None, "preview_item_id": None, "created_at": "2026-07-25T18:12:00Z", }, { "id": "record_test", "dataset_id": "dataset_test", "line_no": 1, "split": "test", "instruction": "测试问题", "input": "输入", "output": "测试答案", "raw": json.dumps({"source_result_id": "result_test"}), "status": "modified", "source_result_id": None, "preview_item_id": None, "created_at": "2026-07-25T18:12:00Z", }, ] self.previews = [{"id": "preview_1"}] self.results: list[dict[str, Any]] = [] def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if normalized.startswith("SELECT COUNT(*) AS count FROM data_process_results"): return _Result(row={"count": len(self.results)}) if normalized.startswith("SELECT id, type, count, created_at FROM datasets"): return _Result(rows=list(self.datasets)) if normalized.startswith("SELECT id, dataset_id, line_no, split"): return _Result(rows=list(self.records)) if normalized.startswith("SELECT id FROM data_process_preview_items"): return _Result(rows=list(self.previews)) if normalized.startswith("INSERT INTO data_process_results"): self.results.append( { "id": params[0], "task_id": params[1], "preview_item_id": params[2], "instruction": params[3], "input": params[4], "output": params[5], "status": params[9], "split": params[10], } ) return _Result() if normalized.startswith("UPDATE dataset_records SET source_result_id="): record = next(item for item in self.records if item["id"] == params[2]) record["source_result_id"] = params[0] record["preview_item_id"] = params[1] return _Result() if normalized.startswith("UPDATE data_process_tasks SET status='completed'"): assert "workflow_step='results'" in normalized self.task.update( { "status": "completed", "progress": 100, "output_dataset_id": params[0], "output_count": params[1], "workflow_step": "results", } ) return _Result() raise AssertionError(f"unexpected SQL: {normalized}") class _LegacyRecoveryStore(DataProcessStore): def __init__(self, conn: _LegacyRecoveryConnection) -> None: self._conn = conn @contextmanager def connect(self) -> Iterator[_LegacyRecoveryConnection]: yield self._conn def _task_in_connection( self, conn: Any, task_id: str, *, for_update: bool = False ) -> dict[str, Any]: assert task_id == "task-1" assert for_update is True return dict(self._conn.task) class _StartGenerationConnection: def __init__(self, *, published_prepared: bool = False, preview_count: int = 1) -> None: config = {"chunk_method": "fixed", "temperature": 0.7} if published_prepared: config["_regeneration_prepared"] = { "prepared": True, "preview_invalidated": False, } self.task = { **_regeneration_task( status="completed" if published_prepared else "pending", config=config, output_dataset_id="dataset_train" if published_prepared else None, output_count=28, results_confirmed=published_prepared, ), "generation_run_id": None, } self.preview_count = preview_count self.results = [{"id": "old-result"}] def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if normalized.startswith("SELECT COUNT(*) AS count FROM data_process_preview_items"): return _Result(row={"count": self.preview_count}) if normalized.startswith("DELETE FROM data_process_results"): self.results.clear() return _Result() assert normalized.startswith("UPDATE data_process_tasks SET config=%s, status='running'") assert "output_dataset_id=NULL" in normalized assert "output_count=0" in normalized assert "results_confirmed=FALSE" in normalized self.task.update( { "config": params[0], "status": "running", "progress": 30, "output_count": 0, "output_dataset_id": None, "failure_reason": None, "started_at": params[1], "completed_at": None, "filtered_count": 0, "duplicate_count": 0, "error_count": 0, "generation_run_id": params[2], "results_confirmed": False, "updated_at": params[3], } ) return _Result(row=dict(self.task)) class _StartGenerationStore(DataProcessStore): def __init__(self, conn: _StartGenerationConnection) -> None: self._conn = conn @contextmanager def connect(self) -> Iterator[_StartGenerationConnection]: yield self._conn def _task_in_connection( self, conn: Any, task_id: str, *, for_update: bool = False ) -> dict[str, Any]: assert task_id == "task-1" assert for_update is True return dict(self._conn.task) def test_decode_row_serializes_postgres_numeric_values_as_json_numbers() -> None: decoded = _decode_row( { "progress": Decimal("100.00"), "duration_seconds": Decimal("389.000000"), } ) assert decoded == {"progress": 100.0, "duration_seconds": 389.0} def test_repeat_task_copies_business_snapshot_with_new_resource_ids() -> None: conn = _RepeatConnection() store = _RepeatStore(conn) request_id = "repeat-request-0001" target_task_id = repeat_task_id("task-source", request_id) repeated = store.repeat_task( "task-source", expected_updated_at="2026-07-28T12:00:00Z", request_id=request_id, file_copies={ "source-old": { "id": "source-new", "storage_object_id": ( f"local://data-process/{target_task_id}/source-new/v1/source.jsonl" ), } }, ) assert repeated["created"] is True assert repeated["task"]["id"] == target_task_id assert repeated["task"]["config"] == {"temperature": 0.3} assert repeated["task"]["results_confirmed"] is False assert repeated["copied_source_file_count"] == 1 assert repeated["copied_preview_count"] == 1 assert conn.created_files == [ { "id": "source-new", "task_id": target_task_id, "storage_object_id": ( f"local://data-process/{target_task_id}/source-new/v1/source.jsonl" ), "content": '{"id":1}\n', } ] assert conn.created_previews[0]["task_id"] == target_task_id assert conn.created_previews[0]["source_file_id"] == "source-new" assert conn.created_previews[0]["edited_content"] == '{"id":1,"checked":true}' def test_decode_row_decodes_aggregated_output_datasets_json() -> None: decoded = _decode_row( { "id": "task-1", "output_datasets": '[{"id":"dataset_train","type":"train"}]', } ) assert decoded == { "id": "task-1", "output_datasets": [{"id": "dataset_train", "type": "train"}], } @pytest.mark.parametrize( ("process_type", "current", "next_config", "expected"), [ ( "structured", {"preprocess_options": ["deduplicate"]}, {"preprocess_options": ["deduplicate"], "temperature": 0.2}, False, ), ( "structured", {"preprocess_options": ["deduplicate"]}, {"preprocess_options": ["clean_invalid"]}, True, ), ( "structured", {"preprocess_options": ["a", "b"]}, {"preprocessOptions": ["b", "a", "a"], "chunk_size": 2048}, False, ), ( "unstructured", {"chunk_method": "fixed"}, {"chunk_method": "fixed", "generation_prompt": "new"}, False, ), ( "unstructured", {"chunk_method": "fixed"}, {"chunk_method": "semantic"}, True, ), ( "unstructured", {"chunk_size": 800, "chunk_overlap": 100}, {"chunk_size": 900, "chunk_overlap": 100}, True, ), ( "unstructured", {"min_chunk_size": 100}, {"min_chunk_size": 120}, True, ), ( "unstructured", {"semantic_breakpoint_percentile": 95}, {"semantic_breakpoint_percentile": 90}, True, ), ( "unstructured", {}, { "preserve_tables": True, "preserve_code_blocks": True, "preserve_lists": True, }, False, ), ( "unstructured", {"preserve_tables": False}, {"preserve_tables": True}, True, ), ( "unstructured", {"preserve_code_blocks": False}, {"preserve_code_blocks": True}, True, ), ( "unstructured", {"preserve_lists": False}, {"preserve_lists": True}, True, ), ( "unstructured", {"preprocess_options": ["deduplicate"]}, {"preprocess_options": ["clean_invalid"]}, True, ), ], ) def test_regeneration_preview_invalidation_matrix( process_type: str, current: dict[str, Any], next_config: dict[str, Any], expected: bool, ) -> None: assert _preview_config_changed(process_type, current, next_config) is expected def _regeneration_task(**updates: Any) -> dict[str, Any]: task = { "id": "task-1", "name": "原任务", "description": "", "process_type": "unstructured", "config": {"chunk_method": "fixed", "temperature": 0.7}, "status": "completed", "progress": 100, "output_dataset_id": "dataset_train", "output_count": 28, "filtered_count": 1, "duplicate_count": 1, "error_count": 0, "failure_reason": None, "generation_run_id": None, "started_at": "2026-07-25T18:00:00Z", "completed_at": "2026-07-25T18:05:00Z", "updated_at": "2026-07-25T18:05:00Z", } task.update(updates) return task def test_list_tasks_exposes_source_and_generation_counts() -> None: page = _TaskListStore(_TaskListConnection()).list_tasks(page=1, page_size=20) assert page["total"] == 1 assert page["page"] == 1 assert page["page_size"] == 20 item = page["items"][0] assert item["status"] == "pending" assert item["source_file_count"] == 2 assert item["output_count"] == 17 assert item["output_dataset_id"] is None def test_start_generation_clears_previous_output_count() -> None: conn = _StartGenerationConnection() task = _StartGenerationStore(conn).start_generation("task-1") assert task["status"] == "running" assert task["output_count"] == 0 assert task["results_confirmed"] is False assert conn.results == [] def test_prepared_published_task_is_only_cleared_when_generation_starts() -> None: conn = _StartGenerationConnection(published_prepared=True) task = _StartGenerationStore(conn).start_generation("task-1") assert task["status"] == "running" assert task["output_dataset_id"] is None assert task["output_count"] == 0 assert "_regeneration_prepared" not in task["config"] assert conn.results == [] def test_prepared_published_task_survives_generation_preflight_failure() -> None: conn = _StartGenerationConnection(published_prepared=True, preview_count=0) with pytest.raises(InvalidStateError, match="preview must be built"): _StartGenerationStore(conn).start_generation("task-1") assert conn.task["status"] == "completed" assert conn.task["output_dataset_id"] == "dataset_train" assert conn.task["output_count"] == 28 assert conn.task["results_confirmed"] is True assert "_regeneration_prepared" in conn.task["config"] assert conn.results == [{"id": "old-result"}] def test_legacy_aborted_regeneration_recovers_results_and_published_state() -> None: conn = _LegacyRecoveryConnection() store = _LegacyRecoveryStore(conn) recovered = store.recover_legacy_aborted_regeneration("task-1") assert recovered == {"recovered": True, "result_count": 2} assert conn.task["status"] == "completed" assert conn.task["progress"] == 100 assert conn.task["output_dataset_id"] == "dataset_train" assert conn.task["output_count"] == 2 assert conn.task["workflow_step"] == "results" assert conn.task["started_at"] is None assert conn.task["completed_at"] is None assert [item["id"] for item in conn.results] == ["result_train", "result_test"] assert conn.results[0]["preview_item_id"] == "preview_1" assert conn.results[1]["preview_item_id"] is None assert conn.records[0]["source_result_id"] == "result_train" assert conn.records[0]["preview_item_id"] == "preview_1" repeated = store.recover_legacy_aborted_regeneration("task-1") assert repeated == {"recovered": False, "result_count": 0} assert len(conn.results) == 2 def test_normal_pending_task_is_not_mistaken_for_legacy_regeneration() -> None: conn = _LegacyRecoveryConnection() conn.datasets = [] conn.records = [] result = _LegacyRecoveryStore(conn).recover_legacy_aborted_regeneration("task-1") assert result == {"recovered": False, "result_count": 0} assert conn.task["status"] == "pending" assert conn.results == [] def _legacy_published_datasets(task_id: str = "task-1") -> list[dict[str, Any]]: specs = ( ("dataset_train", "制度问答-训练集", "train", "train", 22), ("dataset_validation", "制度问答-验证集", "val", "validation", 3), ("dataset_test", "制度问答-测试集", "test", "test", 3), ) dataset_ids = {split: dataset_id for dataset_id, _, _, split, _ in specs} return [ { "id": dataset_id, "name": name, "type": dataset_type, "source": "task", "task_id": task_id, "source_task_id": None, "count": count, "record_count": count, "dataset_split": split, "metadata": json.dumps( { "base_dataset_name": "制度问答", "dataset_split": split, "split_dataset_ids": dataset_ids, } ), "deleted_at": None, } for dataset_id, name, dataset_type, split, count in specs ] def test_prepare_regeneration_preserves_outputs_sources_and_generation_only_preview() -> None: conn = _RegenerationConnection(_regeneration_task()) original_datasets = list(conn.datasets) original_sources = list(conn.sources) result = _RegenerationStore(conn).prepare_regeneration( "task-1", { "name": "新任务名", "description": "更换生成参数", "process_type": "unstructured", "config": {"chunk_method": "fixed", "temperature": 0.2}, "expected_updated_at": "2026-07-25T18:05:00Z", }, ) assert result["preview_invalidated"] is False assert result["published_outputs_preserved"] is True assert result["task"]["output_dataset_id"] == "dataset_train" assert result["task"]["status"] == "completed" assert result["task"]["progress"] == 100 assert result["task"]["output_count"] == 28 assert result["task"]["started_at"] == "2026-07-25T18:00:00Z" assert result["task"]["completed_at"] == "2026-07-25T18:05:00Z" assert "_regeneration_prepared" not in result["task"]["config"] stored_config = json.loads(conn.task["config"]) assert stored_config["_regeneration_prepared"]["prepared"] is True assert conn.results == [{"id": "result_1"}] assert conn.previews == [{"id": "preview_1"}] assert conn.datasets == original_datasets assert conn.sources == original_sources def test_prepare_regeneration_backfills_and_keeps_legacy_task_datasets_visible() -> None: legacy_datasets = _legacy_published_datasets() # 历史 metadata 可能不是合法 JSON,详情查询不能再依赖 metadata::jsonb。 legacy_datasets[0]["metadata"] = "{legacy-invalid-json" deleted_dataset = { **legacy_datasets[0], "id": "dataset_deleted", "name": "已删除训练集", "deleted_at": "2026-07-25T20:00:00Z", } unrelated_dataset = { **legacy_datasets[0], "id": "dataset_unrelated", "name": "其他任务训练集", "task_id": "task-other", } conn = _RegenerationConnection( _regeneration_task(), [*legacy_datasets, deleted_dataset, unrelated_dataset], ) before = _TaskDetailStore(_TaskDetailConnection(conn.task, conn.datasets)).get_task( "task-1" ) assert [item["id"] for item in before["output_datasets"]] == [ "dataset_train", "dataset_validation", "dataset_test", ] result = _RegenerationStore(conn).prepare_regeneration( "task-1", { "name": "原任务", "description": "", "process_type": "unstructured", "config": {"chunk_method": "fixed", "temperature": 0.2}, "expected_updated_at": "2026-07-25T18:05:00Z", }, ) assert result["published_outputs_preserved"] is True assert result["task"]["output_dataset_id"] == "dataset_train" assert len(conn.datasets) == 5 assert all( item["source_task_id"] == "task-1" for item in conn.datasets[:3] ) assert deleted_dataset["source_task_id"] is None assert unrelated_dataset["source_task_id"] is None after = _TaskDetailStore(_TaskDetailConnection(conn.task, conn.datasets)).get_task( "task-1" ) assert [item["id"] for item in after["output_datasets"]] == [ "dataset_train", "dataset_validation", "dataset_test", ] def test_prepare_regeneration_defers_preview_deletion_when_chunk_configuration_changes() -> None: conn = _RegenerationConnection(_regeneration_task(output_dataset_id=None)) result = _RegenerationStore(conn).prepare_regeneration( "task-1", { "name": "原任务", "description": "", "process_type": "unstructured", "config": {"chunk_method": "semantic", "temperature": 0.7}, "expected_updated_at": "2026-07-25T18:05:00Z", }, ) assert result["preview_invalidated"] is True assert result["published_outputs_preserved"] is True assert result["task"]["status"] == "completed" assert result["task"]["progress"] == 100 assert result["task"]["output_count"] == 28 assert conn.previews == [{"id": "preview_1"}] assert conn.results == [{"id": "result_1"}] assert len(conn.datasets) == 3 @pytest.mark.parametrize( ("task_updates", "payload_updates", "error_type", "message"), [ ( {"status": "running"}, {}, ConflictError, "running task cannot be prepared", ), ( {}, {"expected_updated_at": "2026-07-25T17:00:00Z"}, ConflictError, "modified by another request", ), ( {}, {"process_type": "structured"}, InvalidStateError, "process_type cannot be changed", ), ], ) def test_prepare_regeneration_rejects_running_stale_and_type_change_without_mutation( task_updates: dict[str, Any], payload_updates: dict[str, Any], error_type: type[Exception], message: str, ) -> None: conn = _RegenerationConnection(_regeneration_task(**task_updates)) payload = { "name": "原任务", "description": "", "process_type": "unstructured", "config": {"chunk_method": "fixed"}, "expected_updated_at": "2026-07-25T18:05:00Z", **payload_updates, } with pytest.raises(error_type, match=message): _RegenerationStore(conn).prepare_regeneration("task-1", payload) assert conn.results == [{"id": "result_1"}] assert conn.previews == [{"id": "preview_1"}] assert conn.task["output_dataset_id"] == "dataset_train" def test_publish_creates_three_independent_datasets_with_exact_counts() -> None: results = [ { "id": f"result-{index}", "status": "valid", "instruction": f"问题 {index}", "input": "", "output": ( "\n先读取制度条款。\n\n答案 0" if index == 0 else f"答案 {index}" ), "preview_item_id": f"preview-{index}", } for index in range(28) ] conn = _PublishConnection(results) published = _PublishStore(conn).publish( "task-1", { "dataset_name": "制度问答", "storage_type": "local", "format": "alpaca_jsonl", "split": {"train": 80, "validation": 10, "test": 10}, }, ) assert [(item["name"], item["type"], item["count"]) for item in conn.datasets] == [ ("制度问答-训练集", "train", 22), ("制度问答-验证集", "val", 3), ("制度问答-测试集", "test", 3), ] assert len(conn.files) == 3 assert {item["dataset_id"] for item in conn.files} == { item["id"] for item in conn.datasets } assert len(conn.records) == 28 reasoning_record = next( item for item in conn.records if item["output"].startswith("") ) assert reasoning_record["raw"]["output"] == reasoning_record["output"] assert published["dataset"]["type"] == "train" assert len(published["datasets"]) == 3 assert published["split_counts"] == {"train": 22, "validation": 3, "test": 3} original_ids = [item["id"] for item in conn.datasets] republished = _PublishStore(conn).publish( "task-1", { "dataset_name": "制度问答-训练集", "storage_type": "local", "format": "alpaca_jsonl", "split": {"train": 80, "validation": 10, "test": 10}, }, ) assert [item["id"] for item in conn.datasets] == original_ids assert len(conn.datasets) == 3 assert len(conn.files) == 3 assert len(conn.records) == 28 assert republished["created"] is False def test_publish_reuses_legacy_task_id_only_split_datasets() -> None: results = [ { "id": f"result-{index}", "status": "valid", "instruction": f"问题 {index}", "input": "", "output": f"答案 {index}", "preview_item_id": f"preview-{index}", } for index in range(10) ] legacy_datasets = _legacy_published_datasets() original_ids = [item["id"] for item in legacy_datasets] conn = _PublishConnection(results, legacy_datasets) published = _PublishStore(conn).publish( "task-1", { "dataset_name": "不会创建新数据集", "storage_type": "local", "format": "alpaca_jsonl", "split": {"train": 80, "validation": 10, "test": 10}, }, ) assert published["created"] is False assert [item["id"] for item in published["datasets"]] == original_ids assert [item["id"] for item in conn.datasets] == original_ids def test_publish_keeps_all_three_datasets_when_a_small_split_is_empty() -> None: conn = _PublishConnection( [ { "id": "result-only", "status": "valid", "instruction": "唯一问题", "input": "", "output": "唯一答案", "preview_item_id": "preview-only", } ] ) published = _PublishStore(conn).publish( "task-small", { "dataset_name": "小样本", "storage_type": "local", "format": "alpaca_jsonl", "split": {"train": 80, "validation": 10, "test": 10}, }, ) assert [(item["type"], item["count"]) for item in conn.datasets] == [ ("train", 1), ("val", 0), ("test", 0), ] assert len(published["datasets"]) == 3 assert len(conn.files) == 3 @pytest.mark.parametrize( ("output", "expected"), [ ("\n推理步骤\n\n最终答案", True), ("\n最终答案", False), ("只有推理", False), ("没有标签的最终答案", False), ("外层嵌套答案", False), ], ) def test_reasoning_output_validator_requires_one_complete_pair( output: str, expected: bool, ) -> None: assert _reasoning_output_is_valid(output) is expected def test_publish_rejects_invalid_reasoning_output_format() -> None: conn = _PublishConnection( [ { "id": "result-reasoning-invalid", "status": "valid", "instruction": "需要推理的问题", "input": "", "output": "只有最终答案", "preview_item_id": "preview-reasoning-invalid", } ] ) with pytest.raises(InvalidStateError, match="1 invalid results"): _PublishStore(conn, {"output_type": "reasoning"}).publish( "task-reasoning-invalid", { "dataset_name": "无效思维链", "storage_type": "local", "format": "alpaca_jsonl", "split": {"train": 80, "validation": 10, "test": 10}, }, ) def test_source_storage_descriptor_accepts_owned_local_and_legacy_db_references() -> None: task_id = "dpt_task" source_file_id = "dpsf_source" local_reference = ( f"local://data-process/{task_id}/{source_file_id}/v1/source%20100%25.csv" ) reference, metadata = _source_storage_descriptor( { "storage_object_id": local_reference, "metadata": {"storage_backend": "spoofed", "content_type": "text/csv"}, }, task_id, source_file_id, ) assert reference == local_reference assert metadata == {"storage_backend": "local", "content_type": "text/csv"} legacy_reference, legacy_metadata = _source_storage_descriptor( {"metadata": {"legacy": True}}, task_id, source_file_id, ) assert legacy_reference == f"db://data-process/{task_id}/{source_file_id}/v1" assert legacy_metadata == {"storage_backend": "database", "legacy": True} @pytest.mark.parametrize( "reference", [ "local://data-process/dpt_other/dpsf_source/v1/source.txt", "db://data-process/dpt_task/dpsf_other/v1", "/var/tmp/source.txt", ], ) def test_source_storage_descriptor_rejects_unowned_or_unsupported_references( reference: str, ) -> None: with pytest.raises(DataProcessStoreError): _source_storage_descriptor( {"storage_object_id": reference}, "dpt_task", "dpsf_source", ) class _LifecycleConnection: def __init__(self) -> None: self.task: dict[str, Any] = { "id": "task-lifecycle", "status": "running", "generation_run_id": "generation-active", "workflow_step": "generate", "preview_status": "running", "preview_progress": Decimal("40.00"), "preview_run_id": "preview-active", "preview_failure_reason": None, "preview_total_files": 5, "preview_completed_files": 2, "deleted_at": None, "deleted_by": None, } self.last_update_sql = "" def execute(self, sql: str, params: Any = None) -> _Result: normalized = " ".join(sql.split()) if params is not None: assert normalized.count("%s") == len(params) if normalized.startswith("SELECT * FROM data_process_tasks"): row = None if self.task["deleted_at"] is not None else dict(self.task) return _Result(row=row) if normalized.startswith("UPDATE data_process_tasks SET workflow_step="): self.last_update_sql = normalized workflow_step, updated_at, task_id = params assert task_id == self.task["id"] self.task.update(workflow_step=workflow_step, updated_at=updated_at) return _Result(row=dict(self.task)) if normalized.startswith("UPDATE data_process_tasks SET status=CASE"): self.last_update_sql = normalized deleted_at, deleted_by, updated_at, task_id = params assert task_id == self.task["id"] self.task.update( status="stopped", generation_run_id=None, preview_status="cancelled", preview_run_id=None, deleted_at=deleted_at, deleted_by=deleted_by, updated_at=updated_at, ) return _Result() if normalized.startswith("SELECT status, generation_run_id"): row = None if self.task["deleted_at"] is not None else dict(self.task) return _Result(row=row) if normalized.startswith("SELECT preview_status, preview_run_id"): row = None if self.task["deleted_at"] is not None else dict(self.task) return _Result(row=row) raise AssertionError(f"unexpected SQL: {normalized}") class _LifecycleStore(DataProcessStore): def __init__(self, conn: _LifecycleConnection) -> None: self._conn = conn @contextmanager def connect(self) -> Iterator[_LifecycleConnection]: yield self._conn def test_workflow_step_update_does_not_invalidate_active_runs() -> None: conn = _LifecycleConnection() task = _LifecycleStore(conn).update_workflow_step("task-lifecycle", "results") assert task["workflow_step"] == "results" assert task["status"] == "running" assert task["generation_run_id"] == "generation-active" assert task["preview_status"] == "running" assert task["preview_run_id"] == "preview-active" assert "generation_run_id" not in conn.last_update_sql assert "preview_run_id" not in conn.last_update_sql def test_delete_atomically_invalidates_generation_and_preview_runs() -> None: conn = _LifecycleConnection() store = _LifecycleStore(conn) store.delete_task("task-lifecycle", deleted_by="user-1") assert conn.task["status"] == "stopped" assert conn.task["generation_run_id"] is None assert conn.task["preview_status"] == "cancelled" assert conn.task["preview_run_id"] is None assert conn.task["deleted_by"] == "user-1" assert conn.task["deleted_at"] is not None assert store.generation_is_running("task-lifecycle", "generation-active") is False assert store.preview_is_running("task-lifecycle", "preview-active") is False