104 lines
3.2 KiB
Python
104 lines
3.2 KiB
Python
from __future__ import annotations
|
|
|
|
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:
|
|
assert parse_size_bytes("21563 B") == 21563
|
|
assert parse_size_bytes("1.5 KB") == 1536
|
|
assert parse_size_bytes("2 MB") == 2 * 1024**2
|
|
assert parse_size_bytes(4096) == 4096
|
|
assert parse_size_bytes("unknown") == 0
|
|
|
|
|
|
def test_dataset_file_version_summary_uses_active_version_metadata() -> None:
|
|
summary = dataset_file_version_summary(
|
|
{
|
|
"active_version_id": "file-1-v3",
|
|
"current_version_id": "file-1-v1",
|
|
"version_no": 1,
|
|
"versions": (
|
|
'[{"id":"file-1-v1","version":1},'
|
|
'{"id":"file-1-v3","version_no":3}]'
|
|
),
|
|
}
|
|
)
|
|
|
|
assert summary == {
|
|
"active_version_id": "file-1-v3",
|
|
"current_version_id": "file-1-v3",
|
|
"current_version_no": 3,
|
|
"version_count": 2,
|
|
}
|
|
|
|
|
|
def test_dataset_file_version_summary_uses_normalized_version_number_as_fallback() -> None:
|
|
summary = dataset_file_version_summary(
|
|
{
|
|
"active_version_id": "",
|
|
"current_version_id": None,
|
|
"version_no": 1,
|
|
"versions": "[]",
|
|
}
|
|
)
|
|
|
|
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)
|