fix(data-process): 延迟重新生成破坏性变更

This commit is contained in:
caoxiaozhu
2026-07-27 10:43:42 +08:00
parent 03bd0b6d03
commit 53014bb381
3 changed files with 299 additions and 143 deletions

View File

@@ -30,6 +30,7 @@ class FakeDataProcessStore:
self.previews: dict[str, list[dict[str, Any]]] = {}
self.results: dict[str, list[dict[str, Any]]] = {}
self.datasets: dict[str, dict[str, Any]] = {}
self.regeneration_prepared: set[str] = set()
self.sequence = 0
def _id(self, prefix: str) -> str:
@@ -109,27 +110,19 @@ class FakeDataProcessStore:
raise InvalidStateError("data process task was modified by another request")
if payload["process_type"] != task["process_type"]:
raise InvalidStateError("process_type cannot be changed during regeneration")
published_outputs_preserved = bool(task.get("output_dataset_id"))
self.results[task_id] = []
published_outputs_preserved = bool(task.get("output_dataset_id")) or any(
dataset.get("source_task_id") == task_id
for dataset in self.datasets.values()
)
task.update(
{
"name": payload["name"],
"description": payload["description"],
"config": deepcopy(payload["config"]),
"status": "pending",
"progress": 20,
"output_dataset_id": None,
"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,
"updated_at": "2026-07-25T20:00:00Z",
}
)
self.regeneration_prepared.add(task_id)
return {
"task": deepcopy(task),
"preview_invalidated": False,
@@ -340,14 +333,19 @@ class FakeDataProcessStore:
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")
task = self.tasks[task_id]
if task.get("output_dataset_id") and task_id not in self.regeneration_prepared:
raise InvalidStateError("published task cannot be regenerated")
if replace_existing:
self.results[task_id] = []
self.tasks[task_id].update(
task.update(
status="running",
progress=30,
output_dataset_id=None,
output_count=0,
generation_run_id=self._id("dprun"),
)
self.regeneration_prepared.discard(task_id)
return self.get_task(task_id)
def generation_is_running(self, task_id: str, generation_run_id: str) -> bool:
@@ -487,6 +485,8 @@ class FakeDataProcessStore:
def publish(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
task = self.tasks[task_id]
if task_id in self.regeneration_prepared:
raise InvalidStateError("regeneration must start and complete before publishing")
published = [
dataset
for dataset in self.datasets.values()
@@ -969,6 +969,13 @@ def test_regenerate_endpoint_prepares_an_existing_published_task(tmp_path: Path)
"output_count": 2,
}
)
store.datasets["dataset_train"] = {
"id": "dataset_train",
"name": "原训练集",
"type": "train",
"source_task_id": task_id,
"deleted_at": None,
}
store.previews[task_id] = [{"id": "preview_1", "edited_content": "原切片"}]
store.results[task_id] = [{"id": "result_1"}]
@@ -985,16 +992,24 @@ def test_regenerate_endpoint_prepares_an_existing_published_task(tmp_path: Path)
assert response.status_code == 200
data = response.json()["data"]
assert data["task"]["status"] == "pending"
assert data["task"]["output_dataset_id"] is None
assert data["task"]["status"] == "completed"
assert data["task"]["output_dataset_id"] == "dataset_train"
assert data["task"]["output_count"] == 2
assert data["preview_invalidated"] is False
assert data["published_outputs_preserved"] is True
assert store.results[task_id] == []
assert store.results[task_id] == [{"id": "result_1"}]
assert store.previews[task_id][0]["id"] == "preview_1"
detail = client.get(f"/modelTF/data-process/{task_id}").json()["data"]
assert detail["status"] == "completed"
assert detail["output_dataset_id"] == "dataset_train"
assert detail["output_count"] == 2
assert [item["id"] for item in detail["output_datasets"]] == ["dataset_train"]
def test_published_split_datasets_remain_in_detail_after_regeneration(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
client, store, _ = make_client(tmp_path)
task_id = client.post(
@@ -1017,6 +1032,7 @@ def test_published_split_datasets_remain_in_detail_after_regeneration(
"output": "答案",
}
]
store.previews[task_id] = [{"id": "preview_1", "edited_content": "原切片"}]
published = client.post(
f"/modelTF/data-process/{task_id}/publish",
@@ -1039,16 +1055,40 @@ def test_published_split_datasets_remain_in_detail_after_regeneration(
},
)
assert regenerated.status_code == 200
assert regenerated.json()["data"]["task"]["output_dataset_id"] is None
original_output_dataset_id = store.tasks[task_id]["output_dataset_id"]
prepared_task = regenerated.json()["data"]["task"]
assert prepared_task["status"] == "completed"
assert prepared_task["output_dataset_id"] == original_output_dataset_id
assert prepared_task["output_count"] == 1
assert store.results[task_id][0]["id"] == "result_1"
detail = client.get(f"/modelTF/data-process/{task_id}")
assert detail.status_code == 200
detail_data = detail.json()["data"]
assert detail_data["status"] == "pending"
assert detail_data["output_dataset_id"] is None
assert detail_data["status"] == "completed"
assert detail_data["output_dataset_id"] == original_output_dataset_id
assert detail_data["output_count"] == 1
assert len(detail_data["output_datasets"]) == 3
assert {item["id"] for item in detail_data["output_datasets"]} == published_ids
assert set(store.datasets) == published_ids
retained_results = client.get(f"/modelTF/data-process/{task_id}/results").json()["data"]
assert retained_results["total"] == 1
assert retained_results["items"][0]["id"] == "result_1"
monkeypatch.setattr(data_process_endpoint, "_run_generation", lambda *args: None)
started = client.post(f"/modelTF/data-process/{task_id}/generate")
assert started.status_code == 200
running = started.json()["data"]
assert running["status"] == "running"
assert running["output_count"] == 0
assert store.results[task_id] == []
assert set(store.datasets) == published_ids
running_detail = client.get(f"/modelTF/data-process/{task_id}").json()["data"]
assert running_detail["status"] == "running"
assert running_detail["output_dataset_id"] is None
assert running_detail["output_count"] == 0
assert {item["id"] for item in running_detail["output_datasets"]} == published_ids
def test_regenerate_endpoint_validates_snapshot_and_locked_process_type(