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