fix(data-process): 保留重新生成前的已发布数据集
This commit is contained in:
@@ -174,6 +174,7 @@ def _decode_row(row: dict[str, Any] | None) -> dict[str, Any] | None:
|
|||||||
"metadata": {},
|
"metadata": {},
|
||||||
"quality_score": {},
|
"quality_score": {},
|
||||||
"versions": [],
|
"versions": [],
|
||||||
|
"output_datasets": [],
|
||||||
}.items():
|
}.items():
|
||||||
if key in item:
|
if key in item:
|
||||||
item[key] = _json_value(item[key], default)
|
item[key] = _json_value(item[key], default)
|
||||||
@@ -318,7 +319,12 @@ class DataProcessStore:
|
|||||||
) ORDER BY CASE dataset.type
|
) ORDER BY CASE dataset.type
|
||||||
WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4 END), '[]'::json)
|
WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4 END), '[]'::json)
|
||||||
FROM datasets dataset
|
FROM datasets dataset
|
||||||
WHERE dataset.source_task_id=task.id AND dataset.source='task')
|
WHERE dataset.source='task'
|
||||||
|
AND dataset.deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
dataset.source_task_id=task.id
|
||||||
|
OR (dataset.source_task_id IS NULL AND dataset.task_id=task.id)
|
||||||
|
))
|
||||||
AS output_datasets,
|
AS output_datasets,
|
||||||
CASE
|
CASE
|
||||||
WHEN task.started_at IS NOT NULL AND task.completed_at IS NOT NULL
|
WHEN task.started_at IS NOT NULL AND task.completed_at IS NOT NULL
|
||||||
@@ -459,14 +465,30 @@ class DataProcessStore:
|
|||||||
current_config,
|
current_config,
|
||||||
next_config,
|
next_config,
|
||||||
)
|
)
|
||||||
|
now = utcnow()
|
||||||
|
# 002 迁移前发布的数据集只有 task_id。先补齐新关联字段,保证
|
||||||
|
# 解除任务输出指针后,详情和后续重新发布仍能定位原来的三份数据集。
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE datasets
|
||||||
|
SET source_task_id=%s, updated_at=%s
|
||||||
|
WHERE source='task' AND source_task_id IS NULL AND task_id=%s
|
||||||
|
AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(task_id, now, task_id),
|
||||||
|
)
|
||||||
published_row = conn.execute(
|
published_row = conn.execute(
|
||||||
"""
|
"""
|
||||||
SELECT EXISTS(
|
SELECT EXISTS(
|
||||||
SELECT 1 FROM datasets
|
SELECT 1 FROM datasets
|
||||||
WHERE source_task_id=%s AND source='task'
|
WHERE source='task' AND deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
source_task_id=%s
|
||||||
|
OR (source_task_id IS NULL AND task_id=%s)
|
||||||
|
)
|
||||||
) AS exists
|
) AS exists
|
||||||
""",
|
""",
|
||||||
(task_id,),
|
(task_id, task_id),
|
||||||
).fetchone()
|
).fetchone()
|
||||||
published_outputs_preserved = bool(task.get("output_dataset_id")) or bool(
|
published_outputs_preserved = bool(task.get("output_dataset_id")) or bool(
|
||||||
published_row and published_row.get("exists")
|
published_row and published_row.get("exists")
|
||||||
@@ -493,7 +515,6 @@ class DataProcessStore:
|
|||||||
).fetchone()
|
).fetchone()
|
||||||
preview_count = int((preview_row or {}).get("count") or 0)
|
preview_count = int((preview_row or {}).get("count") or 0)
|
||||||
|
|
||||||
now = utcnow()
|
|
||||||
row = conn.execute(
|
row = conn.execute(
|
||||||
"""
|
"""
|
||||||
UPDATE data_process_tasks
|
UPDATE data_process_tasks
|
||||||
@@ -1460,10 +1481,14 @@ class DataProcessStore:
|
|||||||
existing_datasets = conn.execute(
|
existing_datasets = conn.execute(
|
||||||
"""
|
"""
|
||||||
SELECT * FROM datasets
|
SELECT * FROM datasets
|
||||||
WHERE source_task_id=%s AND source='task'
|
WHERE source='task' AND deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
source_task_id=%s
|
||||||
|
OR (source_task_id IS NULL AND task_id=%s)
|
||||||
|
)
|
||||||
ORDER BY created_at, id
|
ORDER BY created_at, id
|
||||||
""",
|
""",
|
||||||
(task_id,),
|
(task_id, task_id),
|
||||||
).fetchall()
|
).fetchall()
|
||||||
existing_by_split: dict[str, dict[str, Any]] = {}
|
existing_by_split: dict[str, dict[str, Any]] = {}
|
||||||
primary_existing = None
|
primary_existing = None
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
from collections.abc import Iterator
|
from collections.abc import Iterator
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
@@ -31,9 +32,13 @@ class _Result:
|
|||||||
|
|
||||||
|
|
||||||
class _PublishConnection:
|
class _PublishConnection:
|
||||||
def __init__(self, results: list[dict[str, Any]]):
|
def __init__(
|
||||||
|
self,
|
||||||
|
results: list[dict[str, Any]],
|
||||||
|
datasets: list[dict[str, Any]] | None = None,
|
||||||
|
):
|
||||||
self.results = results
|
self.results = results
|
||||||
self.datasets: list[dict[str, Any]] = []
|
self.datasets: list[dict[str, Any]] = datasets or []
|
||||||
self.files: list[dict[str, Any]] = []
|
self.files: list[dict[str, Any]] = []
|
||||||
self.records: list[dict[str, Any]] = []
|
self.records: list[dict[str, Any]] = []
|
||||||
|
|
||||||
@@ -47,16 +52,38 @@ class _PublishConnection:
|
|||||||
)
|
)
|
||||||
if normalized.startswith("SELECT * FROM data_process_results"):
|
if normalized.startswith("SELECT * FROM data_process_results"):
|
||||||
return _Result(rows=self.results)
|
return _Result(rows=self.results)
|
||||||
if normalized.startswith("SELECT * FROM datasets WHERE source_task_id"):
|
if normalized.startswith("SELECT * FROM datasets WHERE source='task'"):
|
||||||
return _Result(rows=self.datasets)
|
assert "source_task_id=%s" in normalized
|
||||||
|
assert "source_task_id IS NULL AND task_id=%s" in normalized
|
||||||
|
assert "deleted_at IS NULL" in normalized
|
||||||
|
source_task_id, legacy_task_id = params
|
||||||
|
return _Result(
|
||||||
|
rows=[
|
||||||
|
item
|
||||||
|
for item in self.datasets
|
||||||
|
if item.get("source") == "task"
|
||||||
|
and item.get("deleted_at") is None
|
||||||
|
and (
|
||||||
|
item.get("source_task_id") == source_task_id
|
||||||
|
or (
|
||||||
|
item.get("source_task_id") is None
|
||||||
|
and item.get("task_id") == legacy_task_id
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
if normalized.startswith("INSERT INTO datasets"):
|
if normalized.startswith("INSERT INTO datasets"):
|
||||||
dataset = {
|
dataset = {
|
||||||
"id": params[0],
|
"id": params[0],
|
||||||
"name": params[1],
|
"name": params[1],
|
||||||
"type": params[2],
|
"type": params[2],
|
||||||
|
"source": "task",
|
||||||
|
"task_id": params[4],
|
||||||
|
"source_task_id": params[5],
|
||||||
"count": params[8],
|
"count": params[8],
|
||||||
"record_count": params[9],
|
"record_count": params[9],
|
||||||
"metadata": params[11],
|
"metadata": params[11],
|
||||||
|
"deleted_at": None,
|
||||||
}
|
}
|
||||||
self.datasets.append(dataset)
|
self.datasets.append(dataset)
|
||||||
return _Result(row=dataset)
|
return _Result(row=dataset)
|
||||||
@@ -118,12 +145,34 @@ class _PublishStore(DataProcessStore):
|
|||||||
|
|
||||||
|
|
||||||
class _RegenerationConnection:
|
class _RegenerationConnection:
|
||||||
def __init__(self, task: dict[str, Any]) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
task: dict[str, Any],
|
||||||
|
datasets: list[dict[str, Any]] | None = None,
|
||||||
|
) -> None:
|
||||||
self.task = task
|
self.task = task
|
||||||
self.datasets = [
|
self.datasets = datasets or [
|
||||||
{"id": "dataset_train"},
|
{
|
||||||
{"id": "dataset_validation"},
|
"id": "dataset_train",
|
||||||
{"id": "dataset_test"},
|
"source": "task",
|
||||||
|
"task_id": task["id"],
|
||||||
|
"source_task_id": task["id"],
|
||||||
|
"deleted_at": None,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "dataset_validation",
|
||||||
|
"source": "task",
|
||||||
|
"task_id": task["id"],
|
||||||
|
"source_task_id": task["id"],
|
||||||
|
"deleted_at": None,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "dataset_test",
|
||||||
|
"source": "task",
|
||||||
|
"task_id": task["id"],
|
||||||
|
"source_task_id": task["id"],
|
||||||
|
"deleted_at": None,
|
||||||
|
},
|
||||||
]
|
]
|
||||||
self.sources = [{"id": "source_1"}]
|
self.sources = [{"id": "source_1"}]
|
||||||
self.previews = [{"id": "preview_1"}]
|
self.previews = [{"id": "preview_1"}]
|
||||||
@@ -133,8 +182,35 @@ class _RegenerationConnection:
|
|||||||
normalized = " ".join(sql.split())
|
normalized = " ".join(sql.split())
|
||||||
if params is not None:
|
if params is not None:
|
||||||
assert normalized.count("%s") == len(params)
|
assert normalized.count("%s") == len(params)
|
||||||
|
if normalized.startswith("UPDATE datasets SET source_task_id="):
|
||||||
|
source_task_id, _, legacy_task_id = params
|
||||||
|
for dataset in self.datasets:
|
||||||
|
if (
|
||||||
|
dataset.get("source") == "task"
|
||||||
|
and dataset.get("source_task_id") is None
|
||||||
|
and dataset.get("task_id") == legacy_task_id
|
||||||
|
and dataset.get("deleted_at") is None
|
||||||
|
):
|
||||||
|
dataset["source_task_id"] = source_task_id
|
||||||
|
return _Result()
|
||||||
if normalized.startswith("SELECT EXISTS("):
|
if normalized.startswith("SELECT EXISTS("):
|
||||||
return _Result(row={"exists": bool(self.datasets)})
|
assert "source_task_id=%s" in normalized
|
||||||
|
assert "source_task_id IS NULL AND task_id=%s" in normalized
|
||||||
|
assert "deleted_at IS NULL" in normalized
|
||||||
|
source_task_id, legacy_task_id = params
|
||||||
|
exists = any(
|
||||||
|
dataset.get("source") == "task"
|
||||||
|
and dataset.get("deleted_at") is None
|
||||||
|
and (
|
||||||
|
dataset.get("source_task_id") == source_task_id
|
||||||
|
or (
|
||||||
|
dataset.get("source_task_id") is None
|
||||||
|
and dataset.get("task_id") == legacy_task_id
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for dataset in self.datasets
|
||||||
|
)
|
||||||
|
return _Result(row={"exists": exists})
|
||||||
if normalized.startswith("DELETE FROM data_process_results"):
|
if normalized.startswith("DELETE FROM data_process_results"):
|
||||||
self.results.clear()
|
self.results.clear()
|
||||||
return _Result()
|
return _Result()
|
||||||
@@ -183,6 +259,63 @@ class _RegenerationStore(DataProcessStore):
|
|||||||
return dict(self._conn.task)
|
return dict(self._conn.task)
|
||||||
|
|
||||||
|
|
||||||
|
class _TaskDetailConnection:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
task: dict[str, Any],
|
||||||
|
datasets: list[dict[str, Any]],
|
||||||
|
) -> None:
|
||||||
|
self.task = task
|
||||||
|
self.datasets = datasets
|
||||||
|
|
||||||
|
def execute(self, sql: str, params: Any = None) -> _Result:
|
||||||
|
normalized = " ".join(sql.split())
|
||||||
|
assert normalized.startswith("SELECT task.*")
|
||||||
|
assert "dataset.source_task_id=task.id" in normalized
|
||||||
|
assert "dataset.source_task_id IS NULL AND dataset.task_id=task.id" in normalized
|
||||||
|
assert "dataset.deleted_at IS NULL" in normalized
|
||||||
|
task_id = params[0]
|
||||||
|
visible = [
|
||||||
|
dataset
|
||||||
|
for dataset in self.datasets
|
||||||
|
if dataset.get("source") == "task"
|
||||||
|
and dataset.get("deleted_at") is None
|
||||||
|
and (
|
||||||
|
dataset.get("source_task_id") == task_id
|
||||||
|
or (
|
||||||
|
dataset.get("source_task_id") is None
|
||||||
|
and dataset.get("task_id") == task_id
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
return _Result(
|
||||||
|
row={
|
||||||
|
**self.task,
|
||||||
|
"output_datasets": json.dumps(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"id": item["id"],
|
||||||
|
"name": item["name"],
|
||||||
|
"type": item["type"],
|
||||||
|
"count": item["count"],
|
||||||
|
"dataset_split": item["dataset_split"],
|
||||||
|
}
|
||||||
|
for item in visible
|
||||||
|
]
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _TaskDetailStore(DataProcessStore):
|
||||||
|
def __init__(self, conn: _TaskDetailConnection) -> None:
|
||||||
|
self._conn = conn
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def connect(self) -> Iterator[_TaskDetailConnection]:
|
||||||
|
yield self._conn
|
||||||
|
|
||||||
|
|
||||||
def test_decode_row_serializes_postgres_numeric_values_as_json_numbers() -> None:
|
def test_decode_row_serializes_postgres_numeric_values_as_json_numbers() -> None:
|
||||||
decoded = _decode_row(
|
decoded = _decode_row(
|
||||||
{
|
{
|
||||||
@@ -194,6 +327,20 @@ def test_decode_row_serializes_postgres_numeric_values_as_json_numbers() -> None
|
|||||||
assert decoded == {"progress": 100.0, "duration_seconds": 389.0}
|
assert decoded == {"progress": 100.0, "duration_seconds": 389.0}
|
||||||
|
|
||||||
|
|
||||||
|
def test_decode_row_decodes_aggregated_output_datasets_json() -> None:
|
||||||
|
decoded = _decode_row(
|
||||||
|
{
|
||||||
|
"id": "task-1",
|
||||||
|
"output_datasets": '[{"id":"dataset_train","type":"train"}]',
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decoded == {
|
||||||
|
"id": "task-1",
|
||||||
|
"output_datasets": [{"id": "dataset_train", "type": "train"}],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
("process_type", "current", "next_config", "expected"),
|
("process_type", "current", "next_config", "expected"),
|
||||||
[
|
[
|
||||||
@@ -314,6 +461,37 @@ def _regeneration_task(**updates: Any) -> dict[str, Any]:
|
|||||||
return task
|
return task
|
||||||
|
|
||||||
|
|
||||||
|
def _legacy_published_datasets(task_id: str = "task-1") -> list[dict[str, Any]]:
|
||||||
|
specs = (
|
||||||
|
("dataset_train", "制度问答-训练集", "train", "train", 22),
|
||||||
|
("dataset_validation", "制度问答-验证集", "val", "validation", 3),
|
||||||
|
("dataset_test", "制度问答-测试集", "test", "test", 3),
|
||||||
|
)
|
||||||
|
dataset_ids = {split: dataset_id for dataset_id, _, _, split, _ in specs}
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"id": dataset_id,
|
||||||
|
"name": name,
|
||||||
|
"type": dataset_type,
|
||||||
|
"source": "task",
|
||||||
|
"task_id": task_id,
|
||||||
|
"source_task_id": None,
|
||||||
|
"count": count,
|
||||||
|
"record_count": count,
|
||||||
|
"dataset_split": split,
|
||||||
|
"metadata": json.dumps(
|
||||||
|
{
|
||||||
|
"base_dataset_name": "制度问答",
|
||||||
|
"dataset_split": split,
|
||||||
|
"split_dataset_ids": dataset_ids,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
"deleted_at": None,
|
||||||
|
}
|
||||||
|
for dataset_id, name, dataset_type, split, count in specs
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_prepare_regeneration_preserves_outputs_sources_and_generation_only_preview() -> None:
|
def test_prepare_regeneration_preserves_outputs_sources_and_generation_only_preview() -> None:
|
||||||
conn = _RegenerationConnection(_regeneration_task())
|
conn = _RegenerationConnection(_regeneration_task())
|
||||||
original_datasets = list(conn.datasets)
|
original_datasets = list(conn.datasets)
|
||||||
@@ -344,6 +522,64 @@ def test_prepare_regeneration_preserves_outputs_sources_and_generation_only_prev
|
|||||||
assert conn.sources == original_sources
|
assert conn.sources == original_sources
|
||||||
|
|
||||||
|
|
||||||
|
def test_prepare_regeneration_backfills_and_keeps_legacy_task_datasets_visible() -> None:
|
||||||
|
legacy_datasets = _legacy_published_datasets()
|
||||||
|
deleted_dataset = {
|
||||||
|
**legacy_datasets[0],
|
||||||
|
"id": "dataset_deleted",
|
||||||
|
"name": "已删除训练集",
|
||||||
|
"deleted_at": "2026-07-25T20:00:00Z",
|
||||||
|
}
|
||||||
|
unrelated_dataset = {
|
||||||
|
**legacy_datasets[0],
|
||||||
|
"id": "dataset_unrelated",
|
||||||
|
"name": "其他任务训练集",
|
||||||
|
"task_id": "task-other",
|
||||||
|
}
|
||||||
|
conn = _RegenerationConnection(
|
||||||
|
_regeneration_task(),
|
||||||
|
[*legacy_datasets, deleted_dataset, unrelated_dataset],
|
||||||
|
)
|
||||||
|
|
||||||
|
before = _TaskDetailStore(_TaskDetailConnection(conn.task, conn.datasets)).get_task(
|
||||||
|
"task-1"
|
||||||
|
)
|
||||||
|
assert [item["id"] for item in before["output_datasets"]] == [
|
||||||
|
"dataset_train",
|
||||||
|
"dataset_validation",
|
||||||
|
"dataset_test",
|
||||||
|
]
|
||||||
|
|
||||||
|
result = _RegenerationStore(conn).prepare_regeneration(
|
||||||
|
"task-1",
|
||||||
|
{
|
||||||
|
"name": "原任务",
|
||||||
|
"description": "",
|
||||||
|
"process_type": "unstructured",
|
||||||
|
"config": {"chunk_method": "fixed", "temperature": 0.2},
|
||||||
|
"expected_updated_at": "2026-07-25T18:05:00Z",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["published_outputs_preserved"] is True
|
||||||
|
assert result["task"]["output_dataset_id"] is None
|
||||||
|
assert len(conn.datasets) == 5
|
||||||
|
assert all(
|
||||||
|
item["source_task_id"] == "task-1" for item in conn.datasets[:3]
|
||||||
|
)
|
||||||
|
assert deleted_dataset["source_task_id"] is None
|
||||||
|
assert unrelated_dataset["source_task_id"] is None
|
||||||
|
|
||||||
|
after = _TaskDetailStore(_TaskDetailConnection(conn.task, conn.datasets)).get_task(
|
||||||
|
"task-1"
|
||||||
|
)
|
||||||
|
assert [item["id"] for item in after["output_datasets"]] == [
|
||||||
|
"dataset_train",
|
||||||
|
"dataset_validation",
|
||||||
|
"dataset_test",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_prepare_regeneration_deletes_preview_when_chunk_configuration_changes() -> None:
|
def test_prepare_regeneration_deletes_preview_when_chunk_configuration_changes() -> None:
|
||||||
conn = _RegenerationConnection(_regeneration_task(output_dataset_id=None))
|
conn = _RegenerationConnection(_regeneration_task(output_dataset_id=None))
|
||||||
|
|
||||||
@@ -467,6 +703,37 @@ def test_publish_creates_three_independent_datasets_with_exact_counts() -> None:
|
|||||||
assert republished["created"] is False
|
assert republished["created"] is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_publish_reuses_legacy_task_id_only_split_datasets() -> None:
|
||||||
|
results = [
|
||||||
|
{
|
||||||
|
"id": f"result-{index}",
|
||||||
|
"status": "valid",
|
||||||
|
"instruction": f"问题 {index}",
|
||||||
|
"input": "",
|
||||||
|
"output": f"答案 {index}",
|
||||||
|
"preview_item_id": f"preview-{index}",
|
||||||
|
}
|
||||||
|
for index in range(10)
|
||||||
|
]
|
||||||
|
legacy_datasets = _legacy_published_datasets()
|
||||||
|
original_ids = [item["id"] for item in legacy_datasets]
|
||||||
|
conn = _PublishConnection(results, legacy_datasets)
|
||||||
|
|
||||||
|
published = _PublishStore(conn).publish(
|
||||||
|
"task-1",
|
||||||
|
{
|
||||||
|
"dataset_name": "不会创建新数据集",
|
||||||
|
"storage_type": "local",
|
||||||
|
"format": "alpaca_jsonl",
|
||||||
|
"split": {"train": 80, "validation": 10, "test": 10},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert published["created"] is False
|
||||||
|
assert [item["id"] for item in published["datasets"]] == original_ids
|
||||||
|
assert [item["id"] for item in conn.datasets] == original_ids
|
||||||
|
|
||||||
|
|
||||||
def test_publish_keeps_all_three_datasets_when_a_small_split_is_empty() -> None:
|
def test_publish_keeps_all_three_datasets_when_a_small_split_is_empty() -> None:
|
||||||
conn = _PublishConnection(
|
conn = _PublishConnection(
|
||||||
[
|
[
|
||||||
|
|||||||
Reference in New Issue
Block a user