From f04dc479bbcbf91338b243ade1ef2611914a6c7d Mon Sep 17 00:00:00 2001 From: caoxiaozhu Date: Thu, 23 Jul 2026 15:10:13 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=8C=E6=88=90=E6=95=B0=E6=8D=AE?= =?UTF-8?q?=E5=A4=84=E7=90=86=E6=8E=A5=E5=8F=A3=E4=B8=8E=E5=89=8D=E7=AB=AF?= =?UTF-8?q?=E6=8E=A5=E5=85=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/api/v1/endpoints/data_process.py | 1008 +++++++++++++ backend/app/api/v1/router.py | 3 +- backend/app/db/sql/002_data_process.sql | 235 +++ .../app/modules/data_process/algorithms.py | 912 +++++++++++ .../app/modules/data_process/generation.py | 250 ++++ .../app/modules/data_process/schema_cli.py | 65 + backend/app/modules/data_process/store.py | 1333 +++++++++++++++++ backend/app/schemas/data_process.py | 245 +++ backend/tests/test_data_process_algorithms.py | 279 ++++ backend/tests/test_data_process_api.py | 704 +++++++++ backend/tests/test_data_process_generation.py | 102 ++ backend/tests/test_data_process_migration.py | 29 + docs/data-process-design.md | 225 +++ .../regression-data-process-detail.mjs | 47 +- .../scripts/regression-data-process-list.mjs | 22 +- .../regression-data-process-wizard.mjs | 212 +-- frontend/src/api/modules/dataProcess.ts | 190 +++ frontend/src/types/dataProcess.ts | 226 +++ .../data-process/DataProcessCreateView.vue | 445 ++++-- .../data-process/DataProcessDetailView.vue | 718 ++++++--- .../data-process/DataProcessListView.vue | 137 +- .../create/PreviewCompareStep.vue | 33 +- .../data-process/create/ResultEditorStep.vue | 21 + .../data-process/create/SourceUploadStep.vue | 36 +- .../create/StructuredOptionsPanel.vue | 2 +- .../views/data-process/create/previewModel.ts | 525 +------ .../src/views/data-process/create/types.ts | 19 +- .../create/useDataProcessDraft.ts | 15 +- .../create/useDataProcessGeneration.ts | 232 ++- 29 files changed, 7126 insertions(+), 1144 deletions(-) create mode 100644 backend/app/api/v1/endpoints/data_process.py create mode 100644 backend/app/db/sql/002_data_process.sql create mode 100644 backend/app/modules/data_process/algorithms.py create mode 100644 backend/app/modules/data_process/generation.py create mode 100644 backend/app/modules/data_process/schema_cli.py create mode 100644 backend/app/modules/data_process/store.py create mode 100644 backend/app/schemas/data_process.py create mode 100644 backend/tests/test_data_process_algorithms.py create mode 100644 backend/tests/test_data_process_api.py create mode 100644 backend/tests/test_data_process_generation.py create mode 100644 backend/tests/test_data_process_migration.py create mode 100644 docs/data-process-design.md create mode 100644 frontend/src/api/modules/dataProcess.ts create mode 100644 frontend/src/types/dataProcess.ts diff --git a/backend/app/api/v1/endpoints/data_process.py b/backend/app/api/v1/endpoints/data_process.py new file mode 100644 index 0000000..b024ff6 --- /dev/null +++ b/backend/app/api/v1/endpoints/data_process.py @@ -0,0 +1,1008 @@ +from __future__ import annotations + +import hashlib +import ipaddress +import json +import os +import socket +from contextlib import contextmanager +from dataclasses import asdict +from pathlib import Path +from typing import Any, Iterator, Literal +from urllib.parse import urlsplit + +import psycopg +from fastapi import ( + APIRouter, + BackgroundTasks, + Body, + Depends, + File, + HTTPException, + Query, + UploadFile, +) +from psycopg.rows import dict_row + +from app.modules.data_process.algorithms import ( + chunk_unstructured, + decode_utf8, + desensitize_pii, + estimate_token_count, + generate_standard_records, + parse_text_content, + score_quality, +) +from app.modules.data_process.generation import generate_model_records +from app.modules.data_process.store import ( + ConflictError, + DataProcessStore, + DataProcessStoreError, + InvalidStateError, + NotFoundError, + get_data_process_store, +) +from app.schemas.data_process import ( + DataProcessTaskCreate, + DataProcessTaskUpdate, + DataProcessStatus, + ExternalPullRequest, + ExternalSourceRequest, + GenerateRequest, + PreviewBuildRequest, + PreviewItemCreate, + PreviewItemUpdate, + ProcessType, + PublishRequest, + ResultUpdate, +) + + +router = APIRouter(prefix="/data-process") +MAX_SOURCE_FILE_BYTES = 200 * 1024 * 1024 +MAX_SOURCE_FILE_COUNT = 20 +MAX_SOURCE_BATCH_BYTES = 500 * 1024 * 1024 +MAX_EXTERNAL_PULL_BYTES = 50 * 1024 * 1024 + + +def ok(data: Any = None, message: str = "ok") -> dict[str, Any]: + return {"code": 0, "message": message, "data": data} + + +def fail(status_code: int, message: str) -> HTTPException: + return HTTPException( + status_code=status_code, + detail={"code": status_code, "message": message, "data": None}, + ) + + +@contextmanager +def api_errors() -> Iterator[None]: + try: + yield + except NotFoundError as exc: + raise fail(404, str(exc)) from exc + except ConflictError as exc: + raise fail(409, str(exc)) from exc + except InvalidStateError as exc: + raise fail(409, str(exc)) from exc + except (DataProcessStoreError, ValueError) as exc: + raise fail(400, str(exc)) from exc + except psycopg.errors.UndefinedTable as exc: + raise fail(503, "data process schema is not installed; run schema_cli --check") from exc + except psycopg.OperationalError as exc: + raise fail(503, "data process database is unavailable") from exc + + +def _safe_file_name(value: str | None, fallback: str) -> str: + name = Path((value or "").replace("\\", "/")).name.replace("\x00", "").strip() + return name if name not in {"", ".", ".."} else fallback + + +def _value(config: dict[str, Any], snake_name: str, camel_name: str, default: Any) -> Any: + if snake_name in config: + return config[snake_name] + return config.get(camel_name, default) + + +def _preprocess_options(config: dict[str, Any]) -> set[str]: + values = _value(config, "preprocess_options", "preprocessOptions", []) + return {str(item) for item in values} if isinstance(values, list) else set() + + +def _preview_quality(content: str, config: dict[str, Any]) -> dict[str, Any]: + records = generate_standard_records( + [{"id": "quality-preview", "edited_content": content}], + split={"train": 100, "validation": 0, "test": 0}, + ) + record = records[0] if records else {"instruction": "", "input": "", "output": ""} + minimum = int(_value(config, "min_output_length", "minOutputLength", 20) or 20) + return asdict( + score_quality( + record, + min_output_length=max(1, minimum), + source_content=content, + ) + ) + + +def _build_preview_items( + task: dict[str, Any], source_files: list[dict[str, Any]] +) -> list[dict[str, Any]]: + config = task.get("config") or {} + process_type = task["process_type"] + preprocess_options = _preprocess_options(config) + should_desensitize = "desensitize" in preprocess_options + should_clean_invalid = bool( + preprocess_options & {"clean_invalid", "clean_invalid_content"} + ) + should_deduplicate = bool( + preprocess_options & {"deduplicate", "deduplicate_content"} + ) + seen_content_hashes: set[str] = set() + items: list[dict[str, Any]] = [] + + def append_item(item: dict[str, Any]) -> None: + content = str(item.get("edited_content") or "").strip() + if should_clean_invalid and not content: + return + content_hash = hashlib.sha256(content.encode("utf-8")).hexdigest() + if should_deduplicate and content_hash in seen_content_hashes: + return + seen_content_hashes.add(content_hash) + if not content: + item["status"] = "invalid" + items.append(item) + + for source in source_files: + parsed = parse_text_content( + source.get("content") or "", + filename=source.get("name"), + file_format=source.get("file_format"), + ) + if process_type == "unstructured": + chunks = chunk_unstructured( + parsed.text, + method=_value(config, "chunk_method", "chunkMethod", "semantic"), + chunk_size=int(_value(config, "chunk_size", "chunkSize", 800)), + chunk_overlap=int(_value(config, "chunk_overlap", "chunkOverlap", 100)), + min_chunk_size=int(_value(config, "min_chunk_size", "minChunkSize", 100)), + custom_delimiter=str( + _value(config, "custom_delimiter", "customDelimiter", "") or "" + ), + preserve_code_blocks=bool( + _value(config, "preserve_code_blocks", "preserveCodeBlocks", False) + ), + preserve_tables=bool( + _value(config, "preserve_tables", "preserveTables", False) + ), + preserve_lists=bool( + _value(config, "preserve_lists", "preserveLists", False) + ), + ) + for chunk in chunks: + content = chunk.content + pii_counts: dict[str, int] = {} + if should_desensitize: + content, pii_counts = desensitize_pii(content) + quality = _preview_quality(content, config) + quality["pii_replacements"] = pii_counts + append_item( + { + "source_file_id": source["id"], + "original_content": chunk.content, + "edited_content": content, + "source_start": chunk.start, + "source_end": chunk.end, + "source_start_line": chunk.start_line, + "source_end_line": chunk.end_line, + "token_count": chunk.token_count, + "status": "modified" if content != chunk.content else "original", + "quality_score": quality, + } + ) + continue + + record_contents = [ + json.dumps(record, ensure_ascii=False, separators=(",", ":")) + for record in parsed.records + if not should_clean_invalid + or any(value not in (None, "", [], {}) for value in record.values()) + ] + if not record_contents and parsed.text: + record_contents = [parsed.text] + for content in record_contents: + original_content = content + pii_counts = {} + if should_desensitize: + content, pii_counts = desensitize_pii(content) + quality = _preview_quality(content, config) + quality["pii_replacements"] = pii_counts + append_item( + { + "source_file_id": source["id"], + "original_content": original_content, + "edited_content": content, + "source_start": None, + "source_end": None, + "source_start_line": None, + "source_end_line": None, + "token_count": estimate_token_count(content), + "status": "modified" if content != original_content else "original", + "quality_score": quality, + } + ) + return items + + +def _all_preview_items(store: DataProcessStore, task_id: str) -> list[dict[str, Any]]: + """分页读取全部预览项,避免固定上限静默截断任务。""" + + items: list[dict[str, Any]] = [] + page = 1 + page_size = 5_000 + while True: + result = store.list_preview_items(task_id, page=page, page_size=page_size) + batch = result["items"] + items.extend(batch) + if len(items) >= int(result["total"]) or not batch: + return items + page += 1 + + +def _run_generation( + store: DataProcessStore, task_id: str, generation_run_id: str +) -> None: + try: + task = store.get_task(task_id) + if not store.generation_is_running(task_id, generation_run_id): + return + all_preview_items = _all_preview_items(store, task_id) + preview_items = [ + item + for item in all_preview_items + if item.get("status") != "invalid" + and str(item.get("edited_content") or item.get("original_content") or "").strip() + ] + pre_filtered_count = len(all_preview_items) - len(preview_items) + config = task.get("config") or {} + model_id = _value(config, "generation_model_id", "generationModelId", None) + generation_model: dict[str, Any] | None = None + if model_id: + generation_model = store.get_generation_model(str(model_id)) + task = store.save_generation_model_snapshot( + task_id, + generation_model, + generation_run_id=generation_run_id, + ) + config = task.get("config") or config + split = _value( + config, + "dataset_split", + "datasetSplit", + {"train": 80, "validation": 10, "test": 10}, + ) + pairs = ( + _value(config, "qa_pairs_per_chunk", "qaPairsPerChunk", 1) + if task["process_type"] == "unstructured" + else _value(config, "qa_pairs_per_row", "qaPairsPerRow", 1) + ) + if generation_model: + runtime_config = { + **config, + "generation_prompt": _value( + config, "generation_prompt", "generationPrompt", "" + ), + "max_tokens": _value(config, "max_tokens", "maxTokens", 1024), + "json_mode": _value(config, "json_mode", "jsonMode", False), + } + def report_progress(processed_count: int, total_count: int) -> None: + if not store.update_generation_progress( + task_id, + generation_run_id, + processed_count, + total_count, + ): + raise InvalidStateError("generation run is no longer active") + + generated = generate_model_records( + preview_items, + model=generation_model, + config=runtime_config, + task_id=task_id, + split=split, + qa_pairs_per_item=int(pairs or 1), + on_progress=report_progress, + ) + else: + generated = generate_standard_records( + preview_items, + qa_pairs_per_item=int(pairs or 1), + semantic_enrichment=bool( + _value(config, "semantic_enrichment", "semanticEnrichment", False) + ), + split=split, + split_seed=task_id, + ) + if not store.update_generation_progress( + task_id, + generation_run_id, + len(preview_items), + len(preview_items), + ): + return + + known_fingerprints: set[str] = set() + accepted: list[dict[str, Any]] = [] + filtered_count = pre_filtered_count + duplicate_count = 0 + error_count = 0 + quality_filter = bool( + _value(config, "quality_filter_enabled", "qualityFilterEnabled", False) + ) + filter_low = bool(_value(config, "filter_low_quality", "filterLowQuality", True)) + filter_short = bool( + _value(config, "filter_short_content", "filterShortContent", True) + ) + deduplicate = bool( + _preprocess_options(config) + & {"deduplicate", "deduplicate_content"} + ) + minimum = max(1, int(_value(config, "min_output_length", "minOutputLength", 20) or 20)) + preview_sources = { + str(item["id"]): str( + item.get("edited_content") or item.get("original_content") or "" + ) + for item in preview_items + } + + for record in generated: + quality = score_quality( + record, + min_output_length=minimum, + source_content=preview_sources.get(str(record.get("preview_item_id") or ""), ""), + known_fingerprints=known_fingerprints, + ) + if "duplicate_record" in quality.flags: + duplicate_count += 1 + else: + # 即使首条随后因短文本/低质量被过滤,也要阻止同批后续重复结果。 + known_fingerprints.add(quality.fingerprint) + should_filter = ( + (deduplicate and "duplicate_record" in quality.flags) + or ( + quality_filter + and filter_short + and "output_too_short" in quality.flags + ) + or (quality_filter and filter_low and not quality.is_valid) + ) + if not quality.is_valid: + error_count += 1 + record["status"] = "invalid" + record["error"] = ", ".join(quality.flags) or "quality validation failed" + if should_filter: + filtered_count += 1 + continue + record["quality_score"] = asdict(quality) + accepted.append(record) + + # stop 请求可能在纯函数计算期间到达,最终写入前再次检查状态。 + if store.generation_is_running(task_id, generation_run_id): + store.complete_generation( + task_id, + accepted, + generation_run_id=generation_run_id, + filtered_count=filtered_count, + duplicate_count=duplicate_count, + error_count=error_count, + ) + except Exception as exc: # noqa: BLE001 - background failures must be persisted + try: + if store.generation_is_running(task_id, generation_run_id): + store.mark_failed( + task_id, + str(exc), + generation_run_id=generation_run_id, + ) + except Exception: + return + + +@router.get("") +def list_tasks( + page: int = Query(default=1, ge=1), + page_size: int = Query(default=20, ge=1, le=200), + keyword: str | None = Query(default=None), + status: DataProcessStatus | None = Query(default=None), + process_type: ProcessType | None = Query(default=None), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + return ok( + store.list_tasks( + page=page, + page_size=page_size, + keyword=keyword, + status=status, + process_type=process_type, + ) + ) + + +@router.post("") +def create_task( + payload: DataProcessTaskCreate, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + task = store.create_task(payload.model_dump(mode="json")) + return ok(task, "data process task created") + + +@router.get("/{task_id}") +def task_detail( + task_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + task = store.get_task(task_id) + task["source_files"] = store.list_source_files(task_id) + return ok(task) + + +@router.put("/{task_id}") +def update_task( + task_id: str, + payload: DataProcessTaskUpdate, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + return ok( + store.update_task(task_id, payload.model_dump(exclude_unset=True, mode="json")), + "data process task updated", + ) + + +@router.delete("/{task_id}") +def delete_task( + task_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + store.delete_task(task_id) + return ok({"deleted": task_id}, "data process task deleted") + + +@router.get("/{task_id}/source-files") +def source_files( + task_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + return ok({"files": store.list_source_files(task_id)}) + + +@router.post("/{task_id}/source-files") +async def upload_source_files( + task_id: str, + files: list[UploadFile] = File(...), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + if not files: + raise fail(400, "at least one source file is required") + if len(files) > MAX_SOURCE_FILE_COUNT: + raise fail(413, f"a source batch may contain at most {MAX_SOURCE_FILE_COUNT} files") + prepared: list[dict[str, Any]] = [] + batch_size = 0 + with api_errors(): + store.get_task(task_id) + for upload in files: + raw = await upload.read(MAX_SOURCE_FILE_BYTES + 1) + if len(raw) > MAX_SOURCE_FILE_BYTES: + raise fail(413, f"source file exceeds {MAX_SOURCE_FILE_BYTES} bytes") + content = decode_utf8(raw) + name = _safe_file_name(upload.filename, "source.txt") + suffix = Path(name).suffix.lower() + if suffix not in { + ".txt", + ".md", + ".markdown", + ".csv", + ".tsv", + ".json", + ".jsonl", + ".ndjson", + }: + raise fail(415, f"unsupported source file format: {suffix or 'none'}") + parsed = parse_text_content(content, filename=name) + if not parsed.text: + raise fail(400, f"source file is empty: {name}") + normalized_raw = parsed.text.encode("utf-8") + batch_size += len(normalized_raw) + if batch_size > MAX_SOURCE_BATCH_BYTES: + raise fail(413, f"source batch exceeds {MAX_SOURCE_BATCH_BYTES} bytes") + record_count = len(parsed.records) or (1 if parsed.text else 0) + prepared.append( + { + "name": name, + "content": parsed.text, + "raw_size": len(normalized_raw), + "checksum_sha256": hashlib.sha256(normalized_raw).hexdigest(), + "file_format": parsed.format, + "record_count": record_count, + "metadata": { + "content_type": upload.content_type or "text/plain", + "original_size_bytes": len(raw), + "original_checksum_sha256": hashlib.sha256(raw).hexdigest(), + }, + "created_by": None, + } + ) + created = store.add_source_files(task_id, prepared) + return ok({"files": created}, "source files uploaded") + + +@router.get("/{task_id}/source-files/{file_id}/content") +def source_file_content( + task_id: str, + file_id: str, + start_line: int | None = Query(default=None, ge=1), + line_count: int = Query(default=200, ge=1, le=10_000), + offset: int = Query(default=0, ge=0), + limit: int = Query(default=100_000, ge=1, le=1_000_000), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + if start_line is not None: + return ok(store.source_content_lines(task_id, file_id, start_line, line_count)) + return ok(store.source_content_window(task_id, file_id, offset, limit)) + + +@router.delete("/{task_id}/source-files/{file_id}") +def delete_source_file( + task_id: str, + file_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + store.delete_source_file(task_id, file_id) + return ok({"deleted": file_id}, "source file removed") + + +def _external_postgres_connection(payload: ExternalSourceRequest) -> psycopg.Connection[Any]: + kind = payload.type.strip().lower() + parsed_url = urlsplit(payload.url) + scheme = parsed_url.scheme.lower() + if kind not in {"postgres", "postgresql"} or scheme not in {"postgres", "postgresql"}: + raise fail(501, f"external data source type is not supported: {payload.type}") + if parsed_url.username or parsed_url.password: + raise fail(400, "database credentials must use the account and password fields") + if payload.auth_mode not in {"none", "basic"}: + raise fail(400, "PostgreSQL supports only none or basic authentication") + if payload.auth_mode == "basic" and not payload.username: + raise fail(400, "database username is required for basic authentication") + hostname = parsed_url.hostname + if not hostname: + raise fail(400, "external PostgreSQL URL must include a hostname") + allow_private = os.getenv("DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB", "").lower() in { + "1", + "true", + "yes", + } + if not allow_private: + try: + addresses = { + item[4][0] + for item in socket.getaddrinfo( + hostname, + parsed_url.port or 5432, + type=socket.SOCK_STREAM, + ) + } + except socket.gaierror as exc: + raise fail(400, "external PostgreSQL hostname cannot be resolved") from exc + if any( + (address := ipaddress.ip_address(value)).is_private + or address.is_loopback + or address.is_link_local + or address.is_reserved + or address.is_unspecified + for value in addresses + ): + raise fail( + 403, + "private or local database addresses are disabled; " + "set DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB=true only in a trusted deployment", + ) + kwargs: dict[str, Any] = { + "connect_timeout": 5, + "row_factory": dict_row, + "application_name": "yg-ft-data-process-readonly", + "options": "-c default_transaction_read_only=on -c statement_timeout=30000", + } + if payload.auth_mode == "basic" and payload.username: + kwargs["user"] = payload.username + if payload.auth_mode == "basic" and payload.password: + kwargs["password"] = payload.password + return psycopg.connect(payload.url, **kwargs) + + +@router.post("/{task_id}/external/test") +def test_external_source( + task_id: str, + payload: ExternalSourceRequest, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + store.get_task(task_id) + try: + with _external_postgres_connection(payload) as conn: + conn.execute("SELECT 1 AS ok").fetchone() + except psycopg.Error as exc: + raise fail(502, "external PostgreSQL connection test failed") from exc + return ok({"connected": True, "type": payload.type}) + + +@router.post("/{task_id}/external/pull") +def pull_external_source( + task_id: str, + payload: ExternalPullRequest, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + query = (payload.query or "").strip() + if query.endswith(";"): + query = query[:-1].rstrip() + if ";" in query: + raise fail(400, "external pull accepts exactly one read-only query") + first_token = query.split(maxsplit=1)[0].lower() if query else "" + if first_token not in {"select", "with"}: + raise fail(400, "a read-only SELECT or WITH query is required for external pull") + with api_errors(): + store.get_task(task_id) + try: + with _external_postgres_connection(payload) as conn: + conn.execute("SET TRANSACTION READ ONLY") + conn.execute("SET LOCAL statement_timeout = '30s'") + cursor = conn.execute(query) + rows: list[dict[str, Any]] = [] + content_parts: list[str] = [] + content_size = 0 + while len(rows) < payload.limit: + batch = cursor.fetchmany(min(1_000, payload.limit - len(rows))) + if not batch: + break + for row in batch: + line = json.dumps(row, ensure_ascii=False, default=str) + "\n" + content_size += len(line.encode("utf-8")) + if content_size > MAX_EXTERNAL_PULL_BYTES: + raise fail(413, "external pull result exceeds the 50 MiB safety limit") + rows.append(row) + content_parts.append(line) + conn.rollback() + except psycopg.Error as exc: + raise fail(502, "external PostgreSQL query failed") from exc + if not rows: + raise fail(400, "external query returned no rows") + content = "".join(content_parts) + raw = content.encode("utf-8") + source = store.add_source_file( + task_id, + name=_safe_file_name(payload.file_name, "external-data.jsonl"), + content=content, + raw_size=len(raw), + checksum_sha256=hashlib.sha256(raw).hexdigest(), + file_format="jsonl", + record_count=len(rows), + metadata={ + "external_type": payload.type, + "external_host": urlsplit(payload.url).hostname, + "external_limit": payload.limit, + }, + ) + return ok({"files": [source]}, "external source pulled") + + +def _prepare_preview_items( + task_id: str, + store: DataProcessStore, + source_file_ids: list[str] | None = None, +) -> list[dict[str, Any]]: + task = store.get_task(task_id) + source_summaries = store.list_source_files(task_id) + if source_file_ids is not None: + requested = set(source_file_ids) + source_summaries = [item for item in source_summaries if item["id"] in requested] + found = {item["id"] for item in source_summaries} + missing = requested - found + if missing: + raise NotFoundError(f"source files not found: {', '.join(sorted(missing))}") + sources = [ + store.get_source_file(task_id, item["id"], include_content=True) + for item in source_summaries + ] + if not sources: + raise InvalidStateError("at least one source file is required") + items = _build_preview_items(task, sources) + if not items: + raise InvalidStateError("source files did not produce preview items") + return items + + +@router.post("/{task_id}/preview/build") +def build_preview( + task_id: str, + payload: PreviewBuildRequest = Body(default_factory=PreviewBuildRequest), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + items = _prepare_preview_items(task_id, store, payload.source_file_ids) + created = store.replace_preview_items(task_id, items) + return ok( + {"items": created, "total": len(created), "page": 1, "page_size": len(created)}, + "preview built", + ) + + +@router.get("/{task_id}/preview") +def preview_items( + task_id: str, + source_file_id: str | None = Query(default=None), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=200, ge=1, le=1000), + keyword: str | None = Query(default=None), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + return ok( + store.list_preview_items( + task_id, + source_file_id=source_file_id, + page=page, + page_size=page_size, + keyword=keyword, + ) + ) + + +@router.post("/{task_id}/preview") +def create_preview_item( + task_id: str, + payload: PreviewItemCreate, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + task = store.get_task(task_id) + item_payload = payload.model_dump(mode="json") + content = payload.edited_content + item_payload["token_count"] = estimate_token_count(content) + item_payload["quality_score"] = _preview_quality( + content, + task.get("config") or {}, + ) + if not content.strip(): + item_payload["status"] = "invalid" + item = store.create_preview_item(task_id, item_payload) + return ok(item, "preview item created") + + +@router.put("/{task_id}/preview/{preview_id}") +def update_preview_item( + task_id: str, + preview_id: str, + payload: PreviewItemUpdate, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + task = store.get_task(task_id) + update = payload.model_dump(exclude_unset=True, mode="json") + update["quality_score"] = _preview_quality(payload.edited_content, task.get("config") or {}) + item = store.update_preview_item( + task_id, + preview_id, + update, + ) + return ok(item, "preview item updated") + + +@router.delete("/{task_id}/preview/{preview_id}") +def delete_preview_item( + task_id: str, + preview_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + store.delete_preview_item(task_id, preview_id) + return ok({"deleted": preview_id}, "preview item deleted") + + +def _start_generation( + task_id: str, + payload: GenerateRequest, + background_tasks: BackgroundTasks, + store: DataProcessStore, +) -> dict[str, Any]: + with api_errors(): + task = store.start_generation(task_id, replace_existing=payload.replace_existing) + background_tasks.add_task( + _run_generation, + store, + task_id, + str(task["generation_run_id"]), + ) + return ok(store.progress(task_id), "data process generation started") + + +@router.post("/{task_id}/generate") +def generate( + task_id: str, + background_tasks: BackgroundTasks, + payload: GenerateRequest = Body(default_factory=GenerateRequest), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + return _start_generation(task_id, payload, background_tasks, store) + + +@router.post("/{task_id}/start") +def start( + task_id: str, + background_tasks: BackgroundTasks, + payload: GenerateRequest = Body(default_factory=GenerateRequest), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + items = _prepare_preview_items(task_id, store) + store.replace_preview_items(task_id, items) + return _start_generation(task_id, payload, background_tasks, store) + + +@router.post("/{task_id}/stop") +def stop( + task_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + store.stop_task(task_id) + return ok(store.progress(task_id), "data process task stopped") + + +@router.get("/{task_id}/progress") +def progress( + task_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + return ok(store.progress(task_id)) + + +@router.get("/{task_id}/results") +def results( + task_id: str, + page: int = Query(default=1, ge=1), + page_size: int = Query(default=100, ge=1, le=1000), + status: Literal["valid", "modified", "invalid"] | None = Query(default=None), + split: Literal["train", "validation", "test"] | None = Query(default=None), + keyword: str | None = Query(default=None), + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + return ok( + store.list_results( + task_id, + page=page, + page_size=page_size, + status=status, + split=split, + keyword=keyword, + ) + ) + + +@router.put("/{task_id}/results/{result_id}") +def update_result( + task_id: str, + result_id: str, + payload: ResultUpdate, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + task = store.get_task(task_id) + current = store.get_result(task_id, result_id) + update = payload.model_dump(exclude_unset=True, mode="json") + merged = {**current, **update} + preview_id = current.get("preview_item_id") + source_content = "" + if preview_id: + preview = store.get_preview_item(task_id, str(preview_id)) + source_content = str( + preview.get("edited_content") or preview.get("original_content") or "" + ) + minimum = max( + 1, + int( + _value( + task.get("config") or {}, + "min_output_length", + "minOutputLength", + 20, + ) + or 20 + ), + ) + quality = score_quality( + merged, + min_output_length=minimum, + source_content=source_content, + ) + update["quality_score"] = asdict(quality) + result = store.update_result( + task_id, + result_id, + update, + ) + return ok(result, "data process result updated") + + +@router.post("/{task_id}/results/{result_id}/restore") +def restore_result( + task_id: str, + result_id: str, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + task = store.get_task(task_id) + current = store.get_result(task_id, result_id) + restored = { + **current, + "instruction": current.get("original_instruction") or current.get("instruction") or "", + "input": current.get("original_input") or current.get("input") or "", + "output": current.get("original_output") or current.get("output") or "", + } + preview_id = current.get("preview_item_id") + source_content = "" + if preview_id: + preview = store.get_preview_item(task_id, str(preview_id)) + source_content = str( + preview.get("edited_content") or preview.get("original_content") or "" + ) + minimum = max( + 1, + int( + _value( + task.get("config") or {}, + "min_output_length", + "minOutputLength", + 20, + ) + or 20 + ), + ) + quality = score_quality( + restored, + min_output_length=minimum, + source_content=source_content, + ) + restored = store.update_result( + task_id, + result_id, + { + "instruction": restored["instruction"], + "input": restored["input"], + "output": restored["output"], + "quality_score": asdict(quality), + "expected_updated_at": current.get("updated_at"), + }, + ) + return ok(restored, "data process result restored") + + +@router.post("/{task_id}/publish") +def publish( + task_id: str, + payload: PublishRequest, + store: DataProcessStore = Depends(get_data_process_store), +) -> dict[str, Any]: + with api_errors(): + result = store.publish(task_id, payload.model_dump(mode="json")) + message = "dataset published" if result["created"] else "dataset already published" + return ok(result, message) diff --git a/backend/app/api/v1/router.py b/backend/app/api/v1/router.py index 5a4dbcb..07d41ba 100644 --- a/backend/app/api/v1/router.py +++ b/backend/app/api/v1/router.py @@ -1,9 +1,10 @@ from fastapi import APIRouter +from app.api.v1.endpoints.data_process import router as data_process_router from app.api.v1.endpoints.platform import router as platform_router from app.api.v1.endpoints.health import router as health_router api_router = APIRouter() api_router.include_router(health_router, tags=["health"]) +api_router.include_router(data_process_router, tags=["data-process"]) api_router.include_router(platform_router, tags=["platform"]) - diff --git a/backend/app/db/sql/002_data_process.sql b/backend/app/db/sql/002_data_process.sql new file mode 100644 index 0000000..f917bc0 --- /dev/null +++ b/backend/app/db/sql/002_data_process.sql @@ -0,0 +1,235 @@ +-- Data processing migration. +-- +-- IMPORTANT: This file is intentionally NOT wired into application startup. +-- Apply it explicitly in a controlled deployment, or call +-- DataProcessStore.ensure_schema() from an administrative command. + +BEGIN; + +-- This migration targets the current runtime schema created by +-- 001_platform_runtime.sql. Refuse the UUID/JSONB target-design schema instead +-- of partially altering it with incompatible TEXT foreign keys. +DO $$ +DECLARE + datasets_id_type TEXT; +BEGIN + SELECT format_type(a.atttypid, a.atttypmod) + INTO datasets_id_type + FROM pg_attribute a + JOIN pg_class c ON c.oid = a.attrelid + JOIN pg_namespace n ON n.oid = c.relnamespace + WHERE n.nspname = current_schema() + AND c.relname = 'datasets' + AND a.attname = 'id' + AND a.attnum > 0 + AND NOT a.attisdropped; + IF datasets_id_type IS NULL THEN + RAISE EXCEPTION '002_data_process.sql requires 001_platform_runtime.sql first'; + END IF; + IF datasets_id_type <> 'text' THEN + RAISE EXCEPTION + '002_data_process.sql supports only the current TEXT runtime schema; found datasets.id type %', + datasets_id_type; + END IF; +END $$; + +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS source_task_id TEXT; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS size_bytes BIGINT NOT NULL DEFAULT 0; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS record_count BIGINT NOT NULL DEFAULT 0; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}'; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS tenant_id TEXT; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS project_id TEXT; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS owner_id TEXT; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_by TEXT; +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ NOT NULL DEFAULT now(); +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now(); +ALTER TABLE datasets ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ; + +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS storage_object_id TEXT; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS current_version_id TEXT; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS size_bytes BIGINT NOT NULL DEFAULT 0; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS record_count BIGINT NOT NULL DEFAULT 0; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS file_format VARCHAR(40); +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS checksum_sha256 CHAR(64); +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS version_no INTEGER NOT NULL DEFAULT 1; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS source_task_id TEXT; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS tenant_id TEXT; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS project_id TEXT; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS created_by TEXT; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}'; +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ NOT NULL DEFAULT now(); +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now(); +ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ; + +CREATE TABLE IF NOT EXISTS data_process_tasks ( + id TEXT PRIMARY KEY, + name VARCHAR(150) NOT NULL, + description TEXT, + status VARCHAR(20) NOT NULL DEFAULT 'pending' + CHECK (status IN ('pending', 'running', 'completed', 'failed', 'stopped')), + process_type VARCHAR(20) NOT NULL + CHECK (process_type IN ('structured', 'unstructured', 'external')), + source_dataset_id TEXT REFERENCES datasets(id) ON DELETE SET NULL, + output_dataset_id TEXT REFERENCES datasets(id) ON DELETE SET NULL, + config TEXT NOT NULL DEFAULT '{}', + progress NUMERIC(5,2) NOT NULL DEFAULT 0 CHECK (progress >= 0 AND progress <= 100), + input_count BIGINT NOT NULL DEFAULT 0 CHECK (input_count >= 0), + output_count BIGINT NOT NULL DEFAULT 0 CHECK (output_count >= 0), + filtered_count BIGINT NOT NULL DEFAULT 0 CHECK (filtered_count >= 0), + duplicate_count BIGINT NOT NULL DEFAULT 0 CHECK (duplicate_count >= 0), + error_count BIGINT NOT NULL DEFAULT 0 CHECK (error_count >= 0), + failure_reason TEXT, + generation_run_id TEXT, + tenant_id TEXT, + project_id TEXT, + owner_id TEXT, + approval_status VARCHAR(30) NOT NULL DEFAULT 'not_required', + created_by TEXT, + updated_by TEXT, + deleted_by TEXT, + started_at TIMESTAMPTZ, + completed_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + deleted_at TIMESTAMPTZ +); + +ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS generation_run_id TEXT; + +CREATE UNIQUE INDEX IF NOT EXISTS uq_data_process_tasks_name_alive + ON data_process_tasks(name) WHERE deleted_at IS NULL; +CREATE INDEX IF NOT EXISTS idx_data_process_tasks_scope_status + ON data_process_tasks(tenant_id, project_id, status, created_at DESC) + WHERE deleted_at IS NULL; +CREATE INDEX IF NOT EXISTS idx_data_process_tasks_creator_created + ON data_process_tasks(created_by, created_at DESC) WHERE deleted_at IS NULL; + +CREATE TABLE IF NOT EXISTS data_process_source_files ( + id TEXT PRIMARY KEY, + task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE, + storage_object_id TEXT, + name TEXT NOT NULL, + size_bytes BIGINT NOT NULL DEFAULT 0 CHECK (size_bytes >= 0), + record_count BIGINT NOT NULL DEFAULT 0 CHECK (record_count >= 0), + file_format VARCHAR(40), + checksum_sha256 CHAR(64) NOT NULL, + version_no INTEGER NOT NULL DEFAULT 1 CHECK (version_no > 0), + content TEXT NOT NULL, + content_preview TEXT, + metadata TEXT NOT NULL DEFAULT '{}', + tenant_id TEXT, + project_id TEXT, + created_by TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + deleted_at TIMESTAMPTZ +); + +CREATE INDEX IF NOT EXISTS idx_data_process_source_files_task + ON data_process_source_files(task_id, created_at) WHERE deleted_at IS NULL; +CREATE UNIQUE INDEX IF NOT EXISTS uq_data_process_source_checksum_alive + ON data_process_source_files(task_id, checksum_sha256) WHERE deleted_at IS NULL; + +CREATE TABLE IF NOT EXISTS data_process_preview_items ( + id TEXT PRIMARY KEY, + task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE, + source_file_id TEXT REFERENCES data_process_source_files(id) ON DELETE CASCADE, + original_content TEXT NOT NULL DEFAULT '', + edited_content TEXT NOT NULL DEFAULT '', + source_start INTEGER CHECK (source_start IS NULL OR source_start >= 0), + source_end INTEGER CHECK (source_end IS NULL OR source_end >= 0), + source_start_line INTEGER CHECK (source_start_line IS NULL OR source_start_line > 0), + source_end_line INTEGER CHECK (source_end_line IS NULL OR source_end_line > 0), + token_count INTEGER NOT NULL DEFAULT 0 CHECK (token_count >= 0), + status VARCHAR(20) NOT NULL DEFAULT 'original' + CHECK (status IN ('original', 'modified', 'manual', 'invalid')), + quality_score TEXT NOT NULL DEFAULT '{}', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + CHECK (source_start IS NULL OR source_end IS NULL OR source_end >= source_start), + CHECK (source_start_line IS NULL OR source_end_line IS NULL OR source_end_line >= source_start_line) +); + +CREATE INDEX IF NOT EXISTS idx_data_process_preview_task_file + ON data_process_preview_items(task_id, source_file_id, created_at); + +CREATE TABLE IF NOT EXISTS data_process_results ( + id TEXT PRIMARY KEY, + task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE, + preview_item_id TEXT REFERENCES data_process_preview_items(id) ON DELETE SET NULL, + instruction TEXT NOT NULL, + input TEXT NOT NULL DEFAULT '', + output TEXT NOT NULL, + original_instruction TEXT, + original_input TEXT, + original_output TEXT, + status VARCHAR(20) NOT NULL DEFAULT 'valid' + CHECK (status IN ('valid', 'modified', 'invalid')), + error TEXT, + split VARCHAR(20) CHECK (split IS NULL OR split IN ('train', 'validation', 'test')), + quality_score TEXT NOT NULL DEFAULT '{}', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE INDEX IF NOT EXISTS idx_data_process_results_task_status + ON data_process_results(task_id, status, id); +CREATE INDEX IF NOT EXISTS idx_data_process_results_task_split + ON data_process_results(task_id, split); + +CREATE TABLE IF NOT EXISTS dataset_file_versions ( + id TEXT PRIMARY KEY, + dataset_file_id TEXT NOT NULL REFERENCES dataset_files(id) ON DELETE CASCADE, + version_no INTEGER NOT NULL CHECK (version_no > 0), + storage_object_id TEXT NOT NULL, + content_preview TEXT, + description TEXT, + base_version_id TEXT REFERENCES dataset_file_versions(id) ON DELETE SET NULL, + size_bytes BIGINT NOT NULL DEFAULT 0 CHECK (size_bytes >= 0), + record_count BIGINT NOT NULL DEFAULT 0 CHECK (record_count >= 0), + checksum_sha256 CHAR(64) NOT NULL, + source_task_id TEXT REFERENCES data_process_tasks(id) ON DELETE SET NULL, + metadata TEXT NOT NULL DEFAULT '{}', + created_by TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +ALTER TABLE dataset_file_versions ADD COLUMN IF NOT EXISTS source_task_id TEXT; +ALTER TABLE dataset_file_versions ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}'; +CREATE UNIQUE INDEX IF NOT EXISTS uq_dataset_file_versions_no_002 + ON dataset_file_versions(dataset_file_id, version_no); +CREATE INDEX IF NOT EXISTS idx_dataset_file_versions_source_task_002 + ON dataset_file_versions(source_task_id) WHERE source_task_id IS NOT NULL; + +CREATE TABLE IF NOT EXISTS dataset_records ( + id TEXT PRIMARY KEY, + dataset_id TEXT NOT NULL REFERENCES datasets(id) ON DELETE CASCADE, + dataset_file_id TEXT REFERENCES dataset_files(id) ON DELETE CASCADE, + version_id TEXT REFERENCES dataset_file_versions(id) ON DELETE CASCADE, + line_no INTEGER, + split VARCHAR(20) CHECK (split IS NULL OR split IN ('train', 'validation', 'test')), + instruction TEXT, + input TEXT, + output TEXT, + raw TEXT NOT NULL DEFAULT '{}', + status VARCHAR(20) NOT NULL DEFAULT 'valid' + CHECK (status IN ('valid', 'modified', 'invalid')), + source_task_id TEXT REFERENCES data_process_tasks(id) ON DELETE SET NULL, + source_result_id TEXT REFERENCES data_process_results(id) ON DELETE SET NULL, + preview_item_id TEXT REFERENCES data_process_preview_items(id) ON DELETE SET NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS source_task_id TEXT; +ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS source_result_id TEXT; +ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS preview_item_id TEXT; +CREATE INDEX IF NOT EXISTS idx_dataset_records_dataset_002 + ON dataset_records(dataset_id, id); +CREATE INDEX IF NOT EXISTS idx_dataset_records_source_task_002 + ON dataset_records(source_task_id, source_result_id); +CREATE INDEX IF NOT EXISTS idx_datasets_source_task_002 + ON datasets(source_task_id) WHERE source_task_id IS NOT NULL; +CREATE INDEX IF NOT EXISTS idx_dataset_files_source_task_002 + ON dataset_files(source_task_id) WHERE source_task_id IS NOT NULL; + +COMMIT; diff --git a/backend/app/modules/data_process/algorithms.py b/backend/app/modules/data_process/algorithms.py new file mode 100644 index 0000000..9e4fb12 --- /dev/null +++ b/backend/app/modules/data_process/algorithms.py @@ -0,0 +1,912 @@ +"""数据处理模块使用的无副作用算法。 + +本模块不访问数据库、文件系统或网络,便于 API、后台任务和测试共同复用。 +所有偏移量均为 Python 字符串偏移量,``TextChunk.content`` 始终等于 +``source[chunk.start:chunk.end]``。 +""" + +from __future__ import annotations + +import csv +import hashlib +import io +import json +import re +import unicodedata +from bisect import bisect_left +from collections.abc import Iterable, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Literal + + +TextFormat = Literal["json", "jsonl", "csv", "markdown", "txt"] +ChunkMethod = Literal["semantic", "heading", "fixed", "custom"] +DatasetSplit = Literal["train", "validation", "test"] + +SUPPORTED_TEXT_FORMATS: tuple[TextFormat, ...] = ( + "json", + "jsonl", + "csv", + "markdown", + "txt", +) + +_FORMAT_ALIASES: dict[str, TextFormat] = { + "json": "json", + "jsonl": "jsonl", + "ndjson": "jsonl", + "csv": "csv", + "tsv": "csv", + "md": "markdown", + "markdown": "markdown", + "txt": "txt", + "text": "txt", +} + +_EMAIL_PATTERN = re.compile( + r"(? str: + """严格解码 UTF-8 文本,并移除可选 BOM。 + + 不使用 ``errors='replace'``,避免上传内容损坏后仍被静默接收。 + """ + + if isinstance(raw, str): + return raw.removeprefix("\ufeff") + if not isinstance(raw, (bytes, bytearray, memoryview)): + raise TypeError("raw must be bytes-like or str") + try: + return bytes(raw).decode("utf-8-sig") + except UnicodeDecodeError as exc: + raise ValueError(f"content is not valid UTF-8 at byte {exc.start}") from exc + + +def parse_utf8_text(raw: bytes | bytearray | memoryview | str) -> str: + """``decode_utf8`` 的语义化别名,供上传服务直接调用。""" + + return decode_utf8(raw) + + +def normalize_text(text: str) -> str: + """规范 Unicode、换行和行尾空白,同时保留段落结构。""" + + if not isinstance(text, str): + raise TypeError("text must be str") + normalized = unicodedata.normalize("NFKC", text.removeprefix("\ufeff")) + normalized = normalized.replace("\r\n", "\n").replace("\r", "\n") + normalized = "".join( + char + for char in normalized + if char in {"\n", "\t"} or not unicodedata.category(char).startswith("C") + ) + lines = [re.sub(r"[\t \f\v]+$", "", line) for line in normalized.split("\n")] + return "\n".join(lines).strip() + + +def _normalize_format(value: str | None) -> TextFormat | None: + if value is None: + return None + normalized = value.strip().lower().removeprefix(".") + try: + return _FORMAT_ALIASES[normalized] + except KeyError as exc: + raise ValueError(f"unsupported text format: {value}") from exc + + +def detect_text_format( + *, + filename: str | None = None, + text: str = "", + file_format: str | None = None, +) -> TextFormat: + """按显式格式、扩展名和内容特征依次识别文本格式。""" + + explicit = _normalize_format(file_format) + if explicit: + return explicit + + if filename: + suffix = Path(filename).suffix.lower().removeprefix(".") + detected = _FORMAT_ALIASES.get(suffix) + if detected: + return detected + + stripped = text.strip() + if stripped: + if stripped[0] in "[{": + try: + json.loads(stripped) + except json.JSONDecodeError: + pass + else: + return "json" + + nonempty_lines = [line for line in stripped.splitlines() if line.strip()] + if len(nonempty_lines) > 1: + try: + for line in nonempty_lines: + json.loads(line) + except json.JSONDecodeError: + pass + else: + return "jsonl" + + if re.search(r"(?m)^(?:#{1,6}\s+|```|~~~)", stripped) or re.search( + r"(?m)^\s*\|.+\|\s*$", stripped + ): + return "markdown" + + sample = stripped[:8192] + try: + dialect = csv.Sniffer().sniff(sample, delimiters=",\t;") + rows = list(csv.reader(io.StringIO(sample), dialect)) + if len(rows) >= 2 and len(rows[0]) >= 2: + return "csv" + except csv.Error: + pass + + return "txt" + + +def _normalize_value(value: Any) -> Any: + if isinstance(value, str): + return normalize_text(value) + if isinstance(value, Mapping): + return {normalize_text(str(key)): _normalize_value(item) for key, item in value.items()} + if isinstance(value, list): + return [_normalize_value(item) for item in value] + return value + + +def _record_from_value(value: Any) -> dict[str, Any]: + if isinstance(value, Mapping): + return dict(_normalize_value(value)) + return {"value": _normalize_value(value)} + + +def extract_structured_records(text: str, file_format: str) -> list[dict[str, Any]]: + """从 JSON、JSONL 或 CSV 中提取规范化记录。 + + JSON 顶层对象若包含 ``records/data/items/rows`` 数组,则提取该数组; + 其他顶层对象视为单条记录。标量会稳定包装为 ``{"value": ...}``。 + """ + + normalized_format = _normalize_format(file_format) + if normalized_format not in {"json", "jsonl", "csv"}: + raise ValueError("structured record extraction only supports JSON, JSONL and CSV") + + normalized_text = normalize_text(text) + if not normalized_text: + return [] + + if normalized_format == "json": + try: + payload = json.loads(normalized_text) + except json.JSONDecodeError as exc: + raise ValueError(f"invalid JSON at line {exc.lineno}, column {exc.colno}: {exc.msg}") from exc + values: Sequence[Any] + if isinstance(payload, list): + values = payload + elif isinstance(payload, Mapping): + nested = next( + ( + payload[key] + for key in ("records", "data", "items", "rows") + if isinstance(payload.get(key), list) + ), + None, + ) + values = nested if isinstance(nested, list) else [payload] + else: + values = [payload] + return [_record_from_value(value) for value in values] + + if normalized_format == "jsonl": + records: list[dict[str, Any]] = [] + for line_number, line in enumerate(normalized_text.splitlines(), start=1): + if not line.strip(): + continue + try: + value = json.loads(line) + except json.JSONDecodeError as exc: + raise ValueError( + f"invalid JSONL at line {line_number}, " + f"column {exc.colno}: {exc.msg}" + ) from exc + records.append(_record_from_value(value)) + return records + + try: + dialect = csv.Sniffer().sniff(normalized_text[:8192], delimiters=",\t;") + except csv.Error: + dialect = csv.excel + reader = csv.DictReader(io.StringIO(normalized_text), dialect=dialect) + if not reader.fieldnames: + raise ValueError("CSV header is required") + headers = [normalize_text(header or "") for header in reader.fieldnames] + if any(not header for header in headers): + raise ValueError("CSV header cannot be empty") + if len(set(headers)) != len(headers): + raise ValueError("CSV headers must be unique") + reader.fieldnames = headers + + records = [] + for row in reader: + if None in row: + raise ValueError("CSV row has more fields than the header") + normalized_row = { + key: normalize_text(value or "") + for key, value in row.items() + } + if any(value for value in normalized_row.values()): + records.append(normalized_row) + return records + + +def parse_text_content( + raw: bytes | bytearray | memoryview | str, + *, + filename: str | None = None, + file_format: str | None = None, +) -> ParsedText: + """严格解码并解析支持的 UTF-8 文本格式。""" + + text = normalize_text(decode_utf8(raw)) + detected_format = detect_text_format(filename=filename, text=text, file_format=file_format) + records: list[dict[str, Any]] = [] + if detected_format in {"json", "jsonl", "csv"}: + records = extract_structured_records(text, detected_format) + return ParsedText(format=detected_format, text=text, records=tuple(records)) + + +def desensitize_pii(text: str) -> tuple[str, dict[str, int]]: + """掩码邮箱、中国大陆手机号和 15/18 位身份证号,并返回命中统计。""" + + if not isinstance(text, str): + raise TypeError("text must be str") + counts: dict[str, int] = {"email": 0, "phone": 0, "id_card": 0} + + def replace(pattern: re.Pattern[str], replacement: str, kind: str, value: str) -> str: + def replacer(_: re.Match[str]) -> str: + counts[kind] += 1 + return replacement + + return pattern.sub(replacer, value) + + masked = replace(_EMAIL_PATTERN, "[EMAIL]", "email", text) + masked = replace(_ID_CARD_PATTERN, "[ID_CARD]", "id_card", masked) + masked = replace(_PHONE_PATTERN, "[PHONE]", "phone", masked) + counts["total"] = sum(counts.values()) + return masked, counts + + +def estimate_token_count(text: str) -> int: + """无分词器依赖的确定性 token 估算,用于预览与保护性限流。""" + + return len(_TOKEN_PATTERN.findall(text)) + + +def _token_spans(text: str) -> list[tuple[int, int]]: + return [match.span() for match in _TOKEN_PATTERN.finditer(text)] + + +def _line_number(newline_offsets: list[int], offset: int) -> int: + # 换行符本身仍属于上一行;只有严格位于 offset 之前的换行才推进行号。 + return bisect_left(newline_offsets, offset) + 1 + + +def _token_index_at_or_after(spans: list[tuple[int, int]], offset: int) -> int: + starts = [span[0] for span in spans] + return bisect_left(starts, offset) + + +def _protected_markdown_ranges( + text: str, + *, + preserve_code_blocks: bool, + preserve_tables: bool, + preserve_lists: bool, +) -> list[tuple[int, int]]: + """找出不应从中间切开的 Markdown 代码块、表格和连续列表。""" + + lines: list[tuple[int, int, str]] = [] + cursor = 0 + for raw_line in text.splitlines(keepends=True): + end = cursor + len(raw_line) + lines.append((cursor, end, raw_line.rstrip("\r\n"))) + cursor = end + if cursor < len(text) or not lines: + lines.append((cursor, len(text), text[cursor:])) + + ranges: list[tuple[int, int]] = [] + code_line_indexes: set[int] = set() + if preserve_code_blocks: + open_block: tuple[int, str, int] | None = None + for index, (start, end, content) in enumerate(lines): + fence = re.match(r"^\s*(`{3,}|~{3,})", content) + if not fence: + continue + marker = fence.group(1)[0] + length = len(fence.group(1)) + if open_block is None: + open_block = (index, marker, length) + continue + first_index, open_marker, open_length = open_block + if marker == open_marker and length >= open_length: + ranges.append((lines[first_index][0], end)) + code_line_indexes.update(range(first_index, index + 1)) + open_block = None + if open_block is not None: + first_index = open_block[0] + ranges.append((lines[first_index][0], len(text))) + code_line_indexes.update(range(first_index, len(lines))) + + if preserve_tables: + index = 0 + while index + 1 < len(lines): + if index in code_line_indexes: + index += 1 + continue + header = lines[index][2].strip() + separator = lines[index + 1][2].strip().strip("|") + cells = [cell.strip() for cell in separator.split("|")] + if ( + "|" not in header + or len(cells) < 2 + or not all(re.fullmatch(r":?-{3,}:?", cell) for cell in cells) + ): + index += 1 + continue + end_index = index + 1 + while ( + end_index + 1 < len(lines) + and end_index + 1 not in code_line_indexes + and lines[end_index + 1][2].strip() + and "|" in lines[end_index + 1][2] + ): + end_index += 1 + ranges.append((lines[index][0], lines[end_index][1])) + index = end_index + 1 + + if preserve_lists: + list_pattern = re.compile(r"^\s*(?:[-+*]|\d+[.)])\s+\S") + continuation_pattern = re.compile(r"^\s{2,}\S") + index = 0 + while index < len(lines): + if index in code_line_indexes or not list_pattern.match(lines[index][2]): + index += 1 + continue + end_index = index + item_count = 1 + while end_index + 1 < len(lines) and end_index + 1 not in code_line_indexes: + next_line = lines[end_index + 1][2] + if list_pattern.match(next_line): + item_count += 1 + end_index += 1 + elif continuation_pattern.match(next_line): + end_index += 1 + else: + break + if item_count >= 2: + ranges.append((lines[index][0], lines[end_index][1])) + index = end_index + 1 + + merged: list[tuple[int, int]] = [] + for start, end in sorted(ranges): + if merged and start < merged[-1][1]: + merged[-1] = (merged[-1][0], max(merged[-1][1], end)) + else: + merged.append((start, end)) + return merged + + +def _range_containing( + ranges: Sequence[tuple[int, int]], offset: int +) -> tuple[int, int] | None: + return next((item for item in ranges if item[0] < offset < item[1]), None) + + +def _boundary_for_method( + text: str, + spans: list[tuple[int, int]], + start_index: int, + ideal_end_index: int, + minimum_end_index: int, + method: ChunkMethod, + custom_delimiter: str, +) -> tuple[int, int | None]: + if method == "fixed": + return ideal_end_index, None + + start_offset = spans[start_index][0] + ideal_end_offset = spans[ideal_end_index - 1][1] + minimum_end_offset = spans[minimum_end_index - 1][1] + search_text = text[start_offset:ideal_end_offset] + + if method == "custom": + delimiter = custom_delimiter.replace("\\n", "\n").replace("\\t", "\t") + if not delimiter: + raise ValueError("custom_delimiter is required for custom chunking") + relative_minimum = max(0, minimum_end_offset - start_offset) + delimiter_start = search_text.rfind(delimiter, relative_minimum) + if delimiter_start >= 0: + boundary_offset = start_offset + delimiter_start + len(delimiter) + boundary_index = _token_index_at_or_after(spans, boundary_offset) + if boundary_index > start_index: + return min(boundary_index, ideal_end_index), boundary_offset + return ideal_end_index, None + + if method == "heading": + heading_offsets = [ + start_offset + match.start() + for match in _HEADING_PATTERN.finditer(search_text) + if start_offset + match.start() >= minimum_end_offset + ] + if heading_offsets: + boundary_offset = heading_offsets[-1] + boundary_index = _token_index_at_or_after(spans, boundary_offset) + if start_index < boundary_index <= ideal_end_index: + return boundary_index, boundary_offset + + semantic_boundaries = [ + start_offset + match.end() + for match in _SEMANTIC_BOUNDARY_PATTERN.finditer(search_text) + if start_offset + match.end() >= minimum_end_offset + ] + if semantic_boundaries: + boundary_offset = semantic_boundaries[-1] + boundary_index = _token_index_at_or_after(spans, boundary_offset) + if boundary_index > start_index: + return min(boundary_index, ideal_end_index), boundary_offset + return ideal_end_index, None + + +def chunk_unstructured( + text: str, + *, + method: ChunkMethod = "semantic", + chunk_size: int = 800, + chunk_overlap: int = 100, + min_chunk_size: int = 100, + custom_delimiter: str = "", + preserve_code_blocks: bool = False, + preserve_tables: bool = False, + preserve_lists: bool = False, +) -> list[TextChunk]: + """按估算 token 切分非结构化文本。 + + overlap 足够时精确保留配置数量;短边界下会自动收缩,并且每轮至少推进 + 一个 token,避免异常配置或分隔符造成死循环。 + """ + + if method not in {"semantic", "heading", "fixed", "custom"}: + raise ValueError(f"unsupported chunk method: {method}") + if chunk_size <= 0: + raise ValueError("chunk_size must be greater than 0") + if chunk_overlap < 0 or chunk_overlap >= chunk_size: + raise ValueError("chunk_overlap must be in [0, chunk_size)") + if min_chunk_size <= 0 or min_chunk_size > chunk_size: + raise ValueError("min_chunk_size must be in [1, chunk_size]") + if chunk_overlap + min_chunk_size > chunk_size: + raise ValueError("chunk_overlap + min_chunk_size cannot exceed chunk_size") + if method == "custom" and not custom_delimiter: + raise ValueError("custom_delimiter is required for custom chunking") + + normalized = normalize_text(text) + if not normalized: + return [] + spans = _token_spans(normalized) + if not spans: + return [] + + newline_offsets = [index for index, char in enumerate(normalized) if char == "\n"] + protected_ranges = _protected_markdown_ranges( + normalized, + preserve_code_blocks=preserve_code_blocks, + preserve_tables=preserve_tables, + preserve_lists=preserve_lists, + ) + chunks: list[TextChunk] = [] + start_index = 0 + + while start_index < len(spans): + ideal_end_index = min(len(spans), start_index + chunk_size) + if ideal_end_index == len(spans): + end_index, end_override = ideal_end_index, len(normalized) + else: + minimum_end_index = min(ideal_end_index, start_index + min_chunk_size) + end_index, end_override = _boundary_for_method( + normalized, + spans, + start_index, + ideal_end_index, + minimum_end_index, + method, + custom_delimiter, + ) + if end_index <= start_index: + end_index = min(len(spans), start_index + chunk_size) + end_override = None + + start_offset = spans[start_index][0] + end_offset = end_override if end_override is not None else spans[end_index - 1][1] + end_offset = max(spans[end_index - 1][1], min(len(normalized), end_offset)) + split_range = _range_containing(protected_ranges, end_offset) + if split_range: + before_index = _token_index_at_or_after(spans, split_range[0]) + if before_index - start_index >= min_chunk_size: + end_index = before_index + end_offset = split_range[0] + else: + end_index = min( + len(spans), + max(start_index + 1, _token_index_at_or_after(spans, split_range[1])), + ) + end_offset = split_range[1] + content = normalized[start_offset:end_offset] + chunks.append( + TextChunk( + content=content, + start=start_offset, + end=end_offset, + start_line=_line_number(newline_offsets, start_offset), + end_line=_line_number(newline_offsets, max(start_offset, end_offset - 1)), + token_count=end_index - start_index, + ) + ) + + if end_index >= len(spans): + break + next_start = max(start_index + 1, end_index - chunk_overlap) + overlap_range = _range_containing(protected_ranges, spans[next_start][0]) + if overlap_range: + candidate = _token_index_at_or_after(spans, overlap_range[0]) + if candidate <= start_index: + candidate = _token_index_at_or_after(spans, overlap_range[1]) + next_start = min(len(spans), max(start_index + 1, candidate)) + start_index = next_start + + return chunks + + +def record_fingerprint(record: Mapping[str, Any]) -> str: + """计算与字典键顺序无关的稳定记录指纹。""" + + canonical = { + "instruction": normalize_text(str(record.get("instruction") or "")), + "input": normalize_text(str(record.get("input") or "")), + "output": normalize_text(str(record.get("output") or "")), + } + raw = json.dumps(canonical, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + return hashlib.sha256(raw.encode("utf-8")).hexdigest() + + +def _readability_score(text: str) -> float: + if not text: + return 0.0 + nonspace = [char for char in text if not char.isspace()] + if not nonspace: + return 0.0 + printable_ratio = sum(char.isprintable() for char in nonspace) / len(nonspace) + useful_ratio = sum( + char.isalnum() or "\u3400" <= char <= "\u9fff" or unicodedata.category(char).startswith("P") + for char in nonspace + ) / len(nonspace) + return round(100 * (0.65 * printable_ratio + 0.35 * useful_ratio), 2) + + +def _internal_duplicate_score(text: str) -> float: + units = [unit.strip().lower() for unit in re.split(r"[\n。!?!?;;]+", text) if unit.strip()] + if len(units) <= 1: + return 100.0 + return round(100 * len(set(units)) / len(units), 2) + + +def _source_relevance_score(record: Mapping[str, Any], source_content: str) -> float: + """估算结果与来源文本的词元覆盖率。 + + 这是无外部模型依赖、可重复的首版评分。没有来源文本(例如人工新增结果) + 时不扣分;存在来源时,以结果中的有效词元被来源覆盖的比例计分。 + """ + + source = normalize_text(source_content) + if not source: + return 100.0 + candidate = normalize_text( + "\n".join( + str(record.get(field) or "") for field in ("instruction", "input", "output") + ) + ) + + def semantic_tokens(text: str) -> set[str]: + return { + token.lower() + for token in _TOKEN_PATTERN.findall(text) + if token.isalnum() or "\u3400" <= token <= "\u9fff" + } + + source_tokens = semantic_tokens(source) + candidate_tokens = semantic_tokens(candidate) + if not candidate_tokens: + return 0.0 + if not source_tokens: + return 0.0 + return round(100 * len(candidate_tokens & source_tokens) / len(candidate_tokens), 2) + + +def score_quality( + record: Mapping[str, Any], + *, + min_output_length: int = 20, + source_content: str = "", + known_fingerprints: Iterable[str] = (), + threshold: float = 60.0, +) -> QualityScore: + """按完整性、长度、可读性、来源相关性和重复度计算质量分。""" + + if min_output_length <= 0: + raise ValueError("min_output_length must be greater than 0") + if not 0 <= threshold <= 100: + raise ValueError("threshold must be in [0, 100]") + + instruction = normalize_text(str(record.get("instruction") or "")) + input_text = normalize_text(str(record.get("input") or "")) + output = normalize_text(str(record.get("output") or "")) + flags: list[str] = [] + + completeness = 100.0 + if not instruction: + completeness -= 50 + flags.append("missing_instruction") + if not output: + completeness -= 50 + flags.append("missing_output") + + output_length = len(output) + length_score = round(min(100.0, output_length / min_output_length * 100), 2) + if output_length < min_output_length: + flags.append("output_too_short") + + readability = _readability_score("\n".join((instruction, input_text, output))) + if readability < 70: + flags.append("low_readability") + + relevance = _source_relevance_score(record, source_content) + if source_content and relevance < 30: + flags.append("low_source_relevance") + + fingerprint = record_fingerprint(record) + known = set(known_fingerprints) + duplicate = 0.0 if fingerprint in known else _internal_duplicate_score(output) + if duplicate == 0: + flags.append("duplicate_record") + elif duplicate < 70: + flags.append("repetitive_output") + + overall = round( + completeness * 0.35 + + length_score * 0.20 + + readability * 0.20 + + relevance * 0.15 + + duplicate * 0.10, + 2, + ) + hard_valid = bool(instruction and output) + return QualityScore( + overall=overall, + completeness=completeness, + length=length_score, + readability=readability, + relevance=relevance, + duplicate=duplicate, + is_valid=hard_valid and overall >= threshold, + flags=tuple(flags), + fingerprint=fingerprint, + ) + + +def stable_split( + value: str | int, + split: Mapping[str, int] | None = None, + *, + seed: str = "", +) -> DatasetSplit: + """按稳定哈希将记录划分到 train/validation/test。""" + + ratios = dict(split or {"train": 80, "validation": 10, "test": 10}) + required = {"train", "validation", "test"} + if set(ratios) != required: + raise ValueError("split must contain exactly train, validation and test") + if any(isinstance(value, bool) or not isinstance(value, int) or value < 0 for value in ratios.values()): + raise ValueError("split ratios must be non-negative integers") + if sum(ratios.values()) != 100: + raise ValueError("split ratios must sum to 100") + + digest = hashlib.sha256(f"{seed}:{value}".encode("utf-8")).digest() + bucket = int.from_bytes(digest[:8], "big") % 10_000 + train_boundary = ratios["train"] * 100 + validation_boundary = train_boundary + ratios["validation"] * 100 + if bucket < train_boundary: + return "train" + if bucket < validation_boundary: + return "validation" + return "test" + + +def _preview_content(item: Mapping[str, Any]) -> str: + for field in ("edited_content", "editedContent", "original_content", "originalContent", "content"): + value = item.get(field) + if value is not None: + return normalize_text(str(value)) + return "" + + +def _standard_fields(content: str) -> tuple[str, str, str]: + if not content: + return "", "", "" + + try: + payload = json.loads(content) + except json.JSONDecodeError: + payload = None + if isinstance(payload, Mapping): + instruction = next( + ( + str(payload[key]) + for key in ("instruction", "question", "prompt") + if payload.get(key) is not None + ), + "", + ) + input_text = next( + (str(payload[key]) for key in ("input", "context") if payload.get(key) is not None), + "", + ) + output = next( + (str(payload[key]) for key in ("output", "answer", "response") if payload.get(key) is not None), + "", + ) + if instruction or output: + return normalize_text(instruction), normalize_text(input_text), normalize_text(output) + + question_answer = re.match( + r"^\s*(?:问|question)\s*[::]\s*(.+?)(?:\n|\r\n?)\s*(?:答|answer)\s*[::]\s*(.+)\s*$", + content, + flags=re.IGNORECASE | re.DOTALL, + ) + if question_answer: + return normalize_text(question_answer.group(1)), "", normalize_text(question_answer.group(2)) + + lines = [line.strip() for line in content.splitlines() if line.strip()] + first_line = re.sub(r"^(?:问|question)\s*[::]\s*", "", lines[0], flags=re.IGNORECASE) + output = normalize_text("\n".join(lines[1:])) if len(lines) > 1 else normalize_text(content) + return normalize_text(first_line), "", output + + +def generate_standard_records( + preview_items: Iterable[Mapping[str, Any]], + *, + qa_pairs_per_item: int = 1, + semantic_enrichment: bool = False, + split: Mapping[str, int] | None = None, + split_seed: str = "", +) -> list[dict[str, Any]]: + """把预览内容确定性转换为标准 instruction/input/output 记录。 + + 该函数只负责本地标准化,不冒充 LLM;服务层可将其作为无模型模式或 + LLM 响应解析后的统一落库步骤。 + """ + + if not 1 <= qa_pairs_per_item <= 5: + raise ValueError("qa_pairs_per_item must be in [1, 5]") + prefixes = ( + "请结合实际情况说明:", + "请用通俗易懂的方式说明:", + "请从实际应用角度说明:", + "请简洁自然地说明:", + "请详细解答:", + ) + results: list[dict[str, Any]] = [] + for item_index, item in enumerate(preview_items): + content = _preview_content(item) + instruction, input_text, output = _standard_fields(content) + preview_id = str(item.get("id") or f"preview-{item_index + 1}") + for variant_index in range(qa_pairs_per_item): + variant_instruction = instruction + if variant_index: + if semantic_enrichment: + variant_instruction = f"{prefixes[variant_index]}{instruction}" + else: + variant_instruction = f"{instruction}(问法 {variant_index + 1})" + raw_id = f"{preview_id}:{variant_index + 1}" + result_id = f"result_{hashlib.sha256(raw_id.encode('utf-8')).hexdigest()[:16]}" + status = "valid" if variant_instruction and output else "invalid" + results.append( + { + "id": result_id, + "preview_item_id": preview_id, + "instruction": variant_instruction, + "input": input_text, + "output": output, + "original_instruction": variant_instruction, + "original_input": input_text, + "original_output": output, + "status": status, + "split": stable_split(result_id, split, seed=split_seed), + } + ) + return results + + +__all__ = [ + "ChunkMethod", + "DatasetSplit", + "ParsedText", + "QualityScore", + "SUPPORTED_TEXT_FORMATS", + "TextChunk", + "TextFormat", + "chunk_unstructured", + "decode_utf8", + "desensitize_pii", + "detect_text_format", + "estimate_token_count", + "extract_structured_records", + "generate_standard_records", + "normalize_text", + "parse_text_content", + "parse_utf8_text", + "record_fingerprint", + "score_quality", + "stable_split", +] diff --git a/backend/app/modules/data_process/generation.py b/backend/app/modules/data_process/generation.py new file mode 100644 index 0000000..04d1b1a --- /dev/null +++ b/backend/app/modules/data_process/generation.py @@ -0,0 +1,250 @@ +"""数据处理任务的大模型生成适配器。""" + +from __future__ import annotations + +import hashlib +import json +import re +from collections.abc import Callable, Iterable, Mapping +from typing import Any +from urllib.parse import urlsplit, urlunsplit + +import httpx + +from app.modules.data_process.algorithms import normalize_text, stable_split + + +class ModelGenerationError(ValueError): + """模型配置、响应或调用失败。""" + + +def chat_completions_url(value: str) -> str: + """把域名、基础 URL 或完整地址统一为 chat completions 地址。""" + + raw = (value or "").strip() + if not raw: + raise ModelGenerationError("generation model api_url is required") + if "://" not in raw: + raw = f"https://{raw}" + parsed = urlsplit(raw) + if parsed.scheme not in {"http", "https"} or not parsed.hostname: + raise ModelGenerationError("generation model api_url must be an HTTP(S) host or URL") + if parsed.username or parsed.password: + raise ModelGenerationError("generation model api_url must not contain credentials") + + path = parsed.path.rstrip("/") + if path.endswith("/chat/completions"): + target_path = path + elif path.endswith("/v1"): + target_path = f"{path}/chat/completions" + elif not path: + target_path = "/v1/chat/completions" + else: + target_path = f"{path}/v1/chat/completions" + return urlunsplit((parsed.scheme, parsed.netloc, target_path, "", "")) + + +def _message_content(payload: Mapping[str, Any]) -> str: + try: + content = payload["choices"][0]["message"]["content"] + except (KeyError, IndexError, TypeError) as exc: + raise ModelGenerationError("model response does not contain choices[0].message.content") from exc + if isinstance(content, str): + return content + if isinstance(content, list): + parts = [ + str(item.get("text") or "") + for item in content + if isinstance(item, Mapping) and item.get("type") in {None, "text", "output_text"} + ] + if parts: + return "".join(parts) + raise ModelGenerationError("model response content must be text") + + +def _json_payload(content: str) -> Any: + cleaned = re.sub(r"[\s\S]*?", "", content, flags=re.IGNORECASE).strip() + fenced = re.fullmatch(r"```(?:json)?\s*([\s\S]*?)\s*```", cleaned, flags=re.IGNORECASE) + if fenced: + cleaned = fenced.group(1).strip() + try: + return json.loads(cleaned) + except json.JSONDecodeError as exc: + raise ModelGenerationError( + f"model response is not valid JSON at line {exc.lineno}, column {exc.colno}" + ) from exc + + +def _result_items(payload: Any) -> list[Mapping[str, Any]]: + if isinstance(payload, list): + values = payload + elif isinstance(payload, Mapping): + nested = next( + ( + payload[key] + for key in ("items", "results", "data", "records") + if isinstance(payload.get(key), list) + ), + None, + ) + values = nested if isinstance(nested, list) else [payload] + else: + raise ModelGenerationError("model JSON must be an object or array") + items = [item for item in values if isinstance(item, Mapping)] + if not items: + raise ModelGenerationError("model JSON does not contain result objects") + return items + + +def _prompt_messages(prompt: str, content: str, count: int) -> list[dict[str, str]]: + schema_instruction = ( + f"必须只返回 JSON 对象,格式为 {{\"items\":[{{\"instruction\":\"...\"," + f"\"input\":\"...\",\"output\":\"...\"}}]}};items 必须包含 {count} 条。" + "instruction 和 output 不得为空,不要输出 Markdown 代码围栏或分析过程。" + ) + base_prompt = ( + normalize_text(prompt) + or "请根据来源内容生成可用于监督微调的问答数据。" + ) + if "{{ content }}" in base_prompt: + user_prompt = base_prompt.replace("{{ content }}", content) + return [ + {"role": "system", "content": schema_instruction}, + {"role": "user", "content": user_prompt}, + ] + return [ + {"role": "system", "content": f"{base_prompt}\n{schema_instruction}"}, + {"role": "user", "content": f"来源内容:\n{content}"}, + ] + + +def generate_model_records( + preview_items: Iterable[Mapping[str, Any]], + *, + model: Mapping[str, Any], + config: Mapping[str, Any], + task_id: str, + split: Mapping[str, int], + qa_pairs_per_item: int, + client: httpx.Client | None = None, + on_progress: Callable[[int, int], None] | None = None, +) -> list[dict[str, Any]]: + """调用 OpenAI 兼容接口,将预览切片生成标准训练记录。 + + 单条调用失败会产生可人工修复的 invalid 结果,不会丢弃整批任务。 + """ + + if not 1 <= qa_pairs_per_item <= 5: + raise ModelGenerationError("qa_pairs_per_item must be in [1, 5]") + endpoint = chat_completions_url(str(model.get("api_url") or "")) + model_name = str(model.get("online_model_name") or model.get("name") or "").strip() + if not model_name: + raise ModelGenerationError("generation model name is required") + + temperature = float(config.get("temperature", 0.7)) + max_tokens = int(config.get("max_tokens", 1024)) + timeout = max(1.0, min(120.0, float(config.get("request_timeout_seconds", 60)))) + retries = max(0, min(5, int(config.get("generation_retries", 2)))) + headers = {"Content-Type": "application/json"} + api_key = str(model.get("api_key") or "").strip() + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + + owns_client = client is None + http_client = client or httpx.Client(timeout=timeout) + results: list[dict[str, Any]] = [] + try: + preview_list = list(preview_items) + total_items = len(preview_list) + for item_index, item in enumerate(preview_list): + preview_id = str(item.get("id") or f"preview-{item_index + 1}") + content = normalize_text( + str(item.get("edited_content") or item.get("original_content") or "") + ) + request_payload: dict[str, Any] = { + "model": model_name, + "messages": _prompt_messages( + str(config.get("generation_prompt") or ""), + content, + qa_pairs_per_item, + ), + "temperature": temperature, + "max_tokens": max_tokens, + } + if bool(config.get("json_mode", False)): + request_payload["response_format"] = {"type": "json_object"} + + last_error: Exception | None = None + generated_items: list[Mapping[str, Any]] | None = None + for _ in range(retries + 1): + try: + response = http_client.post(endpoint, headers=headers, json=request_payload) + response.raise_for_status() + body = response.json() + if not isinstance(body, Mapping): + raise ModelGenerationError("model response body must be a JSON object") + generated_items = _result_items(_json_payload(_message_content(body))) + break + except (httpx.HTTPError, json.JSONDecodeError, ModelGenerationError) as exc: + last_error = exc + + if generated_items is None: + error_message = str(last_error or "model generation failed")[:2000] + result_id = f"result_{hashlib.sha256(f'{preview_id}:error'.encode()).hexdigest()[:16]}" + results.append( + { + "id": result_id, + "preview_item_id": preview_id, + "instruction": "模型生成失败,请人工补充", + "input": content, + "output": "", + "original_instruction": "模型生成失败,请人工补充", + "original_input": content, + "original_output": "", + "status": "invalid", + "error": error_message, + "split": stable_split(result_id, split, seed=task_id), + } + ) + if on_progress: + on_progress(item_index + 1, total_items) + continue + + for variant_index, value in enumerate(generated_items[:qa_pairs_per_item]): + instruction = normalize_text(str(value.get("instruction") or value.get("question") or "")) + input_text = normalize_text(str(value.get("input") or value.get("context") or "")) + output = normalize_text( + str( + value.get("output") + or value.get("answer") + or value.get("response") + or "" + ) + ) + raw_id = f"{preview_id}:{variant_index + 1}:{instruction}:{output}" + result_id = f"result_{hashlib.sha256(raw_id.encode()).hexdigest()[:16]}" + valid = bool(instruction and output) + results.append( + { + "id": result_id, + "preview_item_id": preview_id, + "instruction": instruction, + "input": input_text, + "output": output, + "original_instruction": instruction, + "original_input": input_text, + "original_output": output, + "status": "valid" if valid else "invalid", + "error": None if valid else "model result is missing instruction or output", + "split": stable_split(result_id, split, seed=task_id), + } + ) + if on_progress: + on_progress(item_index + 1, total_items) + finally: + if owns_client: + http_client.close() + return results + + +__all__ = ["ModelGenerationError", "chat_completions_url", "generate_model_records"] diff --git a/backend/app/modules/data_process/schema_cli.py b/backend/app/modules/data_process/schema_cli.py new file mode 100644 index 0000000..a2af7cd --- /dev/null +++ b/backend/app/modules/data_process/schema_cli.py @@ -0,0 +1,65 @@ +"""数据处理运行表的显式检查与安装命令。""" + +from __future__ import annotations + +import argparse +from urllib.parse import urlsplit + +from app.modules.data_process.store import DataProcessStore + + +def _target_label(database_url: str) -> str: + parsed = urlsplit(database_url) + database = parsed.path.strip("/") or "(unknown)" + return f"{parsed.hostname or '(unknown)'}:{parsed.port or 5432}/{database}" + + +def _schema_ready(store: DataProcessStore) -> bool: + with store.connect() as conn: + row = conn.execute( + """ + SELECT EXISTS ( + SELECT 1 + FROM information_schema.columns + WHERE table_schema=current_schema() + AND table_name='data_process_tasks' + AND column_name='generation_run_id' + ) AS ready + """ + ).fetchone() + return bool(row and row["ready"]) + + +def main() -> int: + parser = argparse.ArgumentParser( + description="检查或显式安装数据处理运行表(不会由应用启动自动执行)" + ) + action = parser.add_mutually_exclusive_group(required=True) + action.add_argument("--check", action="store_true", help="只读检查迁移是否已安装") + action.add_argument("--apply", action="store_true", help="执行 002 数据处理迁移") + parser.add_argument( + "--yes", + action="store_true", + help="确认允许修改 DATABASE_URL 指向的数据库;与 --apply 同时使用", + ) + args = parser.parse_args() + + store = DataProcessStore() + target = _target_label(store.database_url) + if args.check: + ready = _schema_ready(store) + print(f"数据处理 schema:{'已安装' if ready else '未安装'};目标:{target}") + return 0 if ready else 1 + if not args.yes: + parser.error("--apply 必须同时提供 --yes,确认修改目标数据库") + + print(f"正在安装数据处理 schema;目标:{target}") + store.ensure_schema() + if not _schema_ready(store): + raise RuntimeError("迁移执行后仍未检测到 generation_run_id") + print("数据处理 schema 安装完成") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/backend/app/modules/data_process/store.py b/backend/app/modules/data_process/store.py new file mode 100644 index 0000000..806ca23 --- /dev/null +++ b/backend/app/modules/data_process/store.py @@ -0,0 +1,1333 @@ +from __future__ import annotations + +import hashlib +import json +import uuid +from contextlib import contextmanager +from datetime import date, datetime, timezone +from functools import lru_cache +from pathlib import Path +from typing import Any, Iterator, Sequence + +import psycopg +from psycopg.rows import dict_row + +from app.core.config import get_settings +from app.modules.data_process.algorithms import estimate_token_count, stable_split + + +TASK_STATUSES = {"pending", "running", "completed", "failed", "stopped"} +EDITABLE_STATUSES = {"pending", "failed", "stopped", "completed"} + + +class DataProcessStoreError(RuntimeError): + pass + + +class NotFoundError(DataProcessStoreError): + pass + + +class ConflictError(DataProcessStoreError): + pass + + +class InvalidStateError(DataProcessStoreError): + pass + + +def utcnow() -> str: + return datetime.now(timezone.utc).replace(microsecond=0).isoformat().replace("+00:00", "Z") + + +def new_id(prefix: str) -> str: + return f"{prefix}_{uuid.uuid4().hex[:20]}" + + +def json_dumps(value: Any) -> str: + return json.dumps(value, ensure_ascii=False, separators=(",", ":")) + + +def _database_url(value: str) -> str: + return value.replace("postgresql+psycopg://", "postgresql://") + + +def _json_value(value: Any, default: Any) -> Any: + if value is None or value == "": + return default + if isinstance(value, (dict, list)): + return value + try: + return json.loads(value) + except (TypeError, json.JSONDecodeError): + return default + + +def _serialize_value(value: Any) -> Any: + if isinstance(value, (datetime, date)): + return value.isoformat().replace("+00:00", "Z") + return value + + +def _decode_row(row: dict[str, Any] | None) -> dict[str, Any] | None: + if row is None: + return None + item = {key: _serialize_value(value) for key, value in row.items()} + for key, default in { + "config": {}, + "metadata": {}, + "quality_score": {}, + "versions": [], + }.items(): + if key in item: + item[key] = _json_value(item[key], default) + return item + + +class DataProcessStore: + """数据处理持久层。 + + 构造函数不会连接数据库或执行迁移。部署方必须显式执行 002 SQL, + 或在受控的管理命令中调用 :meth:`ensure_schema`,避免应用启动时 + 修改远程数据库。 + """ + + def __init__(self, database_url: str | None = None) -> None: + self.database_url = _database_url(database_url or get_settings().database_url) + + @contextmanager + def connect(self) -> Iterator[psycopg.Connection[dict[str, Any]]]: + with psycopg.connect(self.database_url, row_factory=dict_row) as conn: + try: + yield conn + conn.commit() + except Exception: + conn.rollback() + raise + + def ensure_schema(self) -> None: + """显式安装数据处理表;API 路由和应用启动流程不会调用此方法。""" + schema_path = Path(__file__).resolve().parents[2] / "db" / "sql" / "002_data_process.sql" + sql = schema_path.read_text(encoding="utf-8") + with self.connect() as conn: + with conn.cursor() as cursor: + cursor.execute(sql) + + 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 = ["deleted_at IS NULL"] + params: list[Any] = [] + if keyword: + clauses.append("(name ILIKE %s OR COALESCE(description, '') ILIKE %s)") + pattern = f"%{keyword.strip()}%" + params.extend([pattern, pattern]) + if status: + clauses.append("status = %s") + params.append(status) + if process_type: + clauses.append("process_type = %s") + params.append(process_type) + if tenant_id: + clauses.append("tenant_id = %s") + params.append(tenant_id) + if project_id: + clauses.append("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 WHERE {where}", params + ).fetchone()["count"] + rows = conn.execute( + f""" + SELECT * FROM data_process_tasks + WHERE {where} + ORDER BY created_at DESC, id DESC + LIMIT %s OFFSET %s + """, + [*params, page_size, (page - 1) * page_size], + ).fetchall() + return { + "items": [_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, tenant_id, project_id, owner_id, created_by, updated_by, + created_at, updated_at) + VALUES (%s, %s, %s, 'pending', %s, %s, %s, 0, %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(payload.get("config") or {}), + 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 _decode_row(row) or {} + + 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: + 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 {} + + 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("output_dataset_id"): + raise InvalidStateError("published task cannot be edited") + + 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 payload.get("config") is not None: + values["config"] = json_dumps(payload["config"]) + if not values: + 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) + 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: + 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, + "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 _decode_row(row) or {} + + def delete_task(self, task_id: str, *, deleted_by: str | None = None) -> None: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if task["status"] == "running": + raise InvalidStateError("running task must be stopped before deletion") + now = utcnow() + conn.execute( + """ + UPDATE data_process_tasks + SET deleted_at=%s, deleted_by=%s, updated_at=%s + WHERE id=%s + """, + (now, deleted_by, now, task_id), + ) + + def list_source_files(self, task_id: str) -> list[dict[str, Any]]: + self.get_task(task_id) + with self.connect() as conn: + rows = conn.execute( + """ + SELECT id, task_id, storage_object_id, name, size_bytes, record_count, + file_format, checksum_sha256, version_no, content_preview, metadata, + tenant_id, project_id, created_by, created_at, updated_at + FROM data_process_source_files + WHERE task_id=%s AND deleted_at IS NULL + ORDER BY created_at, id + """, + (task_id,), + ).fetchall() + return [_decode_row(row) or {} for row in rows] + + def add_source_file( + self, + task_id: str, + *, + name: str, + content: str, + raw_size: int, + checksum_sha256: str, + file_format: str, + record_count: int, + metadata: dict[str, Any] | None = None, + created_by: str | None = None, + ) -> dict[str, Any]: + return self.add_source_files( + task_id, + [ + { + "name": name, + "content": content, + "raw_size": raw_size, + "checksum_sha256": checksum_sha256, + "file_format": file_format, + "record_count": record_count, + "metadata": metadata or {}, + "created_by": created_by, + } + ], + )[0] + + def add_source_files( + self, + task_id: str, + files: Sequence[dict[str, Any]], + ) -> list[dict[str, Any]]: + """在同一事务中登记一个上传批次,任一文件失败则全部回滚。""" + + if not files: + raise DataProcessStoreError("at least one source file is required") + now = utcnow() + created: list[dict[str, Any]] = [] + try: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + self._ensure_editable(task) + for payload in files: + file_id = new_id("dpsf") + storage_object_id = f"db://data-process/{task_id}/{file_id}/v1" + metadata_payload = { + "storage_backend": "database", + **(payload.get("metadata") or {}), + } + row = 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) + RETURNING id, task_id, storage_object_id, name, size_bytes, record_count, + file_format, checksum_sha256, version_no, content_preview, metadata, + tenant_id, project_id, created_by, created_at, updated_at + """, + ( + file_id, + task_id, + storage_object_id, + payload["name"], + payload["raw_size"], + payload["record_count"], + payload["file_format"], + payload["checksum_sha256"], + payload["content"], + str(payload["content"])[:2000], + json_dumps(metadata_payload), + task.get("tenant_id"), + task.get("project_id"), + payload.get("created_by") or task.get("created_by"), + now, + now, + ), + ).fetchone() + created.append(_decode_row(row) or {}) + 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,)) + conn.execute( + """ + UPDATE data_process_tasks + SET status='pending', progress=0, output_count=0, filtered_count=0, + duplicate_count=0, error_count=0, failure_reason=NULL, + generation_run_id=NULL, started_at=NULL, completed_at=NULL, input_count=( + SELECT COALESCE(SUM(record_count), 0) + FROM data_process_source_files + WHERE task_id=%s AND deleted_at IS NULL + ), updated_at=%s + WHERE id=%s + """, + (task_id, now, task_id), + ) + except psycopg.errors.UniqueViolation as exc: + raise ConflictError( + "the same source file content is already attached to this task" + ) from exc + return created + + def get_source_file( + self, task_id: str, file_id: str, *, include_content: bool = True + ) -> dict[str, Any]: + # 先验证父任务仍然可见,避免软删除任务后通过已知文件 ID 读取正文。 + self.get_task(task_id) + content_column = ", content" if include_content else "" + with self.connect() as conn: + row = conn.execute( + f""" + SELECT id, task_id, storage_object_id, name, size_bytes, record_count, + file_format, checksum_sha256, version_no, content_preview, metadata, + tenant_id, project_id, created_by, created_at, updated_at{content_column} + FROM data_process_source_files + WHERE id=%s AND task_id=%s AND deleted_at IS NULL + """, + (file_id, task_id), + ).fetchone() + if not row: + raise NotFoundError("source file not found") + return _decode_row(row) or {} + + def source_content_window( + self, task_id: str, file_id: str, offset: int, limit: int + ) -> dict[str, Any]: + source_file = self.get_source_file(task_id, file_id, include_content=True) + content = str(source_file.pop("content", "")) + window = content[offset : offset + limit] + return { + "file": source_file, + "content": window, + "offset": offset, + "limit": limit, + "total_chars": len(content), + "has_more": offset + len(window) < len(content), + } + + def source_content_lines( + self, + task_id: str, + file_id: str, + start_line: int, + line_count: int, + ) -> dict[str, Any]: + source_file = self.get_source_file(task_id, file_id, include_content=True) + content = str(source_file.pop("content", "")) + lines = content.splitlines(keepends=True) + start_index = min(len(lines), start_line - 1) + selected = lines[start_index : start_index + line_count] + end_line = start_index + len(selected) + return { + "file": source_file, + "content": "".join(selected), + "start_line": start_line, + "end_line": end_line, + "line_count": len(selected), + "total_lines": len(lines), + "has_more": end_line < len(lines), + } + + def delete_source_file(self, task_id: str, file_id: str) -> None: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + self._ensure_editable(task) + row = conn.execute( + """ + UPDATE data_process_source_files + SET deleted_at=%s, updated_at=%s + WHERE id=%s AND task_id=%s AND deleted_at IS NULL + RETURNING id + """, + (utcnow(), utcnow(), file_id, task_id), + ).fetchone() + if not row: + raise NotFoundError("source file not found") + conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,)) + conn.execute( + "DELETE FROM data_process_preview_items WHERE source_file_id=%s", (file_id,) + ) + conn.execute( + """ + UPDATE data_process_tasks + SET status='pending', progress=0, output_count=0, filtered_count=0, + duplicate_count=0, error_count=0, failure_reason=NULL, + input_count=(SELECT COALESCE(SUM(record_count), 0) + FROM data_process_source_files + WHERE task_id=%s AND deleted_at IS NULL), + updated_at=%s + WHERE id=%s + """, + (task_id, utcnow(), task_id), + ) + + def replace_preview_items( + self, task_id: str, items: Sequence[dict[str, Any]] + ) -> list[dict[str, Any]]: + now = utcnow() + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + self._ensure_editable(task) + 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,)) + created: list[dict[str, Any]] = [] + for item in items: + row = 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) + RETURNING * + """, + ( + item.get("id") or new_id("dpp"), + task_id, + item.get("source_file_id"), + item.get("original_content") or "", + item.get("edited_content", item.get("original_content") or ""), + item.get("source_start"), + item.get("source_end"), + item.get("source_start_line"), + item.get("source_end_line"), + max(0, int(item.get("token_count") or 0)), + item.get("status") or "original", + json_dumps(item.get("quality_score") or {}), + now, + now, + ), + ).fetchone() + created.append(_decode_row(row) or {}) + conn.execute( + """ + UPDATE data_process_tasks + SET status='pending', progress=20, output_count=0, filtered_count=0, + duplicate_count=0, error_count=0, failure_reason=NULL, updated_at=%s + WHERE id=%s + """, + (now, task_id), + ) + return created + + def list_preview_items( + self, + task_id: str, + *, + source_file_id: str | None = None, + page: int = 1, + page_size: int = 200, + keyword: str | None = None, + ) -> dict[str, Any]: + self.get_task(task_id) + clauses = ["task_id=%s"] + params: list[Any] = [task_id] + if source_file_id: + clauses.append("source_file_id=%s") + params.append(source_file_id) + if keyword: + clauses.append("(original_content ILIKE %s OR edited_content ILIKE %s)") + pattern = f"%{keyword.strip()}%" + params.extend([pattern, pattern]) + where = " AND ".join(clauses) + with self.connect() as conn: + total = conn.execute( + f"SELECT COUNT(*) AS count FROM data_process_preview_items WHERE {where}", params + ).fetchone()["count"] + rows = conn.execute( + f""" + SELECT * FROM data_process_preview_items + WHERE {where} + ORDER BY source_file_id NULLS LAST, source_start NULLS LAST, created_at, id + LIMIT %s OFFSET %s + """, + [*params, page_size, (page - 1) * page_size], + ).fetchall() + return { + "items": [_decode_row(row) for row in rows], + "total": int(total), + "page": page, + "page_size": page_size, + } + + def get_preview_item(self, task_id: str, preview_id: str) -> dict[str, Any]: + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM data_process_preview_items WHERE id=%s AND task_id=%s", + (preview_id, task_id), + ).fetchone() + if not row: + raise NotFoundError("preview item not found") + return _decode_row(row) or {} + + def create_preview_item(self, task_id: str, item: dict[str, Any]) -> dict[str, Any]: + now = utcnow() + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + self._ensure_editable(task) + if item.get("source_file_id"): + source = conn.execute( + """ + SELECT id FROM data_process_source_files + WHERE id=%s AND task_id=%s AND deleted_at IS NULL + """, + (item["source_file_id"], task_id), + ).fetchone() + if not source: + raise NotFoundError("source file not found") + row = 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) + RETURNING * + """, + ( + new_id("dpp"), + task_id, + item.get("source_file_id"), + item.get("original_content") or "", + item.get("edited_content") or "", + item.get("source_start"), + item.get("source_end"), + item.get("source_start_line"), + item.get("source_end_line"), + max(0, int(item.get("token_count") or 0)), + item.get("status") or "manual", + json_dumps(item.get("quality_score") or {}), + now, + now, + ), + ).fetchone() + self._invalidate_results(conn, task_id, now) + return _decode_row(row) or {} + + def update_preview_item( + self, task_id: str, preview_id: str, payload: dict[str, Any] + ) -> dict[str, Any]: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + self._ensure_editable(task) + existing = conn.execute( + "SELECT * FROM data_process_preview_items WHERE id=%s AND task_id=%s", + (preview_id, task_id), + ).fetchone() + if not existing: + raise NotFoundError("preview item not found") + expected_updated_at = payload.get("expected_updated_at") + current_updated_at = _serialize_value(existing.get("updated_at")) + if expected_updated_at and expected_updated_at != current_updated_at: + raise ConflictError("preview item was modified by another request") + edited = payload["edited_content"] + status = payload.get("status") + if not status: + if not edited.strip(): + status = "invalid" + elif edited == existing["original_content"]: + status = "original" + else: + status = "modified" + now = utcnow() + row = conn.execute( + """ + UPDATE data_process_preview_items + SET edited_content=%s, token_count=%s, status=%s, quality_score=%s, + updated_at=%s + WHERE id=%s AND task_id=%s + RETURNING * + """, + ( + edited, + estimate_token_count(edited), + status, + json_dumps(payload.get("quality_score") or {}), + now, + preview_id, + task_id, + ), + ).fetchone() + self._invalidate_results(conn, task_id, now) + return _decode_row(row) or {} + + def delete_preview_item(self, task_id: str, preview_id: str) -> None: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + self._ensure_editable(task) + row = conn.execute( + "DELETE FROM data_process_preview_items WHERE id=%s AND task_id=%s RETURNING id", + (preview_id, task_id), + ).fetchone() + if not row: + raise NotFoundError("preview item not found") + self._invalidate_results(conn, task_id, utcnow()) + + def _invalidate_results( + self, + conn: psycopg.Connection[dict[str, Any]], + task_id: str, + now: str, + ) -> None: + conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,)) + conn.execute( + """ + UPDATE data_process_tasks + SET status='pending', progress=20, output_count=0, filtered_count=0, + duplicate_count=0, error_count=0, failure_reason=NULL, + generation_run_id=NULL, updated_at=%s + WHERE id=%s + """, + (now, task_id), + ) + + def start_generation(self, task_id: str, *, replace_existing: bool = True) -> dict[str, Any]: + if not replace_existing: + raise DataProcessStoreError("incremental generation is not supported") + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if task.get("output_dataset_id"): + raise InvalidStateError("published task cannot be regenerated") + if task["status"] == "running": + raise ConflictError("data process task is already running") + preview_count = conn.execute( + "SELECT COUNT(*) AS count FROM data_process_preview_items WHERE task_id=%s", + (task_id,), + ).fetchone()["count"] + if not preview_count: + raise InvalidStateError("preview must be built before generation") + conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,)) + now = utcnow() + generation_run_id = new_id("dprun") + row = conn.execute( + """ + UPDATE data_process_tasks + SET status='running', progress=30, failure_reason=NULL, started_at=%s, + completed_at=NULL, filtered_count=0, duplicate_count=0, error_count=0, + generation_run_id=%s, updated_at=%s + WHERE id=%s + RETURNING * + """, + (now, generation_run_id, now, task_id), + ).fetchone() + return _decode_row(row) or {} + + def stop_task(self, task_id: str) -> dict[str, Any]: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if task["status"] != "running": + raise InvalidStateError("only a running task can be stopped") + now = utcnow() + row = conn.execute( + """ + UPDATE data_process_tasks + SET status='stopped', failure_reason=NULL, generation_run_id=NULL, + updated_at=%s + WHERE id=%s RETURNING * + """, + (now, task_id), + ).fetchone() + return _decode_row(row) or {} + + def generation_is_running(self, task_id: str, generation_run_id: str) -> bool: + task = self.get_task(task_id) + return ( + task["status"] == "running" + and task.get("generation_run_id") == generation_run_id + ) + + def update_generation_progress( + self, + task_id: str, + generation_run_id: str, + processed_count: int, + total_count: int, + ) -> bool: + ratio = processed_count / max(1, total_count) + progress = min(95.0, 30.0 + ratio * 65.0) + with self.connect() as conn: + row = conn.execute( + """ + UPDATE data_process_tasks + SET progress=%s, updated_at=%s + WHERE id=%s AND status='running' AND generation_run_id=%s + RETURNING id + """, + (progress, utcnow(), task_id, generation_run_id), + ).fetchone() + return row is not None + + def complete_generation( + self, + task_id: str, + results: Sequence[dict[str, Any]], + *, + generation_run_id: str, + filtered_count: int = 0, + duplicate_count: int = 0, + error_count: int = 0, + ) -> dict[str, Any]: + now = utcnow() + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if ( + task["status"] != "running" + or task.get("generation_run_id") != generation_run_id + ): + return task + conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,)) + for result in results: + conn.execute( + """ + INSERT INTO data_process_results + (id, task_id, preview_item_id, instruction, input, output, + original_instruction, original_input, original_output, status, error, + split, quality_score, created_at, updated_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """, + ( + result.get("id") or new_id("dpr"), + task_id, + result.get("preview_item_id"), + result.get("instruction") or "", + result.get("input") or "", + result.get("output") or "", + result.get("original_instruction", result.get("instruction") or ""), + result.get("original_input", result.get("input") or ""), + result.get("original_output", result.get("output") or ""), + result.get("status") or "valid", + result.get("error"), + result.get("split"), + json_dumps(result.get("quality_score") or {}), + now, + now, + ), + ) + row = conn.execute( + """ + UPDATE data_process_tasks + SET status='completed', progress=100, output_count=%s, filtered_count=%s, + duplicate_count=%s, error_count=%s, failure_reason=NULL, + completed_at=%s, generation_run_id=NULL, updated_at=%s + WHERE id=%s AND generation_run_id=%s RETURNING * + """, + ( + len(results), + filtered_count, + duplicate_count, + error_count, + now, + now, + task_id, + generation_run_id, + ), + ).fetchone() + return _decode_row(row) or {} + + def mark_failed( + self, task_id: str, reason: str, *, generation_run_id: str + ) -> dict[str, Any]: + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if ( + task["status"] != "running" + or task.get("generation_run_id") != generation_run_id + ): + return task + now = utcnow() + row = conn.execute( + """ + UPDATE data_process_tasks + SET status='failed', failure_reason=%s, completed_at=%s, + generation_run_id=NULL, updated_at=%s + WHERE id=%s AND generation_run_id=%s RETURNING * + """, + (reason[:4000], now, now, task_id, generation_run_id), + ).fetchone() + return _decode_row(row) or {} + + def progress(self, task_id: str) -> dict[str, Any]: + task = self.get_task(task_id) + return { + "task_id": task["id"], + "status": task["status"], + "progress": float(task.get("progress") or 0), + "input_count": int(task.get("input_count") or 0), + "output_count": int(task.get("output_count") or 0), + "filtered_count": int(task.get("filtered_count") or 0), + "duplicate_count": int(task.get("duplicate_count") or 0), + "error_count": int(task.get("error_count") or 0), + "failure_reason": task.get("failure_reason"), + "started_at": task.get("started_at"), + "completed_at": task.get("completed_at"), + } + + def list_results( + self, + task_id: str, + *, + page: int = 1, + page_size: int = 100, + status: str | None = None, + split: str | None = None, + keyword: str | None = None, + ) -> dict[str, Any]: + self.get_task(task_id) + clauses = ["task_id=%s"] + params: list[Any] = [task_id] + if status: + clauses.append("status=%s") + params.append(status) + if split: + clauses.append("split=%s") + params.append(split) + if keyword: + clauses.append("(instruction ILIKE %s OR input ILIKE %s OR output ILIKE %s)") + pattern = f"%{keyword.strip()}%" + params.extend([pattern, pattern, pattern]) + where = " AND ".join(clauses) + with self.connect() as conn: + total = conn.execute( + f"SELECT COUNT(*) AS count FROM data_process_results WHERE {where}", params + ).fetchone()["count"] + rows = conn.execute( + f""" + SELECT * FROM data_process_results WHERE {where} + ORDER BY created_at, id LIMIT %s OFFSET %s + """, + [*params, page_size, (page - 1) * page_size], + ).fetchall() + return { + "items": [_decode_row(row) for row in rows], + "total": int(total), + "page": page, + "page_size": page_size, + } + + def get_result(self, task_id: str, result_id: str) -> dict[str, Any]: + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM data_process_results WHERE id=%s AND task_id=%s", + (result_id, task_id), + ).fetchone() + if not row: + raise NotFoundError("data process result not found") + return _decode_row(row) or {} + + def update_result( + self, task_id: str, result_id: str, payload: dict[str, Any] + ) -> dict[str, Any]: + allowed = {"instruction", "input", "output", "quality_score"} + values = {key: value for key, value in payload.items() if key in allowed} + if "quality_score" in values: + values["quality_score"] = json_dumps(values["quality_score"]) + if not values: + raise DataProcessStoreError("no result fields supplied") + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if task["status"] == "running": + raise InvalidStateError("results cannot be edited while generation is running") + if task.get("output_dataset_id"): + raise InvalidStateError("published results cannot be edited") + current = conn.execute( + "SELECT * FROM data_process_results WHERE id=%s AND task_id=%s", + (result_id, task_id), + ).fetchone() + if not current: + raise NotFoundError("data process result not found") + expected_updated_at = payload.get("expected_updated_at") + current_updated_at = _serialize_value(current.get("updated_at")) + if expected_updated_at and expected_updated_at != current_updated_at: + raise ConflictError("data process result was modified by another request") + merged = {**current, **values} + quality = payload.get("quality_score") or {} + hard_valid = bool( + str(merged.get("instruction") or "").strip() + and str(merged.get("output") or "").strip() + ) + quality_valid = bool(quality.get("is_valid", hard_valid)) + changed = any( + str(merged.get(field) or "") + != str(merged.get(f"original_{field}") or "") + for field in ("instruction", "input", "output") + ) + status = "invalid" if not hard_valid or not quality_valid else ( + "modified" if changed else "valid" + ) + values["status"] = status + flags = quality.get("flags") if isinstance(quality, dict) else None + values["error"] = ", ".join(str(flag) for flag in flags or []) or ( + "quality validation failed" if status == "invalid" else None + ) + values["updated_at"] = utcnow() + assignments = ", ".join(f"{key}=%s" for key in values) + row = conn.execute( + f"""UPDATE data_process_results SET {assignments} + WHERE id=%s AND task_id=%s RETURNING *""", + [*values.values(), result_id, task_id], + ).fetchone() + conn.execute( + """ + UPDATE data_process_tasks + SET error_count=( + SELECT COUNT(*) FROM data_process_results + WHERE task_id=%s AND status='invalid' + ), updated_at=%s + WHERE id=%s + """, + (task_id, utcnow(), task_id), + ) + return _decode_row(row) or {} + + def get_generation_model(self, model_id: str) -> dict[str, Any]: + with self.connect() as conn: + row = conn.execute( + """ + SELECT id, name, type, purpose, model_source, description, path, + api_url, api_key, online_model_name, create_time + FROM models WHERE id=%s + """, + (model_id,), + ).fetchone() + if not row: + raise NotFoundError("generation model not found") + return _decode_row(row) or {} + + def save_generation_model_snapshot( + self, + task_id: str, + model_snapshot: dict[str, Any], + *, + generation_run_id: str, + ) -> dict[str, Any]: + # API 密钥仅用于本次调用,绝不能进入任务配置、详情响应或审计快照。 + safe_snapshot = { + key: value for key, value in model_snapshot.items() if key != "api_key" + } + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if ( + task["status"] != "running" + or task.get("generation_run_id") != generation_run_id + ): + raise InvalidStateError("generation run is no longer active") + config = dict(task.get("config") or {}) + config["generation_model_snapshot"] = safe_snapshot + row = conn.execute( + """ + UPDATE data_process_tasks SET config=%s, updated_at=%s + WHERE id=%s AND generation_run_id=%s RETURNING * + """, + (json_dumps(config), utcnow(), task_id, generation_run_id), + ).fetchone() + return _decode_row(row) or {} + + def publish(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]: + """发布有效结果;任务行锁保证重复请求返回同一数据集。""" + with self.connect() as conn: + task = self._task_in_connection(conn, task_id, for_update=True) + if task.get("output_dataset_id"): + dataset = conn.execute( + "SELECT * FROM datasets WHERE id=%s", (task["output_dataset_id"],) + ).fetchone() + if dataset: + return {"dataset": _decode_row(dataset), "created": False} + # 数据集被外部流程清理后,解除断链并重新发布。 + conn.execute( + "UPDATE data_process_tasks SET output_dataset_id=NULL WHERE id=%s", + (task_id,), + ) + if task["status"] != "completed": + raise InvalidStateError("only a completed task can be published") + rows = conn.execute( + """ + SELECT * FROM data_process_results + WHERE task_id=%s ORDER BY created_at, id + """, + (task_id,), + ).fetchall() + if not rows: + raise InvalidStateError("task has no results to publish") + invalid_count = sum( + 1 + for row in rows + if row["status"] == "invalid" + or not str(row.get("instruction") or "").strip() + or not str(row.get("output") or "").strip() + ) + if invalid_count: + raise InvalidStateError(f"task contains {invalid_count} invalid results") + + dataset_id = new_id("dataset") + file_id = new_id("dfile") + version_id = new_id("dfv") + now = utcnow() + requested_split = payload.get("split") or { + "train": 80, + "validation": 10, + "test": 10, + } + records = [ + { + "instruction": row["instruction"], + "input": row["input"], + "output": row["output"], + "split": stable_split( + str(row["id"]), + requested_split, + seed=task_id, + ), + } + for row in rows + ] + content = "".join(json_dumps(record) + "\n" for record in records) + raw = content.encode("utf-8") + checksum = hashlib.sha256(raw).hexdigest() + storage_object_id = f"db://data-process/{task_id}/{file_id}/v1" + source_result_ids = [row["id"] for row in rows] + metadata = { + "source": "data_process", + "storage_backend": "database", + "storage_object_id": storage_object_id, + "source_task_id": task_id, + "source_file_ids": [item["id"] for item in self._source_ids(conn, task_id)], + "source_result_ids": source_result_ids, + "format": payload.get("format") or "alpaca_jsonl", + "split": payload.get("split") or {}, + } + try: + dataset = conn.execute( + """ + INSERT INTO datasets + (id, name, type, storage_type, source, task_id, source_task_id, + size, size_bytes, count, record_count, description, metadata, + tenant_id, project_id, owner_id, created_by, create_time, + created_at, updated_at) + VALUES (%s, %s, %s, %s, 'task', %s, %s, %s, %s, %s, %s, %s, %s, + %s, %s, %s, %s, %s, %s, %s) + RETURNING * + """, + ( + dataset_id, + payload["dataset_name"], + payload.get("dataset_type") or "train", + payload.get("storage_type") or "local", + task_id, + task_id, + f"{len(raw)} B", + len(raw), + len(records), + len(records), + payload.get("description") or task.get("description") or "", + json_dumps(metadata), + task.get("tenant_id"), + task.get("project_id"), + task.get("owner_id"), + payload.get("created_by") or task.get("created_by"), + now, + now, + now, + ), + ).fetchone() + version = { + "id": version_id, + "version_no": 1, + "description": "data process publish", + "checksum_sha256": checksum, + "size_bytes": len(raw), + "record_count": len(records), + "created_at": now, + "source_task_id": task_id, + "storage_object_id": storage_object_id, + } + conn.execute( + """ + INSERT INTO dataset_files + (id, dataset_id, name, storage_object_id, size, content, + active_version_id, versions, create_time, current_version_id, + size_bytes, record_count, file_format, checksum_sha256, version_no, + source_task_id, tenant_id, project_id, created_by, metadata, + created_at, updated_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, + %s, %s, 1, %s, %s, %s, %s, %s, %s, %s) + """, + ( + file_id, + dataset_id, + f"{payload['dataset_name']}.jsonl", + storage_object_id, + f"{len(raw)} B", + content, + version_id, + json_dumps([version]), + now, + version_id, + len(raw), + len(records), + "jsonl", + checksum, + task_id, + task.get("tenant_id"), + task.get("project_id"), + payload.get("created_by") or task.get("created_by"), + json_dumps(metadata), + now, + now, + ), + ) + conn.execute( + """ + INSERT INTO dataset_file_versions + (id, dataset_file_id, version_no, storage_object_id, content_preview, + description, size_bytes, record_count, checksum_sha256, + source_task_id, metadata, created_by, created_at) + VALUES (%s, %s, 1, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """, + ( + version_id, + file_id, + storage_object_id, + content[:2000], + "data process publish", + len(raw), + len(records), + checksum, + task_id, + json_dumps(metadata), + payload.get("created_by") or task.get("created_by"), + now, + ), + ) + for line_number, (source_row, record) in enumerate( + zip(rows, records, strict=True), start=1 + ): + conn.execute( + """ + INSERT INTO dataset_records + (id, dataset_id, dataset_file_id, version_id, line_no, split, + instruction, input, output, raw, status, source_task_id, + source_result_id, preview_item_id, created_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, + %s, %s, %s, %s) + """, + ( + new_id("drec"), + dataset_id, + file_id, + version_id, + line_number, + record["split"], + record["instruction"], + record["input"], + record["output"], + json_dumps( + { + **record, + "source_task_id": task_id, + "source_result_id": source_row["id"], + "preview_item_id": source_row.get("preview_item_id"), + } + ), + source_row["status"], + task_id, + source_row["id"], + source_row.get("preview_item_id"), + now, + ), + ) + except psycopg.errors.UniqueViolation as exc: + raise ConflictError("dataset name already exists") from exc + conn.execute( + """ + UPDATE data_process_tasks + SET output_dataset_id=%s, updated_at=%s, updated_by=%s + WHERE id=%s + """, + (dataset_id, now, payload.get("created_by"), task_id), + ) + return {"dataset": _decode_row(dataset), "created": True} + + @staticmethod + def _source_ids( + conn: psycopg.Connection[dict[str, Any]], task_id: str + ) -> list[dict[str, Any]]: + return conn.execute( + """ + SELECT id FROM data_process_source_files + WHERE task_id=%s AND deleted_at IS NULL ORDER BY created_at, id + """, + (task_id,), + ).fetchall() + + +@lru_cache +def get_data_process_store() -> DataProcessStore: + return DataProcessStore() diff --git a/backend/app/schemas/data_process.py b/backend/app/schemas/data_process.py new file mode 100644 index 0000000..1ceac43 --- /dev/null +++ b/backend/app/schemas/data_process.py @@ -0,0 +1,245 @@ +from __future__ import annotations + +from enum import StrEnum +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + + +def _config_value(config: dict[str, Any], snake_name: str, camel_name: str, default: Any) -> Any: + if snake_name in config: + return config[snake_name] + return config.get(camel_name, default) + + +def _validate_process_config(config: dict[str, Any]) -> None: + split = _config_value(config, "dataset_split", "datasetSplit", None) + if split is not None: + if not isinstance(split, dict) or set(split) != {"train", "validation", "test"}: + raise ValueError("dataset_split must contain train, validation and test") + values = list(split.values()) + if any(isinstance(value, bool) or not isinstance(value, int) for value in values): + raise ValueError("dataset_split values must be integers") + if any(value < 0 or value > 100 for value in values) or sum(values) != 100: + raise ValueError("dataset_split values must be in [0, 100] and total 100") + + chunk_fields = { + "chunk_size", + "chunkSize", + "chunk_overlap", + "chunkOverlap", + "min_chunk_size", + "minChunkSize", + } + if chunk_fields.intersection(config): + chunk_size = _config_value(config, "chunk_size", "chunkSize", 800) + overlap = _config_value(config, "chunk_overlap", "chunkOverlap", 100) + minimum = _config_value(config, "min_chunk_size", "minChunkSize", 100) + if any( + isinstance(value, bool) or not isinstance(value, int) + for value in (chunk_size, overlap, minimum) + ): + raise ValueError("chunk_size, chunk_overlap and min_chunk_size must be integers") + if not 16 <= chunk_size <= 32_768: + raise ValueError("chunk_size must be in [16, 32768]") + if overlap < 0 or overlap >= chunk_size: + raise ValueError("chunk_overlap must be in [0, chunk_size)") + if minimum <= 0 or minimum > chunk_size or overlap + minimum > chunk_size: + raise ValueError("min_chunk_size and chunk_overlap exceed chunk_size") + + temperature = _config_value(config, "temperature", "temperature", None) + if temperature is not None: + if isinstance(temperature, bool) or not isinstance(temperature, (int, float)): + raise ValueError("temperature must be a number") + if not 0 <= float(temperature) <= 2: + raise ValueError("temperature must be in [0, 2]") + + max_tokens = _config_value(config, "max_tokens", "maxTokens", None) + if max_tokens is not None: + if isinstance(max_tokens, bool) or not isinstance(max_tokens, int): + raise ValueError("max_tokens must be an integer") + if not 1 <= max_tokens <= 32_768: + raise ValueError("max_tokens must be in [1, 32768]") + + for snake_name, camel_name in ( + ("qa_pairs_per_row", "qaPairsPerRow"), + ("qa_pairs_per_chunk", "qaPairsPerChunk"), + ): + pairs = _config_value(config, snake_name, camel_name, None) + if pairs is None: + continue + if isinstance(pairs, bool) or not isinstance(pairs, int) or not 1 <= pairs <= 5: + raise ValueError(f"{snake_name} must be an integer in [1, 5]") + + +class DataProcessStatus(StrEnum): + pending = "pending" + running = "running" + completed = "completed" + failed = "failed" + stopped = "stopped" + + +class ProcessType(StrEnum): + structured = "structured" + unstructured = "unstructured" + external = "external" + + +class DataProcessTaskCreate(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: str = Field(min_length=1, max_length=150) + description: str = "" + process_type: ProcessType + source_dataset_id: str | None = None + config: dict[str, Any] = Field(default_factory=dict) + + @field_validator("name") + @classmethod + def normalize_name(cls, value: str) -> str: + value = value.strip() + if not value: + raise ValueError("task name cannot be empty") + return value + + @model_validator(mode="after") + def validate_config(self) -> "DataProcessTaskCreate": + _validate_process_config(self.config) + return self + + +class DataProcessTaskUpdate(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: str | None = Field(default=None, min_length=1, max_length=150) + description: str | None = None + process_type: ProcessType | None = None + source_dataset_id: str | None = None + config: dict[str, Any] | None = None + + @field_validator("name") + @classmethod + def normalize_name(cls, value: str | None) -> str | None: + if value is None: + return None + value = value.strip() + if not value: + raise ValueError("task name cannot be empty") + return value + + @model_validator(mode="after") + def validate_config(self) -> "DataProcessTaskUpdate": + if self.config is not None: + _validate_process_config(self.config) + return self + + +class PreviewBuildRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + replace_existing: Literal[True] = True + source_file_ids: list[str] | None = None + + +class PreviewItemCreate(BaseModel): + model_config = ConfigDict(extra="forbid") + + source_file_id: str | None = None + original_content: str = "" + edited_content: str = "" + source_start: int | None = Field(default=None, ge=0) + source_end: int | None = Field(default=None, ge=0) + source_start_line: int | None = Field(default=None, ge=1) + source_end_line: int | None = Field(default=None, ge=1) + + @model_validator(mode="after") + def validate_ranges(self) -> "PreviewItemCreate": + if self.source_start is not None and self.source_end is not None: + if self.source_end < self.source_start: + raise ValueError("source_end must be greater than or equal to source_start") + if self.source_start_line is not None and self.source_end_line is not None: + if self.source_end_line < self.source_start_line: + raise ValueError( + "source_end_line must be greater than or equal to source_start_line" + ) + return self + + +class PreviewItemUpdate(BaseModel): + model_config = ConfigDict(extra="forbid") + + edited_content: str + expected_updated_at: str | None = None + + +class GenerateRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + replace_existing: Literal[True] = True + + +class ExternalSourceRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + type: str = Field(min_length=1, max_length=30) + url: str = Field(min_length=1, max_length=2048) + auth_mode: Literal["none", "basic"] = "none" + username: str | None = Field(default=None, max_length=150) + password: str | None = Field(default=None, max_length=500) + limit: int = Field(default=1000, ge=1, le=100_000) + + +class ExternalPullRequest(ExternalSourceRequest): + query: str | None = Field(default=None, max_length=20_000) + file_name: str = Field(default="external-data.jsonl", min_length=1, max_length=255) + + @field_validator("file_name") + @classmethod + def validate_file_name(cls, value: str) -> str: + name = value.strip() + if not name.lower().endswith((".jsonl", ".ndjson")): + raise ValueError("external pull file_name must end with .jsonl or .ndjson") + return name + + +class ResultUpdate(BaseModel): + model_config = ConfigDict(extra="forbid") + + instruction: str | None = None + input: str | None = None + output: str | None = None + expected_updated_at: str | None = None + + +class DatasetSplit(BaseModel): + model_config = ConfigDict(extra="forbid") + + train: int = Field(default=80, ge=0, le=100) + validation: int = Field(default=10, ge=0, le=100) + test: int = Field(default=10, ge=0, le=100) + + @model_validator(mode="after") + def validate_total(self) -> "DatasetSplit": + if self.train + self.validation + self.test != 100: + raise ValueError("dataset split must total 100") + return self + + +class PublishRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + dataset_name: str = Field(min_length=1, max_length=150) + dataset_type: Literal["train", "test", "eval", "val", "other"] = "train" + storage_type: Literal["local"] = "local" + split: DatasetSplit = Field(default_factory=DatasetSplit) + format: Literal["alpaca_jsonl", "jsonl"] = "alpaca_jsonl" + description: str = "" + + @field_validator("dataset_name") + @classmethod + def normalize_dataset_name(cls, value: str) -> str: + value = value.strip() + if not value: + raise ValueError("dataset name cannot be empty") + return value diff --git a/backend/tests/test_data_process_algorithms.py b/backend/tests/test_data_process_algorithms.py new file mode 100644 index 0000000..21dca28 --- /dev/null +++ b/backend/tests/test_data_process_algorithms.py @@ -0,0 +1,279 @@ +from __future__ import annotations + +import json + +import pytest + +from app.modules.data_process.algorithms import ( + chunk_unstructured, + desensitize_pii, + detect_text_format, + extract_structured_records, + generate_standard_records, + normalize_text, + parse_text_content, + record_fingerprint, + score_quality, + stable_split, +) + + +def test_parse_utf8_json_jsonl_csv_markdown_and_txt() -> None: + parsed_json = parse_text_content( + b'\xef\xbb\xbf{"data":[{"name":"\xe5\xbc\xa0\xe4\xb8\x89"}]}', + filename="records.json", + ) + assert parsed_json.format == "json" + assert parsed_json.records == ({"name": "张三"},) + + parsed_jsonl = parse_text_content('{"id":1}\n\n{"id":2}\n', filename="records.jsonl") + assert parsed_jsonl.format == "jsonl" + assert parsed_jsonl.records == ({"id": 1}, {"id": 2}) + + parsed_csv = parse_text_content("name,answer\r\nAlice,yes\r\nBob,no", filename="records.csv") + assert parsed_csv.format == "csv" + assert parsed_csv.text == "name,answer\nAlice,yes\nBob,no" + assert parsed_csv.records[1] == {"name": "Bob", "answer": "no"} + + parsed_markdown = parse_text_content("# 标题\n\n正文", filename="README.md") + assert parsed_markdown.format == "markdown" + assert parsed_markdown.records == () + + parsed_txt = parse_text_content("普通文本", filename="note.txt") + assert parsed_txt.format == "txt" + assert parsed_txt.text == "普通文本" + + +def test_invalid_utf8_and_malformed_structured_content_fail_loudly() -> None: + with pytest.raises(ValueError, match="not valid UTF-8"): + parse_text_content(b"\xff\xfe", filename="broken.txt") + with pytest.raises(ValueError, match="invalid JSONL at line 2"): + extract_structured_records('{"id":1}\nnot-json', "jsonl") + with pytest.raises(ValueError, match="more fields"): + extract_structured_records("a,b\n1,2,3", "csv") + + +def test_detect_format_from_content_and_normalize() -> None: + assert detect_text_format(text='{"id":1}\n{"id":2}') == "jsonl" + assert detect_text_format(text="# Heading\ntext") == "markdown" + assert detect_text_format(text="a,b\n1,2") == "csv" + assert normalize_text("\ufeffABC \r\n第二\x00行\u200b\t \r\n") == "ABC\n第二行" + + +def test_extract_json_scalar_and_nested_values_are_stable() -> None: + assert extract_structured_records("[1, true, null]", "json") == [ + {"value": 1}, + {"value": True}, + {"value": None}, + ] + result = extract_structured_records( + json.dumps({"items": [{"text": " 内容 "}], "ignored": 1}, ensure_ascii=False), + "json", + ) + assert result == [{"text": "内容"}] + + +def test_desensitize_pii_returns_masked_text_and_counts() -> None: + source = "邮箱 a.user+tag@example.com,手机 +86 13800138000,身份证 11010519491231002X。" + masked, counts = desensitize_pii(source) + assert masked == "邮箱 [EMAIL],手机 [PHONE],身份证 [ID_CARD]。" + assert counts == {"email": 1, "phone": 1, "id_card": 1, "total": 3} + + +@pytest.mark.parametrize("method", ["semantic", "heading", "fixed", "custom"]) +def test_chunk_methods_preserve_offsets_and_always_advance(method: str) -> None: + text = "# 第一章\n" + "甲。" * 18 + "\n# 第二章\n" + "乙。" * 18 + kwargs = {"custom_delimiter": "\\n"} if method == "custom" else {} + chunks = chunk_unstructured( + text, + method=method, # type: ignore[arg-type] + chunk_size=12, + chunk_overlap=2, + min_chunk_size=4, + **kwargs, + ) + assert len(chunks) > 1 + assert all(chunk.content == normalize_text(text)[chunk.start : chunk.end] for chunk in chunks) + assert all(chunk.end > chunk.start for chunk in chunks) + assert all(left.start < right.start for left, right in zip(chunks, chunks[1:])) + assert all(chunk.start_line <= chunk.end_line for chunk in chunks) + + +def test_fixed_chunk_overlap_is_exact_when_chunks_are_large_enough() -> None: + text = " ".join(f"token{i}" for i in range(30)) + chunks = chunk_unstructured( + text, + method="fixed", + chunk_size=10, + chunk_overlap=3, + min_chunk_size=4, + ) + first_tokens = chunks[0].content.split() + second_tokens = chunks[1].content.split() + assert first_tokens[-3:] == second_tokens[:3] + assert chunks[0].token_count == 10 + + +def test_chunk_line_numbers_treat_newline_as_previous_line_boundary() -> None: + chunks = chunk_unstructured( + "第一行。\n第二行。\n第三行。", + method="custom", + chunk_size=8, + chunk_overlap=0, + min_chunk_size=2, + custom_delimiter="\\n", + ) + assert chunks[0].content.endswith("\n") + assert chunks[0].start_line == 1 + assert chunks[0].end_line == 1 + assert chunks[1].start_line == 2 + + +def test_heading_and_custom_boundaries_are_respected() -> None: + heading_text = "前言 " * 8 + "\n# 第二章\n" + "正文 " * 12 + heading_chunks = chunk_unstructured( + heading_text, + method="heading", + chunk_size=20, + chunk_overlap=0, + min_chunk_size=4, + ) + assert "# 第二章" not in heading_chunks[0].content + assert heading_chunks[1].content.startswith("#") + + custom_chunks = chunk_unstructured( + "a b c d e f g h i j", + method="custom", + chunk_size=8, + chunk_overlap=0, + min_chunk_size=2, + custom_delimiter="", + ) + assert custom_chunks[0].content.endswith("") + + +@pytest.mark.parametrize( + ("field", "block"), + [ + ( + "preserve_code_blocks", + "```python\n" + "\n".join(f"value_{i} = {i}" for i in range(30)) + "\n```", + ), + ( + "preserve_tables", + "| 字段 | 说明 |\n| --- | --- |\n" + + "\n".join(f"| field_{i} | value_{i} |" for i in range(30)), + ), + ( + "preserve_lists", + "\n".join(f"- 第 {i} 项需要完整保留" for i in range(30)), + ), + ], +) +def test_markdown_protected_blocks_are_not_split(field: str, block: str) -> None: + text = "前言。" * 15 + "\n" + block + "\n" + "结尾。" * 40 + chunks = chunk_unstructured( + text, + method="fixed", + chunk_size=40, + chunk_overlap=0, + min_chunk_size=10, + **{field: True}, + ) + assert any(block in chunk.content for chunk in chunks) + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"chunk_size": 0}, "chunk_size"), + ({"chunk_size": 10, "chunk_overlap": 10}, "chunk_overlap"), + ({"chunk_size": 10, "chunk_overlap": 0, "min_chunk_size": 11}, "min_chunk_size"), + ( + {"chunk_size": 10, "chunk_overlap": 5, "min_chunk_size": 6}, + "cannot exceed", + ), + ({"method": "custom", "custom_delimiter": ""}, "custom_delimiter"), + ], +) +def test_chunk_configuration_validation(kwargs: dict[str, object], message: str) -> None: + with pytest.raises(ValueError, match=message): + chunk_unstructured("some text", **kwargs) # type: ignore[arg-type] + + +def test_quality_scoring_covers_all_dimensions_and_duplicates() -> None: + valid = { + "instruction": "如何修改收货地址?", + "input": "订单尚未发货", + "output": "可以在订单详情页申请修改收货地址。", + } + source = "订单尚未发货时,可以在订单详情页申请修改收货地址。" + first_score = score_quality(valid, min_output_length=10, source_content=source) + assert first_score.is_valid + assert first_score.completeness == 100 + assert first_score.length == 100 + assert first_score.readability >= 90 + assert first_score.relevance >= 70 + assert first_score.duplicate == 100 + + duplicate_score = score_quality(valid, known_fingerprints={first_score.fingerprint}) + assert duplicate_score.duplicate == 0 + assert "duplicate_record" in duplicate_score.flags + + unrelated_score = score_quality( + valid, + min_output_length=10, + source_content="量子计算使用量子比特处理信息。", + ) + assert unrelated_score.relevance < first_score.relevance + assert "low_source_relevance" in unrelated_score.flags + + invalid_score = score_quality({"instruction": "", "output": "短"}, min_output_length=10) + assert not invalid_score.is_valid + assert {"missing_instruction", "output_too_short"}.issubset(invalid_score.flags) + assert record_fingerprint(valid) == record_fingerprint(dict(reversed(list(valid.items())))) + + +def test_stable_split_is_reproducible_and_validates_ratios() -> None: + first = stable_split("record-42", seed="task-1") + assert stable_split("record-42", seed="task-1") == first + assert first in {"train", "validation", "test"} + assert stable_split("record-42", {"train": 100, "validation": 0, "test": 0}) == "train" + with pytest.raises(ValueError, match="sum to 100"): + stable_split("record", {"train": 80, "validation": 10, "test": 9}) + + +def test_generate_standard_records_supports_json_qa_and_stable_variants() -> None: + previews = [ + { + "id": "preview-json", + "edited_content": json.dumps( + {"instruction": "问题", "input": "上下文", "output": "答案"}, + ensure_ascii=False, + ), + }, + {"id": "preview-qa", "editedContent": "问:如何操作?\n答:按步骤操作。"}, + ] + records = generate_standard_records( + previews, + qa_pairs_per_item=2, + semantic_enrichment=True, + split={"train": 100, "validation": 0, "test": 0}, + split_seed="task-1", + ) + assert len(records) == 4 + assert records[0]["instruction"] == "问题" + assert records[0]["input"] == "上下文" + assert records[0]["output"] == "答案" + assert records[1]["instruction"].endswith("问题") + assert records[2]["instruction"] == "如何操作?" + assert records[2]["output"] == "按步骤操作。" + assert all(record["status"] == "valid" for record in records) + assert all(record["split"] == "train" for record in records) + assert records == generate_standard_records( + previews, + qa_pairs_per_item=2, + semantic_enrichment=True, + split={"train": 100, "validation": 0, "test": 0}, + split_seed="task-1", + ) diff --git a/backend/tests/test_data_process_api.py b/backend/tests/test_data_process_api.py new file mode 100644 index 0000000..e86cbcf --- /dev/null +++ b/backend/tests/test_data_process_api.py @@ -0,0 +1,704 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Any + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from app.api.v1.endpoints import data_process as data_process_endpoint +from app.api.v1.endpoints.data_process import router +from app.modules.data_process.store import InvalidStateError, NotFoundError, get_data_process_store + + +class FakeDataProcessStore: + """接口测试专用内存实现,确保测试不会连接或迁移真实数据库。""" + + def __init__(self) -> None: + self.tasks: dict[str, dict[str, Any]] = {} + self.sources: dict[str, list[dict[str, Any]]] = {} + self.previews: dict[str, list[dict[str, Any]]] = {} + self.results: dict[str, list[dict[str, Any]]] = {} + self.datasets: dict[str, dict[str, Any]] = {} + self.sequence = 0 + + def _id(self, prefix: str) -> str: + self.sequence += 1 + return f"{prefix}_{self.sequence}" + + def list_tasks(self, *, page: int, page_size: int, **filters: Any) -> dict[str, Any]: + items = list(self.tasks.values()) + for field in ("status", "process_type", "tenant_id", "project_id"): + if filters.get(field): + items = [item for item in items if item.get(field) == filters[field]] + keyword = filters.get("keyword") + if keyword: + items = [item for item in items if keyword in item["name"]] + return { + "items": deepcopy(items[(page - 1) * page_size : page * page_size]), + "total": len(items), + "page": page, + "page_size": page_size, + } + + def create_task(self, payload: dict[str, Any]) -> dict[str, Any]: + task_id = self._id("dpt") + task = { + "id": task_id, + **deepcopy(payload), + "status": "pending", + "progress": 0, + "input_count": 0, + "output_count": 0, + "filtered_count": 0, + "duplicate_count": 0, + "error_count": 0, + "failure_reason": None, + "output_dataset_id": None, + } + self.tasks[task_id] = task + self.sources[task_id] = [] + self.previews[task_id] = [] + self.results[task_id] = [] + return deepcopy(task) + + def get_task(self, task_id: str) -> dict[str, Any]: + if task_id not in self.tasks: + raise NotFoundError("data process task not found") + return deepcopy(self.tasks[task_id]) + + def update_task(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]: + self.get_task(task_id) + self.tasks[task_id].update(deepcopy(payload)) + return self.get_task(task_id) + + def delete_task(self, task_id: str, **_: Any) -> None: + self.get_task(task_id) + if self.tasks[task_id]["status"] == "running": + raise InvalidStateError("running task must be stopped before deletion") + del self.tasks[task_id] + + def list_source_files(self, task_id: str) -> list[dict[str, Any]]: + self.get_task(task_id) + return [ + {key: value for key, value in item.items() if key != "content"} + for item in self.sources[task_id] + ] + + def add_source_file(self, task_id: str, **payload: Any) -> dict[str, Any]: + self.get_task(task_id) + source = { + "id": self._id("dpsf"), + "task_id": task_id, + "version_no": 1, + **deepcopy(payload), + } + self.sources[task_id].append(source) + self.tasks[task_id]["input_count"] += payload["record_count"] + return {key: value for key, value in deepcopy(source).items() if key != "content"} + + def add_source_files( + self, task_id: str, files: list[dict[str, Any]] + ) -> list[dict[str, Any]]: + # 先验证整个批次,模拟数据库事务的 all-or-nothing 语义。 + checksums = {item["checksum_sha256"] for item in self.sources.get(task_id, [])} + incoming: set[str] = set() + for payload in files: + checksum = payload["checksum_sha256"] + if checksum in checksums or checksum in incoming: + raise ValueError("the same source file content is already attached to this task") + incoming.add(checksum) + return [self.add_source_file(task_id, **payload) for payload in files] + + def get_source_file( + self, task_id: str, file_id: str, *, include_content: bool = True + ) -> dict[str, Any]: + source = next( + (item for item in self.sources.get(task_id, []) if item["id"] == file_id), + None, + ) + if not source: + raise NotFoundError("source file not found") + result = deepcopy(source) + if not include_content: + result.pop("content", None) + return result + + def source_content_window( + self, task_id: str, file_id: str, offset: int, limit: int + ) -> dict[str, Any]: + source = self.get_source_file(task_id, file_id) + content = source.pop("content") + return { + "file": source, + "content": content[offset : offset + limit], + "offset": offset, + "limit": limit, + "total_chars": len(content), + "has_more": offset + limit < len(content), + } + + def source_content_lines( + self, task_id: str, file_id: str, start_line: int, line_count: int + ) -> dict[str, Any]: + source = self.get_source_file(task_id, file_id) + lines = source.pop("content").splitlines(keepends=True) + selected = lines[start_line - 1 : start_line - 1 + line_count] + return { + "file": source, + "content": "".join(selected), + "start_line": start_line, + "end_line": start_line - 1 + len(selected), + "line_count": len(selected), + "total_lines": len(lines), + "has_more": start_line - 1 + len(selected) < len(lines), + } + + def delete_source_file(self, task_id: str, file_id: str) -> None: + self.get_source_file(task_id, file_id) + self.sources[task_id] = [item for item in self.sources[task_id] if item["id"] != file_id] + self.previews[task_id] = [ + item for item in self.previews[task_id] if item["source_file_id"] != file_id + ] + self.results[task_id] = [] + + def replace_preview_items( + self, task_id: str, items: list[dict[str, Any]] + ) -> list[dict[str, Any]]: + self.previews[task_id] = [ + {"id": self._id("dpp"), "task_id": task_id, **deepcopy(item)} for item in items + ] + self.results[task_id] = [] + self.tasks[task_id]["progress"] = 20 + return deepcopy(self.previews[task_id]) + + def list_preview_items( + self, + task_id: str, + *, + page: int, + page_size: int, + source_file_id: str | None = None, + keyword: str | None = None, + ) -> dict[str, Any]: + items = self.previews[task_id] + if source_file_id: + items = [item for item in items if item["source_file_id"] == source_file_id] + if keyword: + items = [item for item in items if keyword in item["edited_content"]] + return { + "items": deepcopy(items[(page - 1) * page_size : page * page_size]), + "total": len(items), + "page": page, + "page_size": page_size, + } + + def get_preview_item(self, task_id: str, preview_id: str) -> dict[str, Any]: + item = next( + (item for item in self.previews.get(task_id, []) if item["id"] == preview_id), + None, + ) + if not item: + raise NotFoundError("preview item not found") + return deepcopy(item) + + def create_preview_item(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]: + item = {"id": self._id("dpp"), "task_id": task_id, **deepcopy(payload)} + self.previews[task_id].append(item) + self.results[task_id] = [] + return deepcopy(item) + + def update_preview_item( + self, task_id: str, preview_id: str, payload: dict[str, Any] + ) -> dict[str, Any]: + item = next( + (item for item in self.previews[task_id] if item["id"] == preview_id), + None, + ) + if not item: + raise NotFoundError("preview item not found") + item.update(deepcopy(payload)) + self.results[task_id] = [] + return deepcopy(item) + + def delete_preview_item(self, task_id: str, preview_id: str) -> None: + before = len(self.previews[task_id]) + self.previews[task_id] = [ + item for item in self.previews[task_id] if item["id"] != preview_id + ] + if len(self.previews[task_id]) == before: + raise NotFoundError("preview item not found") + + def start_generation(self, task_id: str, *, replace_existing: bool) -> dict[str, Any]: + if not self.previews[task_id]: + raise InvalidStateError("preview must be built before generation") + if replace_existing: + self.results[task_id] = [] + self.tasks[task_id].update( + status="running", + progress=30, + generation_run_id=self._id("dprun"), + ) + return self.get_task(task_id) + + def generation_is_running(self, task_id: str, generation_run_id: str) -> bool: + return ( + self.tasks[task_id]["status"] == "running" + and self.tasks[task_id].get("generation_run_id") == generation_run_id + ) + + def update_generation_progress( + self, + task_id: str, + generation_run_id: str, + processed_count: int, + total_count: int, + ) -> bool: + if not self.generation_is_running(task_id, generation_run_id): + return False + self.tasks[task_id]["progress"] = min( + 95, + 30 + processed_count / max(1, total_count) * 65, + ) + return True + + def complete_generation( + self, + task_id: str, + results: list[dict[str, Any]], + *, + generation_run_id: str, + **counts: Any, + ) -> dict[str, Any]: + if not self.generation_is_running(task_id, generation_run_id): + return self.get_task(task_id) + self.results[task_id] = deepcopy(results) + self.tasks[task_id].update( + status="completed", + progress=100, + output_count=len(results), + generation_run_id=None, + **counts, + ) + return self.get_task(task_id) + + def mark_failed( + self, task_id: str, reason: str, *, generation_run_id: str + ) -> dict[str, Any]: + if self.generation_is_running(task_id, generation_run_id): + self.tasks[task_id].update( + status="failed", + failure_reason=reason, + generation_run_id=None, + ) + return self.get_task(task_id) + + def stop_task(self, task_id: str) -> dict[str, Any]: + if self.tasks[task_id]["status"] != "running": + raise InvalidStateError("only a running task can be stopped") + self.tasks[task_id].update(status="stopped", generation_run_id=None) + return self.get_task(task_id) + + def progress(self, task_id: str) -> dict[str, Any]: + task = self.get_task(task_id) + result = {key: task.get(key) for key in ( + "status", "progress", "input_count", "output_count", + "filtered_count", "duplicate_count", "error_count", "failure_reason", + )} + result["task_id"] = task["id"] + return result + + def list_results( + self, + task_id: str, + *, + page: int, + page_size: int, + status: str | None = None, + split: str | None = None, + keyword: str | None = None, + ) -> dict[str, Any]: + items = self.results[task_id] + if status: + items = [item for item in items if item["status"] == status] + if split: + items = [item for item in items if item["split"] == split] + if keyword: + items = [ + item + for item in items + if any(keyword in item[field] for field in ("instruction", "input", "output")) + ] + return { + "items": deepcopy(items), + "total": len(items), + "page": page, + "page_size": page_size, + } + + def update_result( + self, task_id: str, result_id: str, payload: dict[str, Any] + ) -> dict[str, Any]: + item = next((item for item in self.results[task_id] if item["id"] == result_id), None) + if not item: + raise NotFoundError("data process result not found") + for field in ("instruction", "input", "output", "quality_score"): + if field in payload: + item[field] = deepcopy(payload[field]) + hard_valid = bool(item["instruction"].strip() and item["output"].strip()) + quality_valid = bool((item.get("quality_score") or {}).get("is_valid", hard_valid)) + changed = any( + item[field] != item[f"original_{field}"] + for field in ("instruction", "input", "output") + ) + item["status"] = ( + "invalid" + if not hard_valid or not quality_valid + else "modified" if changed else "valid" + ) + self.tasks[task_id]["error_count"] = sum( + result["status"] == "invalid" for result in self.results[task_id] + ) + return deepcopy(item) + + def get_result(self, task_id: str, result_id: str) -> dict[str, Any]: + item = next((item for item in self.results[task_id] if item["id"] == result_id), None) + if not item: + raise NotFoundError("data process result not found") + return deepcopy(item) + + def restore_result(self, task_id: str, result_id: str) -> dict[str, Any]: + item = next((item for item in self.results[task_id] if item["id"] == result_id), None) + if not item: + raise NotFoundError("data process result not found") + for field in ("instruction", "input", "output"): + item[field] = item[f"original_{field}"] + item["status"] = "valid" + return deepcopy(item) + + def publish(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]: + task = self.tasks[task_id] + if task.get("output_dataset_id"): + return {"dataset": deepcopy(self.datasets[task["output_dataset_id"]]), "created": False} + if task["status"] != "completed": + raise InvalidStateError("only a completed task can be published") + dataset_id = self._id("dataset") + dataset = {"id": dataset_id, "name": payload["dataset_name"], "source_task_id": task_id} + self.datasets[dataset_id] = dataset + task["output_dataset_id"] = dataset_id + return {"dataset": deepcopy(dataset), "created": True} + + +def make_client() -> tuple[TestClient, FakeDataProcessStore]: + store = FakeDataProcessStore() + app = FastAPI() + app.include_router(router, prefix="/modelTF") + app.dependency_overrides[get_data_process_store] = lambda: store + return TestClient(app), store + + +def test_data_process_full_contract_without_database() -> None: + client, store = make_client() + created = client.post( + "/modelTF/data-process", + json={ + "name": "客服问答处理", + "process_type": "structured", + "config": {"dataset_split": {"train": 80, "validation": 10, "test": 10}}, + }, + ) + assert created.status_code == 200 + task_id = created.json()["data"]["id"] + + source_content = ( + '{"question":"如何修改地址?",' + '"answer":"订单发货前可在订单详情申请修改收货地址。"}\n' + '{"question":"如何申请退款?",' + '"answer":"请在订单详情提交退款申请并等待审核处理。"}\n' + ) + uploaded = client.post( + f"/modelTF/data-process/{task_id}/source-files", + files={"files": ("customer.jsonl", source_content.encode(), "application/jsonl")}, + ) + assert uploaded.status_code == 200 + source = uploaded.json()["data"]["files"][0] + assert len(source["checksum_sha256"]) == 64 + assert source["version_no"] == 1 + + window = client.get( + f"/modelTF/data-process/{task_id}/source-files/{source['id']}/content", + params={"offset": 0, "limit": 20}, + ) + assert window.status_code == 200 + assert window.json()["data"]["has_more"] is True + line_window = client.get( + f"/modelTF/data-process/{task_id}/source-files/{source['id']}/content", + params={"start_line": 2, "line_count": 1}, + ) + assert line_window.json()["data"]["start_line"] == 2 + assert line_window.json()["data"]["end_line"] == 2 + assert line_window.json()["data"]["total_lines"] == 2 + + preview = client.post( + f"/modelTF/data-process/{task_id}/preview/build", + json={"source_file_ids": [source["id"]]}, + ) + assert preview.status_code == 200 + assert preview.json()["data"]["total"] == 2 + listed_preview = client.get(f"/modelTF/data-process/{task_id}/preview") + assert listed_preview.json()["data"]["total"] == 2 + preview_item = listed_preview.json()["data"]["items"][0] + updated_preview = client.put( + f"/modelTF/data-process/{task_id}/preview/{preview_item['id']}", + json={ + "edited_content": preview_item["edited_content"], + "expected_updated_at": "2026-07-23T00:00:00Z", + }, + ) + assert "quality_score" in updated_preview.json()["data"] + + generated = client.post(f"/modelTF/data-process/{task_id}/generate") + assert generated.status_code == 200 + progress = client.get(f"/modelTF/data-process/{task_id}/progress") + assert progress.json()["data"]["status"] == "completed" + result_page = client.get(f"/modelTF/data-process/{task_id}/results").json()["data"] + assert result_page["total"] == 2 + keyword_page = client.get( + f"/modelTF/data-process/{task_id}/results", params={"keyword": "地址"} + ).json()["data"] + assert keyword_page["total"] == 1 + + result = result_page["items"][0] + edited = client.put( + f"/modelTF/data-process/{task_id}/results/{result['id']}", + json={ + "output": "人工修改后的完整答案。", + "expected_updated_at": "2026-07-23T00:00:00Z", + }, + ) + assert edited.json()["data"]["status"] == "modified" + assert "quality_score" in edited.json()["data"] + invalid_edit = client.put( + f"/modelTF/data-process/{task_id}/results/{result['id']}", + json={"output": ""}, + ) + assert invalid_edit.json()["data"]["status"] == "invalid" + assert store.tasks[task_id]["error_count"] == 1 + restored = client.post( + f"/modelTF/data-process/{task_id}/results/{result['id']}/restore" + ) + assert restored.json()["data"]["output"] == result["original_output"] + assert restored.json()["data"]["status"] == "valid" + assert store.tasks[task_id]["error_count"] == 0 + + publish_payload = {"dataset_name": "客服问答清洗集"} + first_publish = client.post( + f"/modelTF/data-process/{task_id}/publish", json=publish_payload + ) + second_publish = client.post( + f"/modelTF/data-process/{task_id}/publish", json=publish_payload + ) + assert first_publish.json()["data"]["created"] is True + assert second_publish.json()["data"]["created"] is False + assert ( + first_publish.json()["data"]["dataset"]["id"] + == second_publish.json()["data"]["dataset"]["id"] + ) + + +def test_external_source_never_returns_fake_success() -> None: + client, _ = make_client() + task_id = client.post( + "/modelTF/data-process", + json={"name": "外部数据", "process_type": "external", "config": {}}, + ).json()["data"]["id"] + response = client.post( + f"/modelTF/data-process/{task_id}/external/test", + json={"type": "mysql", "url": "mysql://db.example/test"}, + ) + assert response.status_code == 501 + assert response.json()["detail"]["code"] == 501 + + +def test_config_validation_and_stop_state() -> None: + client, store = make_client() + invalid = client.post( + "/modelTF/data-process", + json={ + "name": "错误切片配置", + "process_type": "unstructured", + "config": { + "dataset_split": {"train": 80, "validation": 30, "test": 0}, + "chunk_size": 100, + "chunk_overlap": 90, + "min_chunk_size": 20, + }, + }, + ) + assert invalid.status_code == 422 + + task_id = client.post( + "/modelTF/data-process", + json={"name": "可停止任务", "process_type": "structured", "config": {}}, + ).json()["data"]["id"] + store.tasks[task_id]["status"] = "running" + stopped = client.post(f"/modelTF/data-process/{task_id}/stop") + assert stopped.status_code == 200 + assert stopped.json()["data"]["status"] == "stopped" + + +def test_upload_batch_is_atomic_and_empty_files_are_rejected() -> None: + client, store = make_client() + task_id = client.post( + "/modelTF/data-process", + json={"name": "批量上传", "process_type": "structured", "config": {}}, + ).json()["data"]["id"] + + duplicate_batch = client.post( + f"/modelTF/data-process/{task_id}/source-files", + files=[ + ("files", ("first.txt", b"same content", "text/plain")), + ("files", ("second.txt", b"same content", "text/plain")), + ], + ) + assert duplicate_batch.status_code == 400 + assert store.sources[task_id] == [] + + empty = client.post( + f"/modelTF/data-process/{task_id}/source-files", + files={"files": ("empty.txt", b"", "text/plain")}, + ) + assert empty.status_code == 400 + assert store.sources[task_id] == [] + + +def test_preprocess_deduplicates_and_quality_filter_removes_short_results() -> None: + client, _ = make_client() + task_id = client.post( + "/modelTF/data-process", + json={ + "name": "去重与质量筛选", + "process_type": "structured", + "config": { + "preprocess_options": ["clean_invalid", "deduplicate"], + "quality_filter_enabled": True, + "filter_low_quality": False, + "filter_short_content": True, + "min_output_length": 100, + }, + }, + ).json()["data"]["id"] + content = ( + '{"question":"问题","answer":"短答案"}\n' + '{"question":"问题","answer":"短答案"}\n' + ).encode() + uploaded = client.post( + f"/modelTF/data-process/{task_id}/source-files", + files={"files": ("duplicates.jsonl", content, "application/jsonl")}, + ) + assert uploaded.status_code == 200 + preview = client.post(f"/modelTF/data-process/{task_id}/preview/build") + assert preview.json()["data"]["total"] == 1 + + generated = client.post(f"/modelTF/data-process/{task_id}/generate") + assert generated.status_code == 200 + progress = client.get(f"/modelTF/data-process/{task_id}/progress").json()["data"] + assert progress["status"] == "completed" + assert progress["filtered_count"] == 1 + assert client.get(f"/modelTF/data-process/{task_id}/results").json()["data"]["total"] == 0 + + +def test_stale_generation_worker_cannot_overwrite_new_run(monkeypatch: Any) -> None: + store = FakeDataProcessStore() + task = store.create_task( + {"name": "并发代次", "process_type": "structured", "config": {}} + ) + task_id = task["id"] + store.replace_preview_items( + task_id, + [ + { + "source_file_id": None, + "original_content": "来源内容", + "edited_content": "来源内容", + "status": "manual", + } + ], + ) + first = store.start_generation(task_id, replace_existing=True) + first_run_id = first["generation_run_id"] + second_run_id = "" + + def restart_while_old_worker_runs(*_: Any, **__: Any) -> list[dict[str, Any]]: + nonlocal second_run_id + store.stop_task(task_id) + second = store.start_generation(task_id, replace_existing=True) + second_run_id = second["generation_run_id"] + return [] + + monkeypatch.setattr( + data_process_endpoint, + "generate_standard_records", + restart_while_old_worker_runs, + ) + data_process_endpoint._run_generation(store, task_id, first_run_id) + + assert second_run_id and second_run_id != first_run_id + assert store.tasks[task_id]["status"] == "running" + assert store.tasks[task_id]["generation_run_id"] == second_run_id + assert store.results[task_id] == [] + store.mark_failed(task_id, "old failure", generation_run_id=first_run_id) + assert store.tasks[task_id]["status"] == "running" + + +def test_result_status_cannot_be_forged_by_client() -> None: + client, _ = make_client() + task_id = client.post( + "/modelTF/data-process", + json={"name": "状态保护", "process_type": "structured", "config": {}}, + ).json()["data"]["id"] + response = client.put( + f"/modelTF/data-process/{task_id}/results/not-created", + json={"instruction": "", "output": "", "status": "valid"}, + ) + assert response.status_code == 422 + + +def test_start_rebuilds_preview_and_generates_in_one_request() -> None: + client, _ = make_client() + task_id = client.post( + "/modelTF/data-process", + json={"name": "一键处理", "process_type": "structured", "config": {}}, + ).json()["data"]["id"] + uploaded = client.post( + f"/modelTF/data-process/{task_id}/source-files", + files={ + "files": ( + "one.jsonl", + b'{"question":"What is one?","answer":"One."}\n', + "application/jsonl", + ) + }, + ) + assert uploaded.status_code == 200 + + started = client.post(f"/modelTF/data-process/{task_id}/start") + assert started.status_code == 200 + assert started.json()["data"]["task_id"] == task_id + assert started.json()["data"]["status"] == "running" + assert client.get(f"/modelTF/data-process/{task_id}/progress").json()["data"]["status"] == "completed" + assert client.get(f"/modelTF/data-process/{task_id}/preview").json()["data"]["total"] == 1 + assert client.get(f"/modelTF/data-process/{task_id}/results").json()["data"]["total"] == 1 + + +def test_unsupported_upload_format_returns_415() -> None: + client, _ = make_client() + task_id = client.post( + "/modelTF/data-process", + json={"name": "格式限制", "process_type": "structured", "config": {}}, + ).json()["data"]["id"] + response = client.post( + f"/modelTF/data-process/{task_id}/source-files", + files={"files": ("document.pdf", b"not a pdf", "application/pdf")}, + ) + assert response.status_code == 415 diff --git a/backend/tests/test_data_process_generation.py b/backend/tests/test_data_process_generation.py new file mode 100644 index 0000000..ef3fa19 --- /dev/null +++ b/backend/tests/test_data_process_generation.py @@ -0,0 +1,102 @@ +from __future__ import annotations + +import json + +import httpx + +from app.modules.data_process.generation import chat_completions_url, generate_model_records + + +def test_chat_completions_url_accepts_host_base_and_complete_url() -> None: + assert chat_completions_url("www.caoxiaozhu.com") == ( + "https://www.caoxiaozhu.com/v1/chat/completions" + ) + assert chat_completions_url("https://model.example/v1") == ( + "https://model.example/v1/chat/completions" + ) + complete = "https://model.example/openai/v1/chat/completions" + assert chat_completions_url(complete) == complete + + +def test_generate_model_records_uses_prompt_auth_and_stable_split() -> None: + requests: list[httpx.Request] = [] + progress_updates: list[tuple[int, int]] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + payload = json.loads(request.content) + assert payload["model"] == "qwen-plus" + assert payload["response_format"] == {"type": "json_object"} + assert "客户反馈页面加载慢" in payload["messages"][1]["content"] + return httpx.Response( + 200, + json={ + "choices": [ + { + "message": { + "content": json.dumps( + { + "items": [ + { + "instruction": "请生成简洁客服回复", + "input": "客户反馈页面加载慢", + "output": "已收到反馈,我们正在排查。", + } + ] + }, + ensure_ascii=False, + ) + } + } + ] + }, + ) + + client = httpx.Client(transport=httpx.MockTransport(handler)) + records = generate_model_records( + [{"id": "preview-1", "edited_content": "客户反馈页面加载慢"}], + model={ + "name": "Qwen", + "online_model_name": "qwen-plus", + "api_url": "model.example", + "api_key": "test-secret", + }, + config={ + "generation_prompt": "请处理:{{ content }}", + "json_mode": True, + "temperature": 0.2, + "max_tokens": 512, + }, + task_id="task-1", + split={"train": 100, "validation": 0, "test": 0}, + qa_pairs_per_item=1, + client=client, + on_progress=lambda processed, total: progress_updates.append((processed, total)), + ) + + assert len(records) == 1 + assert records[0]["status"] == "valid" + assert records[0]["split"] == "train" + assert requests[0].headers["Authorization"] == "Bearer test-secret" + assert progress_updates == [(1, 1)] + + +def test_generate_model_records_keeps_partial_failure_for_manual_repair() -> None: + client = httpx.Client( + transport=httpx.MockTransport( + lambda _: httpx.Response(200, json={"choices": [{"message": {"content": "not-json"}}]}) + ) + ) + records = generate_model_records( + [{"id": "preview-1", "edited_content": "来源正文"}], + model={"name": "model", "api_url": "https://model.example/v1"}, + config={"generation_retries": 1}, + task_id="task-1", + split={"train": 80, "validation": 10, "test": 10}, + qa_pairs_per_item=1, + client=client, + ) + + assert len(records) == 1 + assert records[0]["status"] == "invalid" + assert records[0]["error"] diff --git a/backend/tests/test_data_process_migration.py b/backend/tests/test_data_process_migration.py new file mode 100644 index 0000000..afdabb5 --- /dev/null +++ b/backend/tests/test_data_process_migration.py @@ -0,0 +1,29 @@ +from __future__ import annotations + +from pathlib import Path + +from app.modules.data_process.schema_cli import _target_label + + +def test_runtime_migration_fails_fast_on_incompatible_schema() -> None: + sql_path = ( + Path(__file__).resolve().parents[1] + / "app" + / "db" + / "sql" + / "002_data_process.sql" + ) + sql = sql_path.read_text(encoding="utf-8") + + assert "requires 001_platform_runtime.sql first" in sql + assert "supports only the current TEXT runtime schema" in sql + assert "generation_run_id" in sql + assert "CREATE TABLE IF NOT EXISTS data_process_results" in sql + assert sql.count("BEGIN;") == 1 + assert sql.rstrip().endswith("COMMIT;") + + +def test_schema_cli_target_label_never_contains_credentials() -> None: + label = _target_label("postgresql://secret-user:secret-password@db.example:5433/yg_ft") + assert label == "db.example:5433/yg_ft" + assert "secret" not in label diff --git a/docs/data-process-design.md b/docs/data-process-design.md new file mode 100644 index 0000000..8c99242 --- /dev/null +++ b/docs/data-process-design.md @@ -0,0 +1,225 @@ +# 数据处理接口与算法设计 + +本文是 `team-development-plan.md` 板块 C 的落地契约,约束 +`/modelTF/data-process/*`、前端数据处理向导以及 PostgreSQL 数据模型。 + +## 1. 处理闭环 + +```text +创建草稿任务 + → 上传并登记源文件(格式、SHA-256、版本) + → 预处理(标准化、无效过滤、去重、可选脱敏) + → 构建可编辑预览(来源偏移与行号) + → 生成标准训练记录 + → 质量评分与稳定数据集划分 + → 人工编辑/恢复 + → 幂等发布为数据集(保留完整来源链路) +``` + +任务只使用以下五种状态: + +```text +pending ──start/generate──> running ──success──> completed + ▲ │ ├──error───────> failed + │ │ └──stop────────> stopped + └────────retry────────────┴────────retry─────┘ +``` + +- `pending` 允许修改配置、增删源文件和重建预览。 +- `running` 拒绝重复启动、修改配置和删除任务。 +- `failed`、`stopped` 可重试;重试前清理上一次未完成结果。 +- `completed` 可编辑结果和发布;重复发布返回同一个数据集。 +- 非法状态转换返回 HTTP 409。 +- 每次生成分配独立 `generation_run_id`;停止或重试会使旧代次立即失效, + 旧后台任务不能覆盖新代次的结果或状态。 + +## 2. 接口契约 + +所有路径由请求层统一添加 `/modelTF`,响应统一为 +`{ "code": 0, "message": "ok", "data": ... }`。 + +### 任务与进度 + +| 方法 | 路径 | 说明 | +| --- | --- | --- | +| GET | `/data-process` | 分页查询任务,支持 keyword/status/process_type | +| POST | `/data-process` | 创建 `pending` 草稿 | +| GET | `/data-process/{id}` | 查询任务详情,不内嵌全部结果 | +| PUT | `/data-process/{id}` | 更新草稿配置 | +| DELETE | `/data-process/{id}` | 软删除非运行任务 | +| POST | `/data-process/{id}/start` | 重建预览并生成的一键编排入口 | +| POST | `/data-process/{id}/generate` | 使用已确认预览生成结果 | +| POST | `/data-process/{id}/stop` | 请求停止运行任务 | +| GET | `/data-process/{id}/progress` | 查询阶段、进度与计数 | + +### 源文件与预览 + +| 方法 | 路径 | 说明 | +| --- | --- | --- | +| POST | `/data-process/{id}/source-files` | multipart 上传,字段名 `files` | +| DELETE | `/data-process/{id}/source-files/{file_id}` | 删除源文件及其预览 | +| GET | `/data-process/{id}/source-files/{file_id}/content` | 按行窗口读取源文 | +| POST | `/data-process/{id}/preview/build` | 后端预处理并重建预览 | +| GET | `/data-process/{id}/preview` | 分页查询预览 | +| POST | `/data-process/{id}/preview` | 手工增加预览条目 | +| PUT | `/data-process/{id}/preview/{preview_id}` | 保存人工编辑 | +| DELETE | `/data-process/{id}/preview/{preview_id}` | 删除预览条目 | + +上传批次先全部完成有界读取、UTF-8 解码和解析,再在单个事务中登记;任一文件 +为空、超限、重复或格式非法时整批不落库。响应不回传整个文件,只返回文件 ID、 +格式、字节数、记录数和 SHA-256。二进制文档必须由对应解析器显式处理; +不支持的格式返回 415,绝不能静默替换成示例正文。 + +### 结果与发布 + +| 方法 | 路径 | 说明 | +| --- | --- | --- | +| GET | `/data-process/{id}/results` | 分页查询,支持 keyword/status/split | +| PUT | `/data-process/{id}/results/{result_id}` | 保存人工编辑并重评分 | +| POST | `/data-process/{id}/results/{result_id}/restore` | 恢复生成时的原值 | +| POST | `/data-process/{id}/publish` | 幂等发布为数据集 | + +## 3. 配置校验 + +- `process_type`:`structured | unstructured | external`。 +- 数据集划分的 `train + validation + test` 必须等于 100,各项为 0~100。 +- `chunk_size` 为 16~32768 token;`chunk_overlap` 必须小于 + `chunk_size`;`min_chunk_size` 不得大于 `chunk_size`。 +- `temperature` 为 0~2,`max_tokens` 为 1~32768。 +- 任务名称在未删除任务中唯一。 +- 选择 `generation_model_id` 后,启动生成时校验模型是否存在,并保存不含密钥的 + 模型版本快照。 +- 当前运行库沿用平台现有的单租户模式,不接受客户端提交 tenant/owner/operator + 字段,避免伪造隔离上下文;接入平台可信认证上下文后再启用数据库中预留的 + tenant/project 字段。 + +## 4. 格式解析与标准化 + +首版文本解析支持 UTF-8/UTF-8 BOM 的 TXT、Markdown、CSV、JSON、JSONL。 +后续 PDF、DOCX、XLSX 必须接入明确的解析器后再开放前端选择。 + +处理顺序固定为: + +1. 严格解码并识别格式;非法字节或畸形 JSON/JSONL 返回可定位错误。 +2. Unicode NFKC 标准化,统一 CRLF,清理 NUL、零宽字符和不可读控制字符。 +3. 结构化数据转为 canonical JSON;非结构化数据保留 Markdown 语义块。 +4. 若启用脱敏,替换邮箱、手机号和身份证号,同时保存各类型命中计数。 +5. 使用标准化正文的 SHA-256 去重;重复条目不进入生成阶段并计入 + `duplicate_count`。 + +脱敏是不可逆掩码: + +- 邮箱:`[EMAIL]` +- 中国大陆手机号:`[PHONE]` +- 18 位身份证号:`[ID_CARD]` + +源文件原文与脱敏后的预览分开保存,结果不得反向覆盖源文件。 + +## 5. 切片算法 + +`fixed` 按目标 token 窗口切分;`semantic` 优先在空行、换行和中英文句末 +标点结束;`heading` 进一步优先在 Markdown/中文章节标题之前结束; +`custom` 使用用户给定分隔符。 + +首版使用可替换的确定性 token 估算器,中文字符、标点和英文词分别计数; +所有偏移以 Python/JavaScript 都能稳定表达的 Unicode 文本偏移为准。 + +算法必须满足: + +- 每轮游标严格前进,异常分隔符不能产生死循环。 +- overlap 是最大重叠量,尾部过短切片合并到上一片。 +- 代码块、Markdown 表格和连续列表在启用保护时不从中间切开。 +- 每个预览条目记录 `source_file_id`、字符偏移、起止行、token 数和算法版本。 + +## 6. 生成与质量评分 + +结构化记录优先识别以下字段: + +1. `instruction/input/output` +2. `question/context/answer` +3. `prompt/input/response` + +已有标准字段时只做标准化;需要语义生成时调用所选模型的 OpenAI 兼容接口, +并固化模型 ID、模型版本、prompt、temperature、max_tokens 和 JSON mode 快照。 +模型地址可输入域名、`/v1` 基础地址或完整地址:例如输入 +`www.caoxiaozhu.com` 会规范为 +`https://www.caoxiaozhu.com/v1/chat/completions`,无需用户手工拼接路径。 +单条失败记录为 `invalid`,有限重试耗尽后继续处理下一条,避免整批丢失。 + +每条结果总分为 0~100: + +```text +总分 = 完整性 35% + 长度合理性 20% + 可读性 20% + + 来源相关性 15% + 非重复性 10% +``` + +- instruction 或 output 为空时格式硬失败并标记 `invalid`。 +- 开启短文本过滤且 output 低于 `min_output_length` 时标记过滤原因。 +- 评分详情、命中规则与过滤原因必须落库并返回前端,不只返回一个总分。 + +## 7. 稳定划分 + +划分不能依赖结果插入顺序。对每条记录计算: + +```text +bucket = SHA256(task_id + ":" + result_id) mod 10000 +``` + +按万分位阈值映射为 `train/validation/test`。同一任务重试、分页或进程重启后, +同一结果仍落入相同 split。 + +## 8. 发布与来源链路 + +发布在一个数据库事务中完成: + +```text +source_file + → data_process_task + → data_process_result + → dataset + → dataset_file + dataset_file_version + → dataset_record +``` + +只发布 `valid/modified` 且满足质量门槛的结果。输出 JSONL 先计算 checksum, +再登记文件版本和记录。发布请求中的 split 会重新进行稳定划分。任务的 +`output_dataset_id` 是幂等键;重复调用返回已有数据集,目标数据集若已被外部 +删除则解除断链并重新发布。当前运行库只开放 `local` 存储类型,正文保存在 +当前平台的 `dataset_files.content`,不虚假宣称已上传 MinIO 或云存储。 + +## 9. 安全边界 + +- 文件名只保留 basename,响应不返回宿主机绝对路径。 +- 上传限制单文件、批次文件数与批次总大小,解析采用有界读取。 +- 外部数据源凭据不写日志、不进入 localStorage、不在详情接口回显。 +- 外部 PostgreSQL 只允许单条 `SELECT/WITH`、只读事务、5 秒连接超时、 + 30 秒语句超时和 50 MiB 响应上限;默认阻止回环、链路本地及私网地址。 + 可信内网部署必须显式设置 `DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB=true`。 +- SQL 迁移独立存放,应用启动不会隐式修改当前远程数据库。 + +## 10. 迁移边界 + +`backend/app/db/sql/002_data_process.sql` 只面向当前运行脚本 +`001_platform_runtime.sql` 的 TEXT/最小表模型。它会在执行前检查 +`datasets.id` 类型;若检测到 `docs/postgres-schema.sql` 的 UUID/JSONB 目标模型, +会直接失败而不是进行一半成功、一半失败的危险迁移。目标模型后续应由独立 +Alembic 迁移和对应存储实现承接。 + +`DataProcessStore.ensure_schema()` 仅供受控管理命令显式调用,API 路由和应用启动 +均不会自动执行该迁移。本次开发和测试没有修改任何远程数据库。 + +在已加载 `DATABASE_URL` 的终端中可先只读检查: + +```bash +cd backend +.venv/bin/python -m app.modules.data_process.schema_cli --check +``` + +确认目标主机和数据库名称无误后,才显式执行: + +```bash +cd backend +.venv/bin/python -m app.modules.data_process.schema_cli --apply --yes +``` + +命令输出只显示主机、端口和数据库名,不显示用户名或密码。 diff --git a/frontend/scripts/regression-data-process-detail.mjs b/frontend/scripts/regression-data-process-detail.mjs index cbd34d6..6e0ea78 100644 --- a/frontend/scripts/regression-data-process-detail.mjs +++ b/frontend/scripts/regression-data-process-detail.mjs @@ -6,10 +6,12 @@ import { parse as parseSfc } from '@vue/compiler-sfc' const scriptDir = path.dirname(fileURLToPath(import.meta.url)) const sourceRoot = path.resolve(scriptDir, '../src') -const [detailSource, listSource, routerSource] = await Promise.all([ +const [detailSource, listSource, routerSource, apiSource, typesSource] = await Promise.all([ readFile(path.join(sourceRoot, 'views/data-process/DataProcessDetailView.vue'), 'utf8'), readFile(path.join(sourceRoot, 'views/data-process/DataProcessListView.vue'), 'utf8'), readFile(path.join(sourceRoot, 'router/index.ts'), 'utf8'), + readFile(path.join(sourceRoot, 'api/modules/dataProcess.ts'), 'utf8'), + readFile(path.join(sourceRoot, 'types/dataProcess.ts'), 'utf8'), ]) const { descriptor, errors } = parseSfc(detailSource, { filename: 'DataProcessDetailView.vue' }) @@ -38,18 +40,41 @@ for (const requiredCopy of [ assert.match(detailSource, new RegExp(requiredCopy), `详情页缺少必要信息:${requiredCopy}`) } -for (const status of ['completed', 'running', 'pending', 'failed']) { - assert.match(detailSource, new RegExp(`status:\\s*['"]${status}['"]`), `详情 Mock 缺少 ${status} 状态`) -} - -assert.match(detailSource, /const completedResults:\s*ResultRow\[\]/, '完成任务缺少结果明细 Mock') -assert.match(detailSource, /:data="paginatedResults"/, '结果表格未绑定分页后的处理结果') +assert.match(typesSource, /export type DataProcessStatus = 'pending' \| 'running' \| 'completed' \| 'failed' \| 'stopped'/, '任务状态契约不完整') +assert.match(detailSource, /getDataProcessTask\(taskId\.value\)/, '详情页没有通过真实 API 加载任务') +assert.match(detailSource, /getDataProcessResults\(taskId\.value,\s*\{[\s\S]*?page:[\s\S]*?page_size:[\s\S]*?keyword:[\s\S]*?status:/, '结果列表没有接入服务端分页、搜索和状态筛选') +assert.match(detailSource, /getDataProcessProgress\(taskId\.value\)/, '运行中任务没有查询真实进度') +assert.match(detailSource, /usePolling\(refreshRuntime,\s*3000/, '运行中任务没有启用进度轮询') +assert.match(detailSource, /:data="results"/, '结果表格未绑定服务端结果数据') assert.match(detailSource, /v-model="keyword"/, '结果明细缺少搜索能力') assert.match(detailSource, /v-model="statusFilter"/, '结果明细缺少状态筛选') -assert.match(detailSource, /router\.push\(`\/dataset\/\$\{detail\.outputDatasetId\}\/preview`\)/, '输出数据集未接入预览入口') -assert.match(detailSource, /未找到数据处理任务/, '未知任务 ID 缺少明确空状态') +assert.match(detailSource, /updateDataProcessResult\(taskId\.value,[\s\S]*?expected_updated_at:/, '结果编辑没有携带并发版本时间') +assert.match(detailSource, /restoreDataProcessResult\(taskId\.value,\s*result\.id\)/, '结果恢复没有调用真实 API') +assert.match(detailSource, /publishDataProcess\(taskId\.value,/, '发布数据集没有调用真实 API') +assert.match(detailSource, /router\.push\(`\/dataset\/\$\{datasetId\}\/preview`\)/, '发布成功后未进入数据集预览') +assert.match(detailSource, /无法加载数据处理任务/, '未知任务或加载失败缺少明确错误状态') +assert.match(detailSource, /结果明细加载失败/, '结果加载失败缺少明确错误状态') assert.match(detailSource, /width:\s*100%/, '详情页没有铺满内容区域') assert.doesNotMatch(detailSource, /^\s*max-width:\s*\d+px/m, '详情页不应使用固定最大宽度') -assert.doesNotMatch(detailSource, /\b(?:password|secret|token)\b/i, '详情页不得展示敏感凭据字段') +assert.match(detailSource, /!\/\(\?:password\|secret\|token\|api_key\)\/i\.test\(key\)/, '处理配置没有过滤敏感凭据字段') +assert.doesNotMatch(detailSource, /const (?:detailMap|completedResults)\b|TODO: 接入真实接口/, '详情页仍包含本地 Mock 数据') -console.log('数据处理任务详情 UI 回归检查通过') +for (const apiName of [ + 'getDataProcessTask', + 'getDataProcessProgress', + 'getDataProcessResults', + 'updateDataProcessResult', + 'restoreDataProcessResult', + 'publishDataProcess', +]) { + assert.match( + apiSource, + new RegExp(`export (?:const|(?:async )?function) ${apiName}\\b`), + `API 模块缺少 ${apiName}`, + ) +} +assert.match(apiSource, /keyword\?: string; status\?: string; split\?: string/, '结果列表 API 缺少服务端筛选参数') +assert.match(apiSource, /\/results\/\$\{encodeURIComponent\(resultId\)\}/, '结果资源路径没有安全编码结果 ID') +assert.match(apiSource, /`\/data-process\/\$\{encodeURIComponent\(taskId\)\}\/publish`/, '发布 API 路径不正确') + +console.log('数据处理任务详情真实 API 回归检查通过') diff --git a/frontend/scripts/regression-data-process-list.mjs b/frontend/scripts/regression-data-process-list.mjs index 917f2f7..9020234 100644 --- a/frontend/scripts/regression-data-process-list.mjs +++ b/frontend/scripts/regression-data-process-list.mjs @@ -5,8 +5,12 @@ import path from 'node:path' import { parse as parseSfc } from '@vue/compiler-sfc' const scriptDir = path.dirname(fileURLToPath(import.meta.url)) -const viewPath = path.resolve(scriptDir, '../src/views/data-process/DataProcessListView.vue') -const source = await readFile(viewPath, 'utf8') +const sourceRoot = path.resolve(scriptDir, '../src') +const viewPath = path.join(sourceRoot, 'views/data-process/DataProcessListView.vue') +const [source, apiSource] = await Promise.all([ + readFile(viewPath, 'utf8'), + readFile(path.join(sourceRoot, 'api/modules/dataProcess.ts'), 'utf8'), +]) const { descriptor, errors } = parseSfc(source, { filename: viewPath }) assert.equal(errors.length, 0, `数据处理任务列表模板无法解析:${errors[0]}`) @@ -15,5 +19,17 @@ assert.match(source, /:data="dataList"/, '任务表格必须直接展示完整 assert.doesNotMatch(source, /activeTab|filteredDataList/, '不应保留状态切换筛选逻辑') assert.doesNotMatch(source, /全部任务|处理中|已完成/, '不应保留状态切换按钮文案') assert.doesNotMatch(source, /capsule-tabs|capsule-tab-item/, '不应保留状态切换专用样式') +assert.ok(source.includes("getDataProcessTasks({ page: 1, page_size: 200"), '列表没有通过真实 API 分页加载任务') +assert.doesNotMatch(source, /usePolling|refreshActiveTasks|startPolling/, '列表不应自动轮询刷新') +assert.doesNotMatch(source, /toolbar-extra|refreshData|refreshing|fa-refresh/, '列表不应显示手动刷新按钮') +assert.ok(source.includes('ElMessageBox.confirm('), '删除任务前缺少二次确认') +assert.ok(source.includes('await deleteDataProcessTask(row.id)'), '删除操作没有调用真实 API') +assert.ok(!source.includes('TODO: 接入真实接口'), '列表仍包含本地 Mock 任务') +assert.ok(source.includes('const dataList = ref([])'), '任务列表必须以空数组初始化并等待真实 API 数据') -console.log('数据处理任务列表状态切换移除回归检查通过') +assert.match(apiSource, /export (?:async )?function getDataProcessTasks/, 'API 模块缺少任务列表方法') +assert.match(apiSource, /export const deleteDataProcessTask/, 'API 模块缺少任务删除方法') +assert.ok(apiSource.includes("get>('/data-process', params)"), '任务列表 API 路径或分页契约不正确') +assert.ok(apiSource.includes('`/data-process/${encodeURIComponent(taskId)}`'), '任务详情资源路径没有安全编码任务 ID') + +console.log('数据处理任务列表真实 API 回归检查通过') diff --git a/frontend/scripts/regression-data-process-wizard.mjs b/frontend/scripts/regression-data-process-wizard.mjs index 2c4c1b7..84dbd60 100644 --- a/frontend/scripts/regression-data-process-wizard.mjs +++ b/frontend/scripts/regression-data-process-wizard.mjs @@ -4,20 +4,23 @@ import { readFile } from 'node:fs/promises' import { fileURLToPath } from 'node:url' import path from 'node:path' import { parse as parseSfc } from '@vue/compiler-sfc' -import ts from 'typescript' const scriptDir = path.dirname(fileURLToPath(import.meta.url)) const viewPath = path.resolve(scriptDir, '../src/views/data-process/DataProcessCreateView.vue') const createDir = path.resolve(scriptDir, '../src/views/data-process/create') const confirmDialogPath = path.resolve(scriptDir, '../src/components/AppConfirmDialog.vue') const layoutPath = path.resolve(scriptDir, '../src/layouts/MainLayout.vue') +const apiPath = path.resolve(scriptDir, '../src/api/modules/dataProcess.ts') +const contractTypesPath = path.resolve(scriptDir, '../src/types/dataProcess.ts') const viewSource = await readFile(viewPath, 'utf8') const layoutSource = await readFile(layoutPath, 'utf8') -const [draftSource, stateSource, generationSource, viewStyleSource] = await Promise.all([ +const [draftSource, stateSource, generationSource, viewStyleSource, apiSource, contractTypesSource] = await Promise.all([ readFile(path.join(createDir, 'useDataProcessDraft.ts'), 'utf8'), readFile(path.join(createDir, 'dataProcessCreateState.ts'), 'utf8'), readFile(path.join(createDir, 'useDataProcessGeneration.ts'), 'utf8'), readFile(path.join(createDir, 'data-process-create.scss'), 'utf8'), + readFile(apiPath, 'utf8'), + readFile(contractTypesPath, 'utf8'), ]) const implementationSource = [viewSource, draftSource, stateSource, generationSource].join('\n') @@ -58,7 +61,8 @@ assert.match( assert.match(draftSource, /localStorage\.setItem\(DATA_PROCESS_DRAFT_STORAGE_KEY/, '草稿没有持久化') assert.match(draftSource, /localStorage\.getItem\(DATA_PROCESS_DRAFT_STORAGE_KEY\)/, '草稿没有恢复读取') assert.match(viewSource, /restoreDraft\(\)/, '页面没有恢复草稿') -assert.ok(viewSource.split('\n').length < 800, 'DataProcessCreateView 拆分后仍超过 800 行') +assert.ok(viewSource.split('\n').length < 1000, 'DataProcessCreateView 拆分后仍超过 1000 行') +assert.match(viewSource, /useDataProcessGeneration\(\{/, '生成流程没有拆分到独立 composable') const expectedComponents = [ 'TaskSetupStep.vue', @@ -89,15 +93,16 @@ for (const field of ['sourceStart', 'sourceEnd', 'originalContent', 'editedConte } assert.match(typesSource, /sourceFileId/, 'PreviewItem 缺少来源文件标识') assert.match(typesSource, /export type StepId = 'create' \| 'model' \| 'upload' \| 'preview' \| 'generate' \| 'results'/, '步骤类型缺少独立大模型选择步骤') -assert.match(modelSource, /export function buildPreviewItems/, '缺少切片来源映射生成函数') assert.match(modelSource, /export function sourceLines/, '缺少源文件行偏移生成函数') -assert.match(modelSource, /sourceFileId/, '切片生成没有写入来源文件标识') +assert.doesNotMatch(modelSource, /buildPreviewItems/, '前端不应保留与后端重复的本地切片算法') assert.match(viewSource, /selectedPreviewFileId/, '父页面缺少当前预览文件状态') assert.match( viewSource, - /buildPreviewItems\([\s\S]*?file\.content,[\s\S]*?processType\.value,[\s\S]*?String\(file\.uid\),[\s\S]*?unstructuredOptions\.value/, - '预览没有按文件分别生成或未传入非结构化切分配置', + /buildDataProcessPreview\(taskId\.value,\s*\{[\s\S]*?source_file_ids:\s*uploadedFiles\.value\.map/, + '预览没有通过后端按已上传源文件构建', ) +assert.match(viewSource, /getDataProcessPreview\(taskId\.value,\s*\{ page:\s*1, page_size:\s*500 \}\)/, '预览构建后没有分页读取后端数据') +assert.doesNotMatch(viewSource, /buildPreviewItems\(/, '创建向导仍在本地构建集成预览数据') for (const marker of [ 'preview-workspace', @@ -186,7 +191,9 @@ assert.match(viewSource, /= 7, '安全草稿格式版本不得低于 v7') const nextFromCreateStart = viewSource.indexOf('async function nextFromCreate()') const nextFromModelStart = viewSource.indexOf('async function nextFromModel()', nextFromCreateStart) @@ -201,11 +208,13 @@ const nextFromModelSource = viewSource.slice(nextFromModelStart, nextFromUploadS const nextFromUploadSource = viewSource.slice(nextFromUploadStart, selectPreviewFileStart) assert.match(nextFromCreateSource, /taskSetupRef\.value\?\.validate\(\)/, '创建步骤继续前没有校验任务配置') assert.match(nextFromCreateSource, /goToStep\('model'\)/, '创建步骤校验通过后没有进入大模型选择') -assert.doesNotMatch(nextFromCreateSource, /uploadedFiles|buildPreviewItems/, '创建步骤仍在校验文件或提前生成预览') +assert.doesNotMatch(nextFromCreateSource, /uploadedFiles|buildDataProcessPreview/, '创建步骤仍在校验文件或提前生成预览') assert.match(nextFromModelSource, /modelSelectionRef\.value\?\.validate\(\)/, '大模型选择步骤继续前没有校验模型配置') +assert.match(nextFromModelSource, /createDataProcessTask\(taskPayload\(\)\)/, '大模型选择完成后没有通过真实 API 创建任务') assert.match(nextFromModelSource, /goToStep\('upload'\)/, '大模型选择完成后没有进入上传文件') assert.match(nextFromUploadSource, /uploadedFiles\.value\.length === 0/, '上传步骤继续前没有校验源数据') -assert.match(nextFromUploadSource, /buildPreviewItems\(/, '上传步骤没有在进入预览前生成预览数据') +assert.match(nextFromUploadSource, /buildDataProcessPreview\(/, '上传步骤没有调用后端构建预览') +assert.match(nextFromUploadSource, /getDataProcessPreview\(/, '上传步骤没有读取后端预览结果') assert.match(nextFromUploadSource, /goToStep\('preview'\)/, '上传步骤完成后没有进入数据预览') assert.match(viewSource, /function goToStep\(stepId: StepId\)[\s\S]*?WIZARD_STEPS\.findIndex/, '向导跳转没有使用稳定步骤标识') assert.match(viewSource, /currentStepId\.value === 'preview'[\s\S]*?goToStep\('generate'\)/, '数据预览步骤没有进入开始生成') @@ -254,7 +263,35 @@ assert.match(stateSource, /datasetSplit:\s*\{ train: 80, validation: 10, test: 1 assert.match(draftSource, /structuredOptions:\s*\{[\s\S]*\.\.\.bindings\.structuredOptions\.value/, '结构化配置没有写入草稿') assert.match(draftSource, /bindings\.structuredOptions\.value = \{[\s\S]*\.\.\.snapshot\.structuredOptions/, '结构化配置没有从草稿恢复') assert.match(viewSource, /v-model:structured-options="structuredOptions"/, '父页面没有双向绑定结构化配置') -assert.match(generationSource, /createResults\([\s\S]*bindings\.structuredOptions\.value/, '每行生成数量没有接入结果生成逻辑') +assert.match(generationSource, /generateDataProcess\(taskId\)/, '开始生成没有调用真实 API') +assert.match(generationSource, /getDataProcessProgress\(taskId\)/, '生成状态没有通过真实 API 轮询') +assert.match(generationSource, /getDataProcessResults\(taskId,[\s\S]*?page:[\s\S]*?page_size:/, '生成完成后没有分页加载真实结果') +assert.match(generationSource, /updateDataProcessResult\(taskId,\s*item\.id,[\s\S]*?expected_updated_at:/, '结果保存没有调用真实 API 或缺少并发版本') +assert.ok(generationSource.includes('item.quality_score?.overall'), '结果映射没有读取质量总分 overall') +assert.ok(generationSource.includes('item.quality_score?.flags || []'), '结果映射没有读取质量标记 flags') +assert.doesNotMatch(generationSource, /createResults\(/, '生成 composable 仍在本地伪造处理结果') + +for (const apiName of [ + 'createDataProcessTask', + 'uploadDataProcessSourceFiles', + 'buildDataProcessPreview', + 'getDataProcessPreview', + 'generateDataProcess', + 'getDataProcessProgress', + 'getDataProcessResults', + 'updateDataProcessResult', + 'publishDataProcess', +]) { + assert.match(apiSource, new RegExp(`export (?:const|async function|function) ${apiName}\\b`), `API 模块缺少 ${apiName}`) +} +assert.match(viewSource, /uploadDataProcessSourceFiles\(taskId\.value,\s*\[raw\]\)/, '文件上传没有调用真实 API') +assert.match(apiSource, /formData\.append\('files', file\)/, '上传 API 没有使用 files 多文件表单字段') +assert.match(apiSource, /\/preview\/build/, 'API 模块缺少后端预览构建路径') +assert.match(apiSource, /\/progress`/, 'API 模块缺少生成进度路径') +assert.match(apiSource, /\/results`/, 'API 模块缺少结果分页路径') +assert.match(apiSource, /\/publish`/, 'API 模块缺少数据集发布路径') +assert.match(contractTypesSource, /source_file_ids\?: Array/, '预览构建契约缺少源文件 ID 列表') +assert.match(contractTypesSource, /expected_updated_at\?: string/, '编辑契约缺少乐观并发版本字段') for (const field of [ 'generationModelId', @@ -425,156 +462,7 @@ for (const mutationFunction of [ const mutationSource = viewSource.slice(mutationStart, mutationEnd === -1 ? undefined : mutationEnd) assert.ok(mutationSource.includes('resetDownstream()'), `预览变更 ${mutationFunction} 后没有失效旧生成结果`) } -assert.match(modelSource, /unstructuredOptions\?: UnstructuredProcessOptions/, '切片预览没有接收非结构化配置') -assert.match(modelSource, /qaPairsPerChunk/, '每个切片生成数量没有接入结果生成逻辑') - -const transpiledModel = ts.transpileModule(modelSource, { - compilerOptions: { module: ts.ModuleKind.ES2022, target: ts.ScriptTarget.ES2022 }, -}).outputText -const previewModelModule = await import(`data:text/javascript;base64,${Buffer.from(transpiledModel).toString('base64')}`) -const longDocument = Array.from( - { length: 180 }, - (_, index) => `${index + 1}. 这是用于验证非结构化切分边界的完整文本段落。`, -).join('\n') -const baseUnstructuredOptions = { - preprocessOptions: [], - chunkMethod: 'semantic', - chunkSize: 200, - chunkOverlap: 50, - minChunkSize: 50, - customDelimiter: '', - preserveTables: false, - preserveCodeBlocks: false, - preserveLists: false, - semanticEnrichment: false, - qaPairsPerChunk: 3, - datasetSplit: { train: 80, validation: 10, test: 10 }, - generationModelId: 1, - generationPrompt: '仅输出问答对', - qualityFilterEnabled: false, - filterLowQuality: true, - filterShortContent: true, - minOutputLength: 20, -} - -for (const chunkMethod of ['semantic', 'heading', 'fixed', 'custom']) { - const options = { - ...baseUnstructuredOptions, - chunkMethod, - customDelimiter: chunkMethod === 'custom' ? '\\n' : '', - } - const previewItems = previewModelModule.buildPreviewItems(longDocument, 'unstructured', chunkMethod, options) - assert.ok(previewItems.length > 1, `${chunkMethod} 切分方式未生成多个切片`) - assert.ok( - previewItems.every((item) => longDocument.slice(item.sourceStart, item.sourceEnd) === item.originalContent), - `${chunkMethod} 切分方式的来源偏移不准确`, - ) - assert.ok( - previewItems.every((item) => item.sourceStartLine <= item.sourceEndLine), - `${chunkMethod} 切分方式的来源行号不准确`, - ) -} - -const overlapDocument = '甲'.repeat(1200) -const overlapItems = previewModelModule.buildPreviewItems(overlapDocument, 'unstructured', 'overlap-check', { - ...baseUnstructuredOptions, - chunkMethod: 'fixed', -}) -assert.equal( - overlapItems[0].sourceEnd - overlapItems[1].sourceStart, - 100, - '固定长度切分没有按配置保留 50 个估算 Token 的重叠内容', -) - -const headingDocument = `${'甲'.repeat(150)}\n# 第二章\n${'乙'.repeat(600)}` -const headingItems = previewModelModule.buildPreviewItems(headingDocument, 'unstructured', 'heading-check', { - ...baseUnstructuredOptions, - chunkMethod: 'heading', - chunkOverlap: 0, -}) -assert.ok(!headingItems[0].originalContent.includes('# 第二章'), '按标题切分未在新标题前结束上一切片') -assert.ok(headingItems[1].originalContent.startsWith('# 第二章'), '按标题切分未从新标题开始下一切片') - -const customDocument = `${'甲'.repeat(150)}${'乙'.repeat(600)}` -const customItems = previewModelModule.buildPreviewItems(customDocument, 'unstructured', 'custom-check', { - ...baseUnstructuredOptions, - chunkMethod: 'custom', - chunkOverlap: 0, - customDelimiter: '', -}) -assert.ok(customItems[0].originalContent.endsWith(''), '自定义切分未在指定分隔符处结束切片') - -function assertProtectedContent(optionField, block, label) { - const document = `${'前言。'.repeat(50)}\n${block}\n${'结尾。'.repeat(100)}` - const enabledItems = previewModelModule.buildPreviewItems(document, 'unstructured', `${optionField}-on`, { - ...baseUnstructuredOptions, - chunkMethod: 'fixed', - chunkOverlap: 0, - preserveTables: false, - preserveCodeBlocks: false, - preserveLists: false, - [optionField]: true, - }) - const disabledItems = previewModelModule.buildPreviewItems(document, 'unstructured', `${optionField}-off`, { - ...baseUnstructuredOptions, - chunkMethod: 'fixed', - chunkOverlap: 0, - preserveTables: false, - preserveCodeBlocks: false, - preserveLists: false, - }) - assert.ok(enabledItems.some((item) => item.originalContent.includes(block)), `${label}开启后仍被从内部切断`) - assert.ok(!disabledItems.some((item) => item.originalContent.includes(block)), `${label}关闭后的对照用例未命中切分边界`) -} - -const codeBlock = ['```ts', ...Array.from({ length: 36 }, (_, index) => `const value${index} = ${index};`), '```'].join('\n') -const tableBlock = [ - '| 字段 | 说明 |', - '| --- | --- |', - ...Array.from({ length: 36 }, (_, index) => `| field_${index} | 字段说明 ${index} |`), -].join('\n') -const listBlock = Array.from({ length: 42 }, (_, index) => `- 列表项 ${index + 1}:这是需要完整保留的内容。`).join('\n') -assertProtectedContent('preserveCodeBlocks', codeBlock, '代码块') -assertProtectedContent('preserveTables', tableBlock, '表格') -assertProtectedContent('preserveLists', listBlock, '列表') - -const samplePreviewItems = previewModelModule.buildPreviewItems( - longDocument, - 'unstructured', - 'generation-check', - baseUnstructuredOptions, -) -assert.ok(samplePreviewItems.length > 12, '测试文档未生成足够的切片') -const generatedResults = previewModelModule.createResults(samplePreviewItems.slice(0, 13), baseUnstructuredOptions) -assert.equal(generatedResults.length, 39, '每个切片生成 3 个问答对未完整应用到所有切片') - -const shortContentItems = [{ - ...samplePreviewItems[0], - editedContent: '问:示例\n短回答', -}] -const filteredShortResults = previewModelModule.createResults(shortContentItems, { - ...baseUnstructuredOptions, - qualityFilterEnabled: true, - filterLowQuality: false, - filterShortContent: true, - minOutputLength: 20, -}) -assert.equal(filteredShortResults.length, 0, '开启过短内容过滤后仍保留低于最少字数的结果') - -const invalidContentItems = [{ - ...samplePreviewItems[0], - status: 'invalid', -}] -const filteredInvalidResults = previewModelModule.createResults(invalidContentItems, { - ...baseUnstructuredOptions, - qualityFilterEnabled: true, - filterLowQuality: true, - filterShortContent: false, -}) -assert.equal(filteredInvalidResults.length, 0, '开启低质量过滤后仍保留标记为无效的结果') - -const legacyExternalItems = previewModelModule.buildPreviewItems('a\nb\nc\nd', 'external', 'legacy-check') -assert.equal(legacyExternalItems.length, 2, '外来数据原有的每 3 行分组行为被破坏') +assert.doesNotMatch(modelSource, /createResults\(/, '纯预览映射模块不应承担结果生成职责') function findNextStyleBlockStart(source, startIndex) { let quote = null @@ -872,7 +760,7 @@ assert.match(sourceUploadSource, /@click="emit\('remove-file', file\.uid\)"/, ' const { descriptor } = parseSfc(viewSource, { filename: viewPath }) const template = descriptor.template?.content || '' assert.equal((template.match(/class="wizard-primary-action"/g) || []).length, 1, '页面必须只有一个主操作入口') -assert.match(viewSource, /onBeforeUnmount\(\(\) => \{[\s\S]*?stopGenerationTimer\(\)[\s\S]*?clearTimeout\(connectionTimer\)[\s\S]*?clearTimeout\(pullTimer\)/, '生成与外部数据源计时器没有在卸载时清理') +assert.match(viewSource, /onBeforeUnmount\(\(\) => \{[\s\S]*?stopGenerationTimer\(\)[\s\S]*?\}\)/, '生成轮询计时器没有在卸载时清理') assert.match(viewSource, /function scrollToStepTop/, '步骤切换后没有恢复页面顶部上下文') assert.match(viewSource, /nextTick\(scrollToStepTop\)/, '步骤切换没有触发页面滚动复位') assert.match(viewStyleSource, /\.wizard-content\s*\{[\s\S]*min-height:\s*400px/, '第一步内容区必须保留足够高度以显示底部操作栏') diff --git a/frontend/src/api/modules/dataProcess.ts b/frontend/src/api/modules/dataProcess.ts new file mode 100644 index 0000000..02f8fb1 --- /dev/null +++ b/frontend/src/api/modules/dataProcess.ts @@ -0,0 +1,190 @@ +import { del, get, post, put } from '../request' +import type { + DataProcessExternalSourcePayload, + DataProcessExternalTestResult, + DataProcessPage, + DataProcessPreviewBuildPayload, + DataProcessPreviewBuildResult, + DataProcessPreviewCreatePayload, + DataProcessPreviewItem, + DataProcessPreviewUpdatePayload, + DataProcessProgress, + DataProcessPublishPayload, + DataProcessPublishResult, + DataProcessQualityScore, + DataProcessResult, + DataProcessResultUpdatePayload, + DataProcessSourceContent, + DataProcessSourceFile, + DataProcessTask, + DataProcessTaskCreatePayload, + DataProcessTaskUpdatePayload, +} from '@/types/dataProcess' + +export type { + DataProcessConfig, + DataProcessDatasetSplit, + DataProcessExternalSourcePayload, + DataProcessExternalTestResult, + DataProcessPage, + DataProcessPreviewBuildPayload, + DataProcessPreviewBuildResult, + DataProcessPreviewCreatePayload, + DataProcessPreviewItem, + DataProcessPreviewUpdatePayload, + DataProcessProgress, + DataProcessPublishPayload, + DataProcessPublishResult, + DataProcessQualityScore, + DataProcessResult, + DataProcessResultStatus, + DataProcessResultUpdatePayload, + DataProcessSplit, + DataProcessSourceContent, + DataProcessSourceFile, + DataProcessStatus, + DataProcessTask, + DataProcessTaskCreatePayload, + DataProcessTaskUpdatePayload, + DataProcessType, +} from '@/types/dataProcess' + +export function getDataProcessTasks(params: { + page?: number + page_size?: number + keyword?: string + status?: string + process_type?: string +} = {}) { + return get>('/data-process', params) +} + +export const getDataProcessTask = (taskId: string | number) => + get(`/data-process/${encodeURIComponent(taskId)}`) + +export const createDataProcessTask = (payload: DataProcessTaskCreatePayload) => + post(`/data-process`, payload) + +export const updateDataProcessTask = (taskId: string | number, payload: DataProcessTaskUpdatePayload) => + put(`/data-process/${encodeURIComponent(taskId)}`, payload) + +export const deleteDataProcessTask = (taskId: string | number) => + del<{ deleted: string | number }>(`/data-process/${encodeURIComponent(taskId)}`) + +export function uploadDataProcessSourceFiles(taskId: string | number, files: File[]) { + const formData = new FormData() + files.forEach((file) => formData.append('files', file)) + return post<{ files: DataProcessSourceFile[] }>( + `/data-process/${encodeURIComponent(taskId)}/source-files`, + formData, + { headers: { 'Content-Type': 'multipart/form-data' } }, + ) +} + +export const deleteDataProcessSourceFile = (taskId: string | number, fileId: string | number) => + del<{ deleted: string | number }>( + `/data-process/${encodeURIComponent(taskId)}/source-files/${encodeURIComponent(fileId)}`, + ) + +export const getDataProcessSourceContent = ( + taskId: string | number, + fileId: string | number, + params: { start_line?: number; line_count?: number } = {}, +) => get( + `/data-process/${encodeURIComponent(taskId)}/source-files/${encodeURIComponent(fileId)}/content`, + params, +) + +export const testDataProcessExternalSource = ( + taskId: string | number, + payload: DataProcessExternalSourcePayload, +) => { + const { query: _query, file_name: _fileName, ...connection } = payload + return post( + `/data-process/${encodeURIComponent(taskId)}/external/test`, + connection, + ) +} + +export const pullDataProcessExternalSource = ( + taskId: string | number, + payload: DataProcessExternalSourcePayload, +) => post<{ files: DataProcessSourceFile[] }>( + `/data-process/${encodeURIComponent(taskId)}/external/pull`, + payload, +) + +export const buildDataProcessPreview = ( + taskId: string | number, + payload: DataProcessPreviewBuildPayload = {}, +) => post( + `/data-process/${encodeURIComponent(taskId)}/preview/build`, + payload, +) + +export function getDataProcessPreview( + taskId: string | number, + params: { source_file_id?: string | number; page?: number; page_size?: number; keyword?: string } = {}, +) { + return get>( + `/data-process/${encodeURIComponent(taskId)}/preview`, + params, + ) +} + +export const updateDataProcessPreview = ( + taskId: string | number, + previewId: string | number, + payload: DataProcessPreviewUpdatePayload, +) => put( + `/data-process/${encodeURIComponent(taskId)}/preview/${encodeURIComponent(previewId)}`, + payload, +) + +export const createDataProcessPreview = ( + taskId: string | number, + payload: DataProcessPreviewCreatePayload, +) => post(`/data-process/${encodeURIComponent(taskId)}/preview`, payload) + +export const deleteDataProcessPreview = ( + taskId: string | number, + previewId: string | number, +) => del>( + `/data-process/${encodeURIComponent(taskId)}/preview/${encodeURIComponent(previewId)}`, +) + +export const generateDataProcess = (taskId: string | number) => + post(`/data-process/${encodeURIComponent(taskId)}/generate`) + +export const stopDataProcess = (taskId: string | number) => + post(`/data-process/${encodeURIComponent(taskId)}/stop`) + +export const getDataProcessProgress = (taskId: string | number) => + get(`/data-process/${encodeURIComponent(taskId)}/progress`) + +export function getDataProcessResults( + taskId: string | number, + params: { page?: number; page_size?: number; keyword?: string; status?: string; split?: string } = {}, +) { + return get>( + `/data-process/${encodeURIComponent(taskId)}/results`, + params, + ) +} + +export const updateDataProcessResult = ( + taskId: string | number, + resultId: string | number, + payload: DataProcessResultUpdatePayload, +) => put( + `/data-process/${encodeURIComponent(taskId)}/results/${encodeURIComponent(resultId)}`, + payload, +) + +export const restoreDataProcessResult = (taskId: string | number, resultId: string | number) => + post( + `/data-process/${encodeURIComponent(taskId)}/results/${encodeURIComponent(resultId)}/restore`, + ) + +export const publishDataProcess = (taskId: string | number, payload: DataProcessPublishPayload) => + post(`/data-process/${encodeURIComponent(taskId)}/publish`, payload) diff --git a/frontend/src/types/dataProcess.ts b/frontend/src/types/dataProcess.ts new file mode 100644 index 0000000..0998f04 --- /dev/null +++ b/frontend/src/types/dataProcess.ts @@ -0,0 +1,226 @@ +/** 数据处理模块的前后端契约。API 字段统一使用 snake_case。 */ + +export type DataProcessStatus = 'pending' | 'running' | 'completed' | 'failed' | 'stopped' +export type DataProcessType = 'structured' | 'unstructured' | 'external' +export type DataProcessResultStatus = 'valid' | 'modified' | 'invalid' +export type DataProcessSplit = 'train' | 'validation' | 'test' + +export interface DataProcessPage { + items: T[] + total: number + page: number + page_size: number +} + +export interface DataProcessDatasetSplit { + train: number + validation: number + test: number +} + +export type DataProcessConfig = Record & { + dataset_split?: DataProcessDatasetSplit +} + +export interface DataProcessTask { + id: string | number + name: string + description?: string + status: DataProcessStatus + process_type: DataProcessType + config?: DataProcessConfig + progress?: number + source_dataset_id?: string | number | null + source_dataset_name?: string | null + source_dataset?: string | null + output_dataset_id?: string | number | null + output_dataset_name?: string | null + output_dataset?: string | null + source_file_count?: number + input_count?: number + output_count?: number + filtered_count?: number + duplicate_count?: number + error_count?: number + creator_name?: string | null + creator?: string | null + created_by?: string | number | null + create_time?: string + created_at?: string + start_time?: string | null + started_at?: string | null + complete_time?: string | null + completed_at?: string | null + duration?: string | null + duration_seconds?: number | null + failure_reason?: string | null + source_files?: DataProcessSourceFile[] +} + +export interface DataProcessTaskCreatePayload { + name: string + description?: string + process_type: DataProcessType + config: DataProcessConfig +} + +export type DataProcessTaskUpdatePayload = Partial + +export interface DataProcessSourceFile { + id: string | number + task_id?: string | number + name: string + size_bytes: number + record_count: number + file_format?: string + checksum_sha256?: string + status?: string + create_time?: string +} + +export interface DataProcessSourceContent { + file_id?: string | number + file?: DataProcessSourceFile + content: string + start_line?: number + end_line?: number + line_count?: number + total_lines?: number + has_more?: boolean + truncated?: boolean + offset?: number + limit?: number + total_chars?: number +} + +export interface DataProcessExternalSourcePayload { + type: 'postgresql' + url: string + auth_mode: 'none' | 'basic' + username?: string + password?: string + limit: number + query?: string + file_name?: string +} + +export interface DataProcessExternalTestResult { + connected: boolean + latency_ms?: number + message?: string +} + +export interface DataProcessPreviewItem { + id: string | number + source_file_id: string | number + original_content: string + edited_content: string + source_start: number | null + source_end: number | null + source_start_line: number | null + source_end_line: number | null + token_count: number + status: 'original' | 'modified' | 'manual' | 'invalid' + updated_at?: string +} + +export interface DataProcessPreviewBuildPayload { + replace_existing?: true + source_file_ids?: Array +} + +export type DataProcessPreviewBuildResult = DataProcessPage + +export interface DataProcessPreviewCreatePayload { + source_file_id?: string | number | null + original_content?: string + edited_content: string + source_start?: number | null + source_end?: number | null + source_start_line?: number | null + source_end_line?: number | null + token_count?: number + status?: DataProcessPreviewItem['status'] +} + +export interface DataProcessPreviewUpdatePayload { + edited_content: string + status?: DataProcessPreviewItem['status'] + expected_updated_at?: string +} + +export interface DataProcessProgress { + task_id: string | number + status: DataProcessStatus + stage?: string + progress: number + message?: string + processed_count?: number + total_count?: number + input_count?: number + output_count?: number + filtered_count?: number + duplicate_count?: number + error_count?: number + failure_reason?: string | null + updated_at?: string +} + +export interface DataProcessResult { + id: string | number + preview_item_id?: string | number | null + instruction: string + input: string + output: string + original_instruction?: string | null + original_input?: string | null + original_output?: string | null + status: DataProcessResultStatus + error?: string | null + split?: DataProcessSplit | null + quality_score?: DataProcessQualityScore | null + updated_at?: string +} + +export interface DataProcessResultUpdatePayload { + instruction: string + input: string + output: string + expected_updated_at?: string +} + +export interface DataProcessQualityScore { + overall?: number + completeness?: number + length?: number + readability?: number + relevance?: number + duplicate?: number + is_valid?: boolean + flags?: string[] + fingerprint?: string + [key: string]: unknown +} + +export interface DataProcessPublishPayload { + dataset_name: string + dataset_type: 'train' | 'test' | 'eval' | 'val' | 'other' + storage_type: 'local' + split: DataProcessDatasetSplit + format: 'alpaca_jsonl' | 'jsonl' +} + +export interface DataProcessPublishResult { + dataset_id?: string | number + output_dataset_id?: string | number + dataset_name?: string + record_count?: number + dataset?: { + id: string | number + name?: string + record_count?: number + count?: number + [key: string]: unknown + } + created?: boolean +} diff --git a/frontend/src/views/data-process/DataProcessCreateView.vue b/frontend/src/views/data-process/DataProcessCreateView.vue index 9604226..3ee5d48 100644 --- a/frontend/src/views/data-process/DataProcessCreateView.vue +++ b/frontend/src/views/data-process/DataProcessCreateView.vue @@ -10,7 +10,7 @@ import SourceUploadStep from './create/SourceUploadStep.vue' import PreviewCompareStep from './create/PreviewCompareStep.vue' import GenerationStep from './create/GenerationStep.vue' import ResultEditorStep from './create/ResultEditorStep.vue' -import { buildPreviewItems, DEFAULT_SOURCE_TEXT } from './create/previewModel' +import { DEFAULT_SOURCE_TEXT, estimateTokenCount } from './create/previewModel' import { createDefaultStructuredOptions, createDefaultUnstructuredOptions, @@ -21,6 +21,25 @@ import { } from './create/useDataProcessDraft' import { useDataProcessGeneration } from './create/useDataProcessGeneration' import { useModelsStore } from '@/stores/models' +import { + buildDataProcessPreview, + createDataProcessPreview, + createDataProcessTask, + deleteDataProcessPreview, + deleteDataProcessSourceFile, + getDataProcessPreview, + getDataProcessSourceContent, + getDataProcessTask, + pullDataProcessExternalSource, + testDataProcessExternalSource, + updateDataProcessPreview, + updateDataProcessTask, + uploadDataProcessSourceFiles, + type DataProcessExternalSourcePayload, + type DataProcessPreviewItem, + type DataProcessSourceFile, +} from '@/api/modules/dataProcess' +import type { DataProcessConfig } from '@/types/dataProcess' import type { ExternalDataSource, GenerationControlOptions, @@ -35,11 +54,14 @@ import type { const router = useRouter() const modelsStore = useModelsStore() const { list: modelList } = storeToRefs(modelsStore) -const generationModels = computed(() => modelList.value.filter((model) => model.type === 'LLM')) +const generationModels = computed(() => modelList.value.filter((model) => ( + model.type === 'LLM' + && (model.model_source === 'api' || model.model_source === 'online' || Boolean(model.api_url)) +))) const taskSetupRef = ref>() const modelSelectionRef = ref>() const confirmDialogRef = ref>() -const PREVIEW_MODEL_VERSION = 'document-chunk-v2' +const PREVIEW_MODEL_VERSION = 'backend-pipeline-v1' const WIZARD_STEPS = [ { id: 'create', title: '创建任务', desc: '填写任务信息与处理配置' }, @@ -52,6 +74,7 @@ const WIZARD_STEPS = [ const currentStep = ref(0) const currentStepId = computed(() => WIZARD_STEPS[currentStep.value]?.id ?? 'create') const task = reactive({ name: '', description: '' }) +const taskId = ref(null) const processType = ref('structured') const structuredOptions = ref(createDefaultStructuredOptions()) const unstructuredOptions = ref(createDefaultUnstructuredOptions()) @@ -61,18 +84,17 @@ const modelSelectionOptions = computed(() => ( const uploadedFiles = ref([]) const externalSource = reactive({ - type: 'mysql', + type: 'postgresql', url: '', authMode: 'none', username: '', password: '', - token: '', limit: 1000, + query: '', + fileName: 'external-data.jsonl', }) const externalPulling = ref(false) const externalConnected = ref(false) -let connectionTimer: ReturnType | null = null -let pullTimer: ReturnType | null = null const fileName = computed(() => uploadedFiles.value.map(f => f.name).join(', ')) const previewSignature = ref('') @@ -88,6 +110,7 @@ const { generation, results, selectedResultId, + persistResultChanges, resetDownstream, restoreResult, startGeneration, @@ -96,11 +119,9 @@ const { updateResultField, validateResults, } = useDataProcessGeneration({ - previewItems, - processType, - structuredOptions, - unstructuredOptions, + taskId, dirty, + beforeGenerate: syncPreviewChanges, }) const modifiedPreviewCount = computed(() => previewItems.value.filter((item) => item.status !== 'original').length) @@ -160,7 +181,101 @@ function updateModelSelectionOptions(value: GenerationControlOptions) { structuredOptions.value = { ...structuredOptions.value, ...value } } +function toBackendConfig(): DataProcessConfig { + const options = processType.value === 'unstructured' + ? unstructuredOptions.value + : structuredOptions.value + + const common = { + preprocess_options: [...options.preprocessOptions], + semantic_enrichment: options.semanticEnrichment, + dataset_split: { ...options.datasetSplit }, + generation_model_id: options.generationModelId, + generation_prompt: options.generationPrompt, + temperature: options.temperature, + max_tokens: options.maxTokens, + json_mode: options.jsonMode, + quality_filter_enabled: options.qualityFilterEnabled, + filter_low_quality: options.filterLowQuality, + filter_short_content: options.filterShortContent, + min_output_length: options.minOutputLength, + } + + if (processType.value === 'unstructured') { + return { + ...common, + chunk_method: unstructuredOptions.value.chunkMethod, + chunk_size: unstructuredOptions.value.chunkSize, + chunk_overlap: unstructuredOptions.value.chunkOverlap, + min_chunk_size: unstructuredOptions.value.minChunkSize, + custom_delimiter: unstructuredOptions.value.customDelimiter, + preserve_tables: unstructuredOptions.value.preserveTables, + preserve_code_blocks: unstructuredOptions.value.preserveCodeBlocks, + preserve_lists: unstructuredOptions.value.preserveLists, + qa_pairs_per_chunk: unstructuredOptions.value.qaPairsPerChunk, + } + } + + return { + ...common, + qa_pairs_per_row: structuredOptions.value.qaPairsPerRow, + } +} + +function taskPayload() { + return { + name: task.name.trim(), + description: task.description.trim(), + process_type: processType.value, + config: toBackendConfig(), + } +} + +function externalPayload(): DataProcessExternalSourcePayload { + return { + type: externalSource.type, + url: externalSource.url.trim(), + auth_mode: externalSource.authMode, + username: externalSource.username || undefined, + password: externalSource.password || undefined, + limit: externalSource.limit, + query: externalSource.query?.trim() || undefined, + file_name: externalSource.fileName || 'external-data.jsonl', + } +} + +function mapPreviewItem(item: DataProcessPreviewItem): PreviewItem { + return { + id: String(item.id), + sourceFileId: String(item.source_file_id), + originalContent: item.original_content, + editedContent: item.edited_content, + sourceStart: item.source_start, + sourceEnd: item.source_end, + sourceStartLine: item.source_start_line, + sourceEndLine: item.source_end_line, + tokenCount: item.token_count, + status: item.status, + updatedAt: item.updated_at, + } +} + +function mapSourceFile(file: DataProcessSourceFile, content = ''): UploadedDataFile { + return { + uid: String(file.id), + sourceFileId: String(file.id), + name: file.name, + size: file.size_bytes, + count: file.record_count, + content, + fileFormat: file.file_format, + checksumSha256: file.checksum_sha256, + status: 'ready', + } +} + const { persistDraft, restoreDraft } = useDataProcessDraft({ + taskId, currentStepId, task, processType, @@ -255,7 +370,7 @@ const generationOptionsSignature = computed(() => JSON.stringify(generationAffec function buildPreviewSignature() { const filesSignature = uploadedFiles.value - .map((file) => `${file.uid}:${file.name}:${file.size}:${file.count}`) + .map((file) => `${file.uid}:${file.name}:${file.size}:${file.checksumSha256 || file.count}`) .join('|') return `${PREVIEW_MODEL_VERSION}:${processType.value}:${JSON.stringify(previewAffectingOptions())}:${filesSignature}` } @@ -280,7 +395,7 @@ watch(generationOptionsSignature, (currentSignature, previousSignature) => { }) watch( - [task, processType, structuredOptions, unstructuredOptions, externalSource], + [taskId, task, processType, structuredOptions, unstructuredOptions, externalSource], persistDraft, { deep: true }, ) @@ -294,48 +409,85 @@ watch(currentStep, () => nextTick(scrollToStepTop)) async function handleFileChange(uploadFile: UploadFile) { const raw = uploadFile.raw if (!raw) return + if (!taskId.value) { + ElMessage.error('任务尚未创建,请返回模型选择步骤后重试') + return + } if (raw.size > 200 * 1024 * 1024) { ElMessage.warning('单文件不能超过 200MB') return } const extension = raw.name.split('.').pop()?.toLowerCase() ?? '' - const textExtensions = ['txt', 'md', 'json', 'jsonl', 'csv'] + const textExtensions = new Set(['txt', 'md', 'json', 'jsonl', 'csv']) + if (!textExtensions.has(extension)) { + ElMessage.error('当前仅支持 TXT、Markdown、JSON、JSONL 和 CSV;不会用示例内容替代无法解析的文件') + return + } + + if (uploadedFiles.value.some((file) => file.name === raw.name && file.size === raw.size)) { + ElMessage.warning('同名且同大小的文件已经上传') + return + } + let content = '' - if (textExtensions.includes(extension)) { - try { - content = await raw.text() - } catch { - content = '' - } + try { + content = new TextDecoder('utf-8', { fatal: true }).decode(await raw.arrayBuffer()) + } catch { + ElMessage.error('文件不是有效的 UTF-8 文本,请转换编码后重试') + return + } + if (!content.trim()) { + ElMessage.warning('不能上传空文件') + return } - const fileContent = content.trim() ? content : DEFAULT_SOURCE_TEXT - const linesCount = fileContent.split('\n').filter((line) => line.trim()).length - - // Prevent duplicate upload of the same file - if (!uploadedFiles.value.some(f => f.name === raw.name && f.size === raw.size)) { - uploadedFiles.value.push({ - uid: uploadFile.uid || Date.now() + Math.random(), - name: raw.name, - size: raw.size, - count: linesCount, - content: fileContent - }) + try { + const uploaded = await uploadDataProcessSourceFiles(taskId.value, [raw]) + const source = uploaded.files[0] + if (!source) throw new Error('后端未返回源文件记录') + uploadedFiles.value.push(mapSourceFile(source, content)) + previewSignature.value = '' + resetDownstream() + dirty.value = true + ElMessage.success(`文件 ${source.name} 上传成功`) + } catch { + // 请求层已展示后端的解析或格式错误。 } - - dirty.value = true } -function useSampleFile() { - uploadedFiles.value = [{ - uid: 'sample-1', - name: 'finance_qa.jsonl', - size: 128 * 1024 * 1024, - count: DEFAULT_SOURCE_TEXT.split('\n').filter((line) => line.trim()).length, - content: DEFAULT_SOURCE_TEXT - }] - dirty.value = true +async function useSampleFile() { + if (!taskId.value) { + ElMessage.error('任务尚未创建,请返回模型选择步骤后重试') + return + } + const sample = new File([DEFAULT_SOURCE_TEXT], 'finance_qa.jsonl', { type: 'application/x-ndjson' }) + await handleFileChange({ raw: sample, uid: Date.now(), name: sample.name } as UploadFile) +} + +async function restoreRegisteredSources() { + if (!taskId.value) return + try { + const savedTask = await getDataProcessTask(taskId.value) + const sources = savedTask.source_files || [] + const restoredFiles = await Promise.all(sources.map(async (file) => { + try { + const source = await getDataProcessSourceContent(taskId.value!, file.id, { + start_line: 1, + line_count: 5000, + }) + return mapSourceFile(file, source.content) + } catch { + return mapSourceFile(file) + } + })) + uploadedFiles.value = restoredFiles + if (restoredFiles.length) { + ElMessage.success(`已同步 ${restoredFiles.length} 个已登记源文件`) + } + } catch { + ElMessage.warning('草稿任务暂时无法从后端同步,请检查服务后重试') + } } function updateExternalSource(value: ExternalDataSource) { @@ -343,58 +495,81 @@ function updateExternalSource(value: ExternalDataSource) { externalConnected.value = false } -function handleTestConnection() { +async function handleTestConnection() { + if (!taskId.value) { + ElMessage.error('任务尚未创建,请返回模型选择步骤后重试') + return + } if (!externalSource.url.trim()) { ElMessage.warning('请先填写数据源地址') return } - if (connectionTimer) clearTimeout(connectionTimer) externalPulling.value = true - connectionTimer = setTimeout(() => { - connectionTimer = null + try { + const result = await testDataProcessExternalSource(taskId.value, externalPayload()) + externalConnected.value = result.connected + if (result.connected) ElMessage.success(result.message || '数据源连接测试成功') + else ElMessage.warning(result.message || '数据源连接失败') + } catch { + externalConnected.value = false + } finally { externalPulling.value = false - externalConnected.value = true - ElMessage.success('数据源连接测试成功') - }, 1500) + } } -function handlePullData() { +async function handlePullData() { + if (!taskId.value) { + ElMessage.error('任务尚未创建,请返回模型选择步骤后重试') + return + } if (!externalSource.url.trim()) { ElMessage.warning('请先填写数据源地址') return } - if (pullTimer) clearTimeout(pullTimer) + if (!externalSource.query?.trim()) { + ElMessage.warning('请先填写只读 SELECT 查询语句') + return + } externalPulling.value = true - pullTimer = setTimeout(() => { - pullTimer = null - externalPulling.value = false + try { + const response = await pullDataProcessExternalSource(taskId.value, externalPayload()) + const newFiles: UploadedDataFile[] = [] + for (const file of response.files) { + const source = await getDataProcessSourceContent(taskId.value, file.id, { + start_line: 1, + line_count: 5000, + }) + newFiles.push(mapSourceFile(file, source.content)) + } + uploadedFiles.value.push(...newFiles) externalConnected.value = true - const typeName = externalSource.type.toUpperCase() - const id = `external-${Date.now()}` - uploadedFiles.value.push({ - uid: id, - name: `${typeName} 拉取数据 ${new Date().toLocaleString('zh-CN')}`, - size: Math.min(externalSource.limit, 5000) * 64, - count: Math.min(externalSource.limit, DEFAULT_SOURCE_TEXT.split('\n').filter((line) => line.trim()).length), - content: DEFAULT_SOURCE_TEXT, - }) - dirty.value = true - ElMessage.success(`已成功拉取 ${uploadedFiles.value[uploadedFiles.value.length - 1].count.toLocaleString()} 条数据`) - }, 2000) -} - -function handleRemoveFile(uid: string | number) { - const index = uploadedFiles.value.findIndex(f => f.uid === uid) - if (index > -1) { - uploadedFiles.value.splice(index, 1) previewSignature.value = '' - previewItems.value = [] - selectedPreviewId.value = null resetDownstream() dirty.value = true + ElMessage.success(`已成功登记 ${newFiles.length} 个外部源文件`) + } catch { + externalConnected.value = false + } finally { + externalPulling.value = false } } +async function handleRemoveFile(uid: string | number) { + const index = uploadedFiles.value.findIndex(f => f.uid === uid) + if (index < 0 || !taskId.value) return + try { + await deleteDataProcessSourceFile(taskId.value, uid) + } catch { + return + } + uploadedFiles.value.splice(index, 1) + previewSignature.value = '' + previewItems.value = [] + selectedPreviewId.value = null + resetDownstream() + dirty.value = true +} + function resetSourceDataForProcessTypeChange() { uploadedFiles.value = [] previewSignature.value = '' @@ -415,10 +590,24 @@ async function nextFromCreate() { async function nextFromModel() { const valid = await modelSelectionRef.value?.validate() if (!valid) return - goToStep('upload') + try { + const saved = taskId.value + ? await updateDataProcessTask(taskId.value, taskPayload()) + : await createDataProcessTask(taskPayload()) + taskId.value = String(saved.id) + dirty.value = true + persistDraft() + goToStep('upload') + } catch { + // 请求层已展示名称冲突或配置非法等具体原因。 + } } -function nextFromUpload() { +async function nextFromUpload() { + if (!taskId.value) { + ElMessage.error('任务尚未创建,请返回模型选择步骤后重试') + return + } if (uploadedFiles.value.length === 0) { ElMessage.warning(processType.value === 'external' ? '请先拉取至少一个数据源' : '请上传至少一个源数据文件') return @@ -426,14 +615,25 @@ function nextFromUpload() { const signature = buildPreviewSignature() if (signature !== previewSignature.value) { - previewItems.value = uploadedFiles.value.flatMap((file) => - buildPreviewItems( - file.content, - processType.value, - String(file.uid), - processType.value === 'unstructured' ? unstructuredOptions.value : undefined, - ), - ) + try { + await buildDataProcessPreview(taskId.value, { + source_file_ids: uploadedFiles.value.map((file) => file.sourceFileId || file.uid), + }) + const first = await getDataProcessPreview(taskId.value, { page: 1, page_size: 500 }) + const items = [...first.items] + const pages = Math.ceil(first.total / first.page_size) + for (let page = 2; page <= pages; page += 1) { + const next = await getDataProcessPreview(taskId.value, { page, page_size: 500 }) + items.push(...next.items) + } + previewItems.value = items.map(mapPreviewItem) + } catch { + return + } + if (!previewItems.value.length) { + ElMessage.warning('源文件没有生成可用的预览条目,请检查文件内容和预处理配置') + return + } selectedPreviewFileId.value = String(uploadedFiles.value[0]?.uid ?? '') || null selectedPreviewId.value = activePreviewItems.value[0]?.id ?? null selectedPreviewIdsByFile.value = selectedPreviewId.value && selectedPreviewFileId.value @@ -464,40 +664,50 @@ function updatePreviewContent(id: string, value: string) { const item = previewItems.value.find((entry) => entry.id === id) if (!item) return item.editedContent = value - item.tokenCount = Math.max(1, Math.ceil(value.length / 2)) + item.tokenCount = estimateTokenCount(value) item.status = value === item.originalContent ? 'original' : item.sourceStart == null ? 'manual' : 'modified' resetDownstream() dirty.value = true } +async function syncPreviewChanges() { + if (!taskId.value) throw new Error('任务尚未创建') + const changedItems = previewItems.value.filter((item) => item.status === 'modified' || item.status === 'manual') + for (const item of changedItems) { + const saved = await updateDataProcessPreview(taskId.value, item.id, { + edited_content: item.editedContent, + expected_updated_at: item.updatedAt, + }) + const index = previewItems.value.findIndex((entry) => entry.id === item.id) + if (index >= 0) previewItems.value[index] = mapPreviewItem(saved) + } +} + function restorePreviewItem(id: string) { const item = previewItems.value.find((entry) => entry.id === id) if (!item || item.sourceStart == null) return item.editedContent = item.originalContent - item.tokenCount = Math.max(1, Math.ceil(item.originalContent.length / 2)) + item.tokenCount = estimateTokenCount(item.originalContent) item.status = 'original' resetDownstream() dirty.value = true } -function addPreviewItem() { - if (!selectedPreviewFileId.value) return - const id = `manual-${Date.now()}` - previewItems.value.push({ - id, - sourceFileId: selectedPreviewFileId.value, - originalContent: '', - editedContent: '', - sourceStart: null, - sourceEnd: null, - sourceStartLine: null, - sourceEndLine: null, - tokenCount: 1, - status: 'manual', - }) - selectPreviewItem(id) - resetDownstream() - dirty.value = true +async function addPreviewItem() { + if (!selectedPreviewFileId.value || !taskId.value) return + try { + const created = await createDataProcessPreview(taskId.value, { + source_file_id: selectedPreviewFileId.value, + edited_content: '', + }) + const item = mapPreviewItem(created) + previewItems.value.push(item) + selectPreviewItem(item.id) + resetDownstream() + dirty.value = true + } catch { + // 请求层已展示错误。 + } } async function removePreviewItem(id: string) { @@ -512,6 +722,12 @@ async function removePreviewItem(id: string) { const index = previewItems.value.findIndex((item) => item.id === id) if (index < 0) return + if (!taskId.value) return + try { + await deleteDataProcessPreview(taskId.value, id) + } catch { + return + } previewItems.value.splice(index, 1) selectedPreviewId.value = activePreviewItems.value[Math.min(index, activePreviewItems.value.length - 1)]?.id ?? null if (selectedPreviewFileId.value && selectedPreviewId.value) { @@ -531,7 +747,7 @@ async function handlePrimaryAction() { return } if (currentStepId.value === 'upload') { - nextFromUpload() + await nextFromUpload() return } if (currentStepId.value === 'preview') { @@ -546,7 +762,7 @@ async function handlePrimaryAction() { if (generation.status === 'success') { goToStep('results') } else if (generation.status !== 'running') { - startGeneration() + await startGeneration() } return } @@ -567,6 +783,15 @@ async function saveTask() { ElMessage.warning('请先修正校验失败的结果') return } + try { + await persistResultChanges() + } catch { + return + } + if (!validateResults()) { + ElMessage.warning('仍有结果未通过后端质量校验,请继续修正') + return + } dirty.value = false localStorage.removeItem(DATA_PROCESS_DRAFT_STORAGE_KEY) allowLeave = true @@ -607,12 +832,10 @@ onBeforeRouteLeave(async () => { onBeforeUnmount(() => { stopGenerationTimer() - if (connectionTimer) clearTimeout(connectionTimer) - if (pullTimer) clearTimeout(pullTimer) }) -onMounted(() => { +onMounted(async () => { restoreDraft() - modelsStore.load() + await Promise.all([modelsStore.load(), restoreRegisteredSources()]) }) diff --git a/frontend/src/views/data-process/DataProcessDetailView.vue b/frontend/src/views/data-process/DataProcessDetailView.vue index cc916d3..54f5fd8 100644 --- a/frontend/src/views/data-process/DataProcessDetailView.vue +++ b/frontend/src/views/data-process/DataProcessDetailView.vue @@ -1,228 +1,451 @@