diff --git a/backend/app/db/platform_store.py b/backend/app/db/platform_store.py index c011839..0d872bf 100644 --- a/backend/app/db/platform_store.py +++ b/backend/app/db/platform_store.py @@ -1151,12 +1151,25 @@ class PlatformStore: def datasets(self) -> list[dict[str, Any]]: with self.connect() as conn: - rows = conn.execute("SELECT * FROM datasets ORDER BY create_time DESC").fetchall() + rows = conn.execute( + """SELECT dataset.*, task.name AS task_name + FROM datasets dataset + LEFT JOIN data_process_tasks task + ON task.id=COALESCE(dataset.source_task_id, dataset.task_id) + ORDER BY dataset.create_time DESC""" + ).fetchall() return [self._dataset(conn, row) for row in rows] def dataset(self, dataset_id: str) -> dict[str, Any]: with self.connect() as conn: - row = conn.execute("SELECT * FROM datasets WHERE id=?", (dataset_id,)).fetchone() + row = conn.execute( + """SELECT dataset.*, task.name AS task_name + FROM datasets dataset + LEFT JOIN data_process_tasks task + ON task.id=COALESCE(dataset.source_task_id, dataset.task_id) + WHERE dataset.id=?""", + (dataset_id,), + ).fetchone() if not row: raise KeyError(dataset_id) return self._dataset(conn, row) diff --git a/backend/tests/test_platform_dataset_metadata.py b/backend/tests/test_platform_dataset_metadata.py index 11a9ef2..1a873a1 100644 --- a/backend/tests/test_platform_dataset_metadata.py +++ b/backend/tests/test_platform_dataset_metadata.py @@ -1,6 +1,46 @@ from __future__ import annotations -from app.db.platform_store import dataset_file_version_summary, parse_size_bytes +from contextlib import contextmanager +from typing import Any, Iterator + +from app.db.platform_store import PlatformStore, dataset_file_version_summary, parse_size_bytes + + +class _DatasetCursor: + def __init__(self, rows: list[dict[str, Any]]) -> None: + self.rows = rows + + def fetchall(self) -> list[dict[str, Any]]: + return self.rows + + +class _DatasetConnection: + def __init__(self) -> None: + self.queries: list[str] = [] + + def execute(self, sql: str, params: tuple[Any, ...] | None = None) -> _DatasetCursor: + self.queries.append(sql) + if "FROM datasets dataset" in sql: + return _DatasetCursor( + [ + { + "id": "dataset-train", + "name": "cash-数据集-训练集", + "type": "train", + "storage_type": "local", + "source": "task", + "task_id": "task-cash", + "source_task_id": "task-cash", + "task_name": "cash", + "size": "0 B", + "size_bytes": 0, + "metadata": "{}", + } + ] + ) + if "FROM dataset_files" in sql or "FROM dataset_records" in sql: + return _DatasetCursor([]) + raise AssertionError(f"unexpected query: {sql}") def test_parse_size_bytes_supports_legacy_units() -> None: @@ -44,3 +84,20 @@ def test_dataset_file_version_summary_uses_normalized_version_number_as_fallback assert summary["current_version_no"] == 1 assert summary["version_count"] == 0 + + +def test_dataset_list_exposes_source_task_name() -> None: + store = PlatformStore.__new__(PlatformStore) + conn = _DatasetConnection() + + @contextmanager + def connect() -> Iterator[_DatasetConnection]: + yield conn + + store.connect = connect # type: ignore[method-assign] + + [dataset] = store.datasets() + + assert dataset["task_name"] == "cash" + assert dataset["name"] == "cash-数据集-训练集" + assert any("task.name AS task_name" in query for query in conn.queries) diff --git a/frontend/scripts/regression-dataset-task-tab.mjs b/frontend/scripts/regression-dataset-task-tab.mjs index c3c47a4..028059d 100644 --- a/frontend/scripts/regression-dataset-task-tab.mjs +++ b/frontend/scripts/regression-dataset-task-tab.mjs @@ -18,6 +18,7 @@ const [typesSource, dataSource, viewSource, tablePageSource] = await Promise.all assert.match(typesSource, /export type DatasetSource = 'upload' \| 'task'/) assert.match(typesSource, /source\?: DatasetSource/) +assert.match(typesSource, /task_name\?: string/) assert.equal((dataSource.match(/source: 'upload'/g) || []).length, 6) assert.equal((dataSource.match(/source: 'task'/g) || []).length, 4) @@ -31,6 +32,7 @@ for (const name of [ const datasetLine = dataSource.split('\n').find((line) => line.includes(`name: '${name}'`)) assert.ok(datasetLine, `缺少数据任务 Mock:${name}`) assert.match(datasetLine, /source: 'task'/, `${name} 必须标记为数据任务来源`) + assert.match(datasetLine, /task_name: '/, `${name} 必须提供来源任务名称`) } assert.match(viewSource, /route\.query\.tab === 'task' \? 'task' : 'upload'/) @@ -42,7 +44,9 @@ assert.match( ) assert.match(viewSource, /return dataList\.value\.filter\(\(item\) => item\.source !== 'task'\)/) assert.doesNotMatch(viewSource, /数据任务产生的数据集[\s\S]*?return \[\]/) -assert.match(viewSource, /:search-fields="\['task_id', 'name'\]"/, '数据任务搜索应支持任务 ID') +assert.match(viewSource, /:search-fields="activeTab === 'task' \? \['task_id', 'task_name'\] : \['name'\]"/, '搜索字段应与当前页签展示的名称一致') +assert.match(viewSource, /activeTab === 'task' \? '训练任务名称' : '数据集名称'/, '数据任务页应展示训练任务名称表头') +assert.match(viewSource, /activeTab === 'task' \? \(row\.task_name \|\| '-'\) : row\.name/, '数据任务页应展示接口返回的来源任务名称') assert.match(viewSource, /formatMegabytes\(row\.size_bytes, row\.size\)/, '数据集大小应统一转换为 MB') assert.match(viewSource, / ({ ...dataset, files: dataset.files?.length diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 5ab7c8c..afe8427 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -80,6 +80,7 @@ export interface DatasetItem { storage_type: DatasetStorage | string source?: DatasetSource task_id?: string | number + task_name?: string size?: string | number size_bytes?: number count?: number diff --git a/frontend/src/views/dataset/DatasetListView.vue b/frontend/src/views/dataset/DatasetListView.vue index 5aa125b..6a901b8 100644 --- a/frontend/src/views/dataset/DatasetListView.vue +++ b/frontend/src/views/dataset/DatasetListView.vue @@ -149,7 +149,7 @@ onMounted(loadData) :data="filteredDataList" :loading="loading" searchable - :search-fields="['task_id', 'name']" + :search-fields="activeTab === 'task' ? ['task_id', 'task_name'] : ['name']" :multi-select="activeTab === 'task' && batchMode" :show-batch-bar="false" :create-text="activeTab === 'upload' ? '上传数据集' : ''" @@ -199,7 +199,15 @@ onMounted(loadData)