"""数据处理存储层 - 任务管理。""" from __future__ import annotations from typing import Any from collections.abc import Sequence import psycopg from .base import ( StoreBase, utcnow, new_id, repeat_task_id, json_dumps, _json_value, _decode_row, _public_task, _business_config, _preview_config_value, _preview_config_changed, _preview_config_projection, _normalized_preprocess_options, _regeneration_marker, _is_regeneration_prepared, _task_output_type, _task_reasoning_detail, _reasoning_output_is_valid, _dpo_fields_are_valid, _source_storage_descriptor, _serialize_value, NotFoundError, ConflictError, InvalidStateError, EDITABLE_STATUSES, ACTIVE_PREVIEW_STATUSES, WORKFLOW_STEPS, _REGENERATION_MARKER_KEY, _REPEAT_SOURCE_TASK_KEY, _REPEAT_REQUEST_KEY, _INTERNAL_CONFIG_KEYS, ) from ..algorithms import estimate_token_count class TasksMixin: """任务管理 Mixin。""" def list_tasks( self, *, page: int = 1, page_size: int = 20, keyword: str | None = None, status: str | None = None, process_type: str | None = None, tenant_id: str | None = None, project_id: str | None = None, ) -> dict[str, Any]: clauses = ["task.deleted_at IS NULL"] params: list[Any] = [] if keyword: clauses.append("(task.name ILIKE %s OR COALESCE(task.description, '') ILIKE %s)") pattern = f"%{keyword.strip()}%" params.extend([pattern, pattern]) if status: clauses.append("task.status = %s") params.append(status) if process_type: clauses.append("task.process_type = %s") params.append(process_type) if tenant_id: clauses.append("task.tenant_id = %s") params.append(tenant_id) if project_id: clauses.append("task.project_id = %s") params.append(project_id) where = " AND ".join(clauses) with self.connect() as conn: total = conn.execute( f"SELECT COUNT(*) AS count FROM data_process_tasks task WHERE {where}", params, ).fetchone()["count"] rows = conn.execute( f""" SELECT task.*, (SELECT COUNT(*) FROM data_process_source_files source_file WHERE source_file.task_id=task.id AND source_file.deleted_at IS NULL) AS source_file_count FROM data_process_tasks task WHERE {where} ORDER BY task.created_at DESC, task.id DESC LIMIT %s OFFSET %s """, [*params, page_size, (page - 1) * page_size], ).fetchall() return { "items": [_public_task(_decode_row(row)) for row in rows], "total": int(total), "page": page, "page_size": page_size, } def create_task(self, payload: dict[str, Any]) -> dict[str, Any]: task_id = new_id("dpt") now = utcnow() try: with self.connect() as conn: row = conn.execute( """ INSERT INTO data_process_tasks (id, name, description, status, process_type, source_dataset_id, config, progress, results_confirmed, tenant_id, project_id, owner_id, created_by, updated_by, created_at, updated_at) VALUES (%s, %s, %s, 'pending', %s, %s, %s, 0, FALSE, %s, %s, %s, %s, %s, %s, %s) RETURNING * """, ( task_id, payload["name"], payload.get("description") or "", payload["process_type"], payload.get("source_dataset_id"), json_dumps(_business_config(payload.get("config"))), payload.get("tenant_id"), payload.get("project_id"), payload.get("owner_id"), payload.get("created_by"), payload.get("created_by"), now, now, ), ).fetchone() except psycopg.errors.UniqueViolation as exc: raise ConflictError("data process task name already exists") from exc return _public_task(_decode_row(row)) or {} @staticmethod def _repeat_response( conn: psycopg.Connection[dict[str, Any]], row: dict[str, Any], *, source_task_id: str, created: bool, ) -> dict[str, Any]: task_id = str(row["id"]) counts = conn.execute( """ SELECT (SELECT COUNT(*) FROM data_process_source_files WHERE task_id=%s AND deleted_at IS NULL) AS source_file_count, (SELECT COUNT(*) FROM data_process_preview_items WHERE task_id=%s) AS preview_count """, (task_id, task_id), ).fetchone() or {} task = _public_task(_decode_row(row)) or {} task["source_file_count"] = int(counts.get("source_file_count") or 0) task["preview_count"] = int(counts.get("preview_count") or 0) return { "task": task, "source_task_id": source_task_id, "created": created, "copied_source_file_count": task["source_file_count"], "copied_preview_count": task["preview_count"], } def find_repeated_task( self, source_task_id: str, request_id: str, ) -> dict[str, Any] | None: """查找同一幂等请求已创建的新任务。""" task_id = repeat_task_id(source_task_id, request_id) with self.connect() as conn: row = conn.execute( "SELECT * FROM data_process_tasks WHERE id=%s", (task_id,), ).fetchone() if row is None: return None decoded = _decode_row(row) or {} config = decoded.get("config") or {} if ( config.get(_REPEAT_SOURCE_TASK_KEY) != source_task_id or config.get(_REPEAT_REQUEST_KEY) != request_id ): raise ConflictError("再次生成请求与现有任务冲突") if decoded.get("deleted_at"): raise ConflictError("此次再次生成创建的任务已被删除,请重新发起") return self._repeat_response( conn, row, source_task_id=source_task_id, created=False, ) def repeat_task( self, source_task_id: str, *, expected_updated_at: str, request_id: str, file_copies: dict[str, dict[str, str]], ) -> dict[str, Any]: """复制已确认任务的配置、源文件和预览,结果与发布数据保持独立。""" task_id = repeat_task_id(source_task_id, request_id) now = utcnow() try: with self.connect() as conn: existing = conn.execute( "SELECT * FROM data_process_tasks WHERE id=%s FOR UPDATE", (task_id,), ).fetchone() if existing is not None: decoded = _decode_row(existing) or {} config = decoded.get("config") or {} if ( config.get(_REPEAT_SOURCE_TASK_KEY) != source_task_id or config.get(_REPEAT_REQUEST_KEY) != request_id ): raise ConflictError("再次生成请求与现有任务冲突") if decoded.get("deleted_at"): raise ConflictError("此次再次生成创建的任务已被删除,请重新发起") return self._repeat_response( conn, existing, source_task_id=source_task_id, created=False, ) source_task = self._task_in_connection( conn, source_task_id, for_update=True, ) if ( source_task.get("status") != "completed" or source_task.get("results_confirmed") is False ): raise InvalidStateError("只有已完成并确认结果的任务可以再次生成") if source_task.get("preview_status") in ACTIVE_PREVIEW_STATUSES: raise ConflictError("源任务仍在处理切分,暂时不能再次生成") if expected_updated_at != _serialize_value(source_task.get("updated_at")): raise ConflictError("源任务已被其他操作修改,请刷新后重试") source_files = conn.execute( """ SELECT * FROM data_process_source_files WHERE task_id=%s AND deleted_at IS NULL ORDER BY created_at, id """, (source_task_id,), ).fetchall() source_file_ids = {str(row["id"]) for row in source_files} if source_file_ids != set(file_copies): raise ConflictError("源文件快照已变化,请刷新后重试") previews = conn.execute( """ SELECT * FROM data_process_preview_items WHERE task_id=%s ORDER BY source_file_id NULLS LAST, source_start NULLS LAST, created_at, id """, (source_task_id,), ).fetchall() if not previews: raise InvalidStateError("源任务没有可用于再次生成的切分结果") suffix = f"(再次生成-{task_id[-6:]})" base_name = str(source_task.get("name") or "数据处理任务") repeated_name = f"{base_name[: max(1, 150 - len(suffix))]}{suffix}" repeated_config = _business_config(source_task.get("config") or {}) repeated_config[_REPEAT_SOURCE_TASK_KEY] = source_task_id repeated_config[_REPEAT_REQUEST_KEY] = request_id input_count = sum(int(row.get("record_count") or 0) for row in source_files) task_row = conn.execute( """ INSERT INTO data_process_tasks (id, name, description, status, process_type, source_dataset_id, output_dataset_id, config, progress, input_count, output_count, filtered_count, duplicate_count, error_count, failure_reason, generation_run_id, results_confirmed, workflow_step, preview_status, preview_progress, preview_run_id, preview_failure_reason, preview_total_files, preview_completed_files, tenant_id, project_id, owner_id, approval_status, created_by, updated_by, created_at, updated_at) VALUES (%s, %s, %s, 'pending', %s, %s, NULL, %s, 20, %s, 0, 0, 0, 0, NULL, NULL, FALSE, 'preview', 'completed', 100, NULL, NULL, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) RETURNING * """, ( task_id, repeated_name, source_task.get("description") or "", source_task["process_type"], source_task.get("source_dataset_id"), json_dumps(repeated_config), input_count, len(source_files), len(source_files), source_task.get("tenant_id"), source_task.get("project_id"), source_task.get("owner_id"), source_task.get("approval_status") or "not_required", source_task.get("created_by"), source_task.get("created_by"), now, now, ), ).fetchone() file_id_map: dict[str, str] = {} for source in source_files: old_file_id = str(source["id"]) copy = file_copies[old_file_id] new_file_id = str(copy["id"]) storage_object_id, metadata = _source_storage_descriptor( { "storage_object_id": copy["storage_object_id"], "metadata": _json_value(source.get("metadata"), {}), }, task_id, new_file_id, ) file_id_map[old_file_id] = new_file_id conn.execute( """ INSERT INTO data_process_source_files (id, task_id, storage_object_id, name, size_bytes, record_count, file_format, checksum_sha256, version_no, content, content_preview, metadata, tenant_id, project_id, created_by, created_at, updated_at) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, 1, %s, %s, %s, %s, %s, %s, %s, %s) """, ( new_file_id, task_id, storage_object_id, source["name"], source.get("size_bytes") or 0, source.get("record_count") or 0, source.get("file_format"), source["checksum_sha256"], source.get("content") or "", source.get("content_preview"), json_dumps(metadata), source_task.get("tenant_id"), source_task.get("project_id"), source.get("created_by") or source_task.get("created_by"), now, now, ), ) for preview in previews: old_source_file_id = preview.get("source_file_id") conn.execute( """ INSERT INTO data_process_preview_items (id, task_id, source_file_id, original_content, edited_content, source_start, source_end, source_start_line, source_end_line, token_count, status, quality_score, created_at, updated_at) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) """, ( new_id("dpp"), task_id, file_id_map.get(str(old_source_file_id)) if old_source_file_id else None, preview.get("original_content") or "", preview.get("edited_content") or "", preview.get("source_start"), preview.get("source_end"), preview.get("source_start_line"), preview.get("source_end_line"), max(0, int(preview.get("token_count") or 0)), preview.get("status") or "original", json_dumps(_json_value(preview.get("quality_score"), {})), now, now, ), ) return self._repeat_response( conn, task_row or {}, source_task_id=source_task_id, created=True, ) except psycopg.errors.UniqueViolation as exc: raise ConflictError("再次生成任务名称或请求发生冲突,请重试") from exc def get_task(self, task_id: str, *, for_update: bool = False) -> dict[str, Any]: lock = " FOR UPDATE" if for_update else "" with self.connect() as conn: if for_update: row = conn.execute( f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}", (task_id,), ).fetchone() else: row = conn.execute( """ SELECT task.*, (SELECT COUNT(*) FROM data_process_source_files source WHERE source.task_id=task.id AND source.deleted_at IS NULL) AS source_file_count, (SELECT COUNT(*) FROM data_process_preview_items preview WHERE preview.task_id=task.id) AS preview_count, (SELECT COALESCE(json_agg(json_build_object( 'id', dataset.id, 'name', dataset.name, 'type', dataset.type, 'count', dataset.count, 'dataset_split', CASE dataset.type WHEN 'train' THEN 'train' WHEN 'val' THEN 'validation' WHEN 'test' THEN 'test' ELSE NULL END ) ORDER BY CASE dataset.type WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4 END), '[]'::json) FROM datasets dataset WHERE dataset.source='task' AND dataset.deleted_at IS NULL AND ( dataset.source_task_id=task.id OR (dataset.source_task_id IS NULL AND dataset.task_id=task.id) )) AS output_datasets, CASE WHEN task.started_at IS NOT NULL AND task.completed_at IS NOT NULL THEN EXTRACT(EPOCH FROM (task.completed_at - task.started_at)) ELSE NULL END AS duration_seconds FROM data_process_tasks task WHERE task.id=%s AND task.deleted_at IS NULL """, (task_id,), ).fetchone() if not row: raise NotFoundError("data process task not found") return _public_task(_decode_row(row)) or {} def _task_in_connection( self, conn: psycopg.Connection[dict[str, Any]], task_id: str, *, for_update: bool = False, ) -> dict[str, Any]: lock = " FOR UPDATE" if for_update else "" row = conn.execute( f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}", (task_id,), ).fetchone() if not row: raise NotFoundError("data process task not found") return _decode_row(row) or {} @staticmethod def _ensure_editable(task: dict[str, Any]) -> None: if task["status"] not in EDITABLE_STATUSES: raise InvalidStateError(f"task cannot be edited while status is {task['status']}") if task.get("preview_status") in ACTIVE_PREVIEW_STATUSES: raise InvalidStateError("task cannot be edited while preview is running") if task.get("output_dataset_id") and not _is_regeneration_prepared(task): raise InvalidStateError("published task cannot be edited") def update_workflow_step(self, task_id: str, workflow_step: str) -> dict[str, Any]: """独立保存向导位置,不触发配置或结果失效逻辑。""" if workflow_step not in WORKFLOW_STEPS: raise ValueError("invalid data process workflow step") with self.connect() as conn: row = conn.execute( """ UPDATE data_process_tasks SET workflow_step=%s, updated_at=%s WHERE id=%s AND deleted_at IS NULL RETURNING * """, (workflow_step, utcnow(), task_id), ).fetchone() if not row: raise NotFoundError("data process task not found") return _public_task(_decode_row(row)) or {} def recover_legacy_aborted_regeneration(self, task_id: str) -> dict[str, Any]: """恢复旧版在真正开始生成前误删的上一轮结果。 旧实现会在 ``POST /regenerate`` 时立即把已发布任务置为 pending、 清空结果并解除输出指针。三个已发布数据集仍是独立完整产物,因此只在 这个特征完全匹配时,使用其记录恢复结果和任务状态。该操作幂等,不会 触碰正常的新建待生成任务或已经开始的新一轮生成。 """ with self.connect() as conn: task = self._task_in_connection(conn, task_id, for_update=True) if ( task.get("status") != "pending" or task.get("generation_run_id") or task.get("output_dataset_id") or int(task.get("output_count") or 0) != 0 ): return {"recovered": False, "result_count": 0} result_count = int( ( conn.execute( "SELECT COUNT(*) AS count FROM data_process_results WHERE task_id=%s", (task_id,), ).fetchone() or {} ).get("count") or 0 ) if result_count: return {"recovered": False, "result_count": result_count} datasets = conn.execute( """ SELECT id, type, count, created_at FROM datasets WHERE source='task' AND deleted_at IS NULL AND ( source_task_id=%s OR (source_task_id IS NULL AND task_id=%s) ) ORDER BY CASE type WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4 END, created_at, id """, (task_id, task_id), ).fetchall() train_dataset = next( (dataset for dataset in datasets if dataset.get("type") == "train"), None, ) if not train_dataset: return {"recovered": False, "result_count": 0} dataset_ids = [str(dataset["id"]) for dataset in datasets] records = conn.execute( """ SELECT id, dataset_id, line_no, split, instruction, input, output, raw, status, source_result_id, preview_item_id, created_at FROM dataset_records WHERE dataset_id = ANY(%s) ORDER BY created_at, dataset_id, line_no NULLS LAST, id """, (dataset_ids,), ).fetchall() if not records: return {"recovered": False, "result_count": 0} preview_rows = conn.execute( "SELECT id FROM data_process_preview_items WHERE task_id=%s", (task_id,), ).fetchall() preview_ids = {str(row["id"]) for row in preview_rows} used_result_ids: set[str] = set() recovered_count = 0 for record in records: raw = _json_value(record.get("raw"), {}) raw = raw if isinstance(raw, dict) else {} candidate_id = str( record.get("source_result_id") or raw.get("source_result_id") or "" ) result_id = ( candidate_id if candidate_id and candidate_id not in used_result_ids else new_id("dpr") ) used_result_ids.add(result_id) candidate_preview_id = str( record.get("preview_item_id") or raw.get("preview_item_id") or "" ) preview_item_id = ( candidate_preview_id if candidate_preview_id in preview_ids else None ) instruction = str(record.get("instruction") or raw.get("instruction") or "") input_text = str(record.get("input") or raw.get("input") or "") chosen = str(raw.get("chosen") or "") rejected = str(raw.get("rejected") or "") output = str( record.get("output") or raw.get("output") or chosen or "" ) split = str(record.get("split") or raw.get("split") or "") or None status = str(record.get("status") or "valid") if status not in {"valid", "modified", "invalid"}: status = "valid" created_at = record.get("created_at") or utcnow() conn.execute( """ INSERT INTO data_process_results (id, task_id, preview_item_id, instruction, input, output, chosen, rejected, original_instruction, original_input, original_output, original_chosen, original_rejected, status, error, split, quality_score, created_at, updated_at) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, NULL, %s, '{}', %s, %s) """, ( result_id, task_id, preview_item_id, instruction, input_text, output, chosen, rejected, instruction, input_text, output, chosen, rejected, status, split, created_at, created_at, ), ) conn.execute( """ UPDATE dataset_records SET source_result_id=%s, preview_item_id=%s WHERE id=%s """, (result_id, preview_item_id, record["id"]), ) recovered_count += 1 now = utcnow() conn.execute( """ UPDATE data_process_tasks SET status='completed', progress=100, output_dataset_id=%s, output_count=%s, filtered_count=0, duplicate_count=0, error_count=(SELECT COUNT(*) FROM data_process_results WHERE task_id=%s AND status='invalid'), failure_reason=NULL, results_confirmed=TRUE, workflow_step='results', updated_at=%s WHERE id=%s """, ( train_dataset["id"], recovered_count, task_id, now, task_id, ), ) return {"recovered": True, "result_count": recovered_count} def update_task(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]: allowed = { "name", "description", "process_type", "source_dataset_id", } values: dict[str, Any] = {key: value for key, value in payload.items() if key in allowed} if not values and payload.get("config") is None: return self.get_task(task_id) try: with self.connect() as conn: task = self._task_in_connection(conn, task_id, for_update=True) self._ensure_editable(task) regeneration_prepared = _is_regeneration_prepared(task) if regeneration_prepared and any( key in payload and payload.get(key) != task.get(key) for key in ("process_type", "source_dataset_id") ): raise InvalidStateError( "process type and source dataset cannot change during regeneration" ) if payload.get("config") is not None: next_config = _business_config(payload["config"]) current_config = dict(task.get("config") or {}) for key in _INTERNAL_CONFIG_KEYS: if key in current_config: next_config[key] = current_config[key] values["config"] = json_dumps(next_config) invalidates_results = ( ("config" in payload and payload.get("config") != task.get("config")) or ( "process_type" in payload and payload.get("process_type") != task.get("process_type") ) or ( "source_dataset_id" in payload and payload.get("source_dataset_id") != task.get("source_dataset_id") ) ) if invalidates_results and not regeneration_prepared: values.update( { "status": "pending", "progress": 0, "output_count": 0, "filtered_count": 0, "duplicate_count": 0, "error_count": 0, "failure_reason": None, "generation_run_id": None, "results_confirmed": False, "preview_status": "idle", "preview_progress": 0, "preview_run_id": None, "preview_failure_reason": None, "preview_total_files": 0, "preview_completed_files": 0, "started_at": None, "completed_at": None, } ) conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,)) conn.execute( "DELETE FROM data_process_preview_items WHERE task_id=%s", (task_id,) ) if ( "process_type" in payload and payload.get("process_type") != task.get("process_type") ): conn.execute( """ UPDATE data_process_source_files SET deleted_at=%s, updated_at=%s WHERE task_id=%s AND deleted_at IS NULL """, (utcnow(), utcnow(), task_id), ) values["input_count"] = 0 values["updated_at"] = utcnow() assignments = ", ".join(f"{key}=%s" for key in values) row = conn.execute( f"UPDATE data_process_tasks SET {assignments} WHERE id=%s RETURNING *", [*values.values(), task_id], ).fetchone() except psycopg.errors.UniqueViolation as exc: raise ConflictError("data process task name already exists") from exc return _public_task(_decode_row(row)) or {} def prepare_regeneration( self, task_id: str, payload: dict[str, Any], ) -> dict[str, Any]: """非破坏性地保存重新生成配置。 准备阶段保留任务当前状态、结果、切片及已发布数据集。真正开始 生成时,才在同一事务内切换运行状态并清理上一轮结果。 """ # 先修复曾被旧版 prepare 提前清空的任务,再建立新的非破坏性草稿标记。 self.recover_legacy_aborted_regeneration(task_id) try: with self.connect() as conn: task = self._task_in_connection(conn, task_id, for_update=True) if task["status"] == "running": raise ConflictError("running task cannot be prepared for regeneration") if task.get("preview_status") in ACTIVE_PREVIEW_STATUSES: raise ConflictError("running preview cannot be prepared for regeneration") current_updated_at = _serialize_value(task.get("updated_at")) if payload["expected_updated_at"] != current_updated_at: raise ConflictError("data process task was modified by another request") process_type = str(payload["process_type"]) if process_type != str(task["process_type"]): raise InvalidStateError("process_type cannot be changed during regeneration") current_config = dict(task.get("config") or {}) next_config = _business_config(payload.get("config")) for key in (_REPEAT_SOURCE_TASK_KEY, _REPEAT_REQUEST_KEY): if key in current_config: next_config[key] = current_config[key] preview_invalidated = _preview_config_changed( process_type, current_config, next_config, ) now = utcnow() next_config[_REGENERATION_MARKER_KEY] = { "prepared": True, "preview_invalidated": preview_invalidated, "prepared_at": now, } # 002 迁移前发布的数据集只有 task_id。先补齐新关联字段,保证 # 解除任务输出指针后,详情和后续重新发布仍能定位原来的三份数据集。 conn.execute( """ UPDATE datasets SET source_task_id=%s, updated_at=%s WHERE source='task' AND source_task_id IS NULL AND task_id=%s AND deleted_at IS NULL """, (task_id, now, task_id), ) published_row = conn.execute( """ SELECT EXISTS( SELECT 1 FROM datasets WHERE source='task' AND deleted_at IS NULL AND ( source_task_id=%s OR (source_task_id IS NULL AND task_id=%s) ) ) AS exists """, (task_id, task_id), ).fetchone() published_outputs_preserved = bool(task.get("output_dataset_id")) or bool( published_row and published_row.get("exists") ) row = conn.execute( """ UPDATE data_process_tasks SET name=%s, description=%s, config=%s, updated_at=%s WHERE id=%s RETURNING * """, ( payload["name"], payload.get("description") or "", json_dumps(next_config), now, task_id, ), ).fetchone() except psycopg.errors.UniqueViolation as exc: raise ConflictError("data process task name already exists") from exc return { "task": _public_task(_decode_row(row)) or {}, "preview_invalidated": preview_invalidated, "published_outputs_preserved": published_outputs_preserved, } def delete_task(self, task_id: str, *, deleted_by: str | None = None) -> None: with self.connect() as conn: self._task_in_connection(conn, task_id, for_update=True) now = utcnow() conn.execute( """ UPDATE data_process_tasks SET status=CASE WHEN status='running' THEN 'stopped' ELSE status END, generation_run_id=NULL, preview_status=CASE WHEN preview_status IN ('queued', 'running') THEN 'cancelled' ELSE preview_status END, preview_run_id=NULL, deleted_at=%s, deleted_by=%s, updated_at=%s WHERE id=%s """, (now, deleted_by, now, task_id), )