feat: 完成数据处理接口与前端接入

This commit is contained in:
caoxiaozhu
2026-07-23 15:10:13 +08:00
parent f453234057
commit f04dc479bb
29 changed files with 7126 additions and 1144 deletions

View File

@@ -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