diff --git a/backend/app/modules/data_process/store.py b/backend/app/modules/data_process/store.py index e501550..6acb405 100644 --- a/backend/app/modules/data_process/store.py +++ b/backend/app/modules/data_process/store.py @@ -3,11 +3,13 @@ from __future__ import annotations import hashlib import json import uuid +from collections.abc import Iterator, Sequence from contextlib import contextmanager -from datetime import date, datetime, timezone +from datetime import UTC, date, datetime +from decimal import Decimal from functools import lru_cache from pathlib import Path -from typing import Any, Iterator, Sequence +from typing import Any import psycopg from psycopg.rows import dict_row @@ -36,7 +38,7 @@ class InvalidStateError(DataProcessStoreError): def utcnow() -> str: - return datetime.now(timezone.utc).replace(microsecond=0).isoformat().replace("+00:00", "Z") + return datetime.now(UTC).replace(microsecond=0).isoformat().replace("+00:00", "Z") def new_id(prefix: str) -> str: @@ -65,6 +67,8 @@ def _json_value(value: Any, default: Any) -> Any: def _serialize_value(value: Any) -> Any: if isinstance(value, (datetime, date)): return value.isoformat().replace("+00:00", "Z") + if isinstance(value, Decimal): + return float(value) return value @@ -136,9 +140,8 @@ class DataProcessStore: """显式安装数据处理表;API 路由和应用启动流程不会调用此方法。""" schema_path = Path(__file__).resolve().parents[2] / "db" / "sql" / "002_data_process.sql" sql = schema_path.read_text(encoding="utf-8") - with self.connect() as conn: - with conn.cursor() as cursor: - cursor.execute(sql) + with self.connect() as conn, conn.cursor() as cursor: + cursor.execute(sql) def list_tasks( self, @@ -227,10 +230,30 @@ class DataProcessStore: def get_task(self, task_id: str, *, for_update: bool = False) -> dict[str, Any]: lock = " FOR UPDATE" if for_update else "" with self.connect() as conn: - row = conn.execute( - f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}", - (task_id,), - ).fetchone() + if for_update: + row = conn.execute( + f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}", + (task_id,), + ).fetchone() + else: + row = conn.execute( + """ + SELECT task.*, + (SELECT COUNT(*) FROM data_process_source_files source + WHERE source.task_id=task.id AND source.deleted_at IS NULL) + AS source_file_count, + (SELECT COUNT(*) FROM data_process_preview_items preview + WHERE preview.task_id=task.id) AS preview_count, + CASE + WHEN task.started_at IS NOT NULL AND task.completed_at IS NOT NULL + THEN EXTRACT(EPOCH FROM (task.completed_at - task.started_at)) + ELSE NULL + END AS duration_seconds + FROM data_process_tasks task + WHERE task.id=%s AND task.deleted_at IS NULL + """, + (task_id,), + ).fetchone() if not row: raise NotFoundError("data process task not found") return _decode_row(row) or {} diff --git a/backend/tests/test_data_process_store.py b/backend/tests/test_data_process_store.py index 0e7ba29..004504b 100644 --- a/backend/tests/test_data_process_store.py +++ b/backend/tests/test_data_process_store.py @@ -1,13 +1,27 @@ from __future__ import annotations +from decimal import Decimal + import pytest from app.modules.data_process.store import ( DataProcessStoreError, + _decode_row, _source_storage_descriptor, ) +def test_decode_row_serializes_postgres_numeric_values_as_json_numbers() -> None: + decoded = _decode_row( + { + "progress": Decimal("100.00"), + "duration_seconds": Decimal("389.000000"), + } + ) + + assert decoded == {"progress": 100.0, "duration_seconds": 389.0} + + def test_source_storage_descriptor_accepts_owned_local_and_legacy_db_references() -> None: task_id = "dpt_task" source_file_id = "dpsf_source"