2026-07-23 15:10:13 +08:00
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
|
|
import hashlib
|
|
|
|
|
|
import json
|
|
|
|
|
|
import uuid
|
2026-07-24 16:28:47 +08:00
|
|
|
|
from collections.abc import Iterator, Sequence
|
2026-07-23 15:10:13 +08:00
|
|
|
|
from contextlib import contextmanager
|
2026-07-24 16:28:47 +08:00
|
|
|
|
from datetime import UTC, date, datetime
|
|
|
|
|
|
from decimal import Decimal
|
2026-07-23 15:10:13 +08:00
|
|
|
|
from functools import lru_cache
|
|
|
|
|
|
from pathlib import Path
|
2026-07-24 16:28:47 +08:00
|
|
|
|
from typing import Any
|
2026-07-23 15:10:13 +08:00
|
|
|
|
|
|
|
|
|
|
import psycopg
|
|
|
|
|
|
from psycopg.rows import dict_row
|
|
|
|
|
|
|
|
|
|
|
|
from app.core.config import get_settings
|
2026-07-24 20:43:47 +08:00
|
|
|
|
from app.modules.data_process.algorithms import estimate_token_count, stable_split_assignments
|
2026-07-23 15:10:13 +08:00
|
|
|
|
|
|
|
|
|
|
TASK_STATUSES = {"pending", "running", "completed", "failed", "stopped"}
|
|
|
|
|
|
EDITABLE_STATUSES = {"pending", "failed", "stopped", "completed"}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class DataProcessStoreError(RuntimeError):
|
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class NotFoundError(DataProcessStoreError):
|
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ConflictError(DataProcessStoreError):
|
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class InvalidStateError(DataProcessStoreError):
|
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def utcnow() -> str:
|
2026-07-24 16:28:47 +08:00
|
|
|
|
return datetime.now(UTC).replace(microsecond=0).isoformat().replace("+00:00", "Z")
|
2026-07-23 15:10:13 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def new_id(prefix: str) -> str:
|
|
|
|
|
|
return f"{prefix}_{uuid.uuid4().hex[:20]}"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def json_dumps(value: Any) -> str:
|
|
|
|
|
|
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _database_url(value: str) -> str:
|
|
|
|
|
|
return value.replace("postgresql+psycopg://", "postgresql://")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _json_value(value: Any, default: Any) -> Any:
|
|
|
|
|
|
if value is None or value == "":
|
|
|
|
|
|
return default
|
|
|
|
|
|
if isinstance(value, (dict, list)):
|
|
|
|
|
|
return value
|
|
|
|
|
|
try:
|
|
|
|
|
|
return json.loads(value)
|
|
|
|
|
|
except (TypeError, json.JSONDecodeError):
|
|
|
|
|
|
return default
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _serialize_value(value: Any) -> Any:
|
|
|
|
|
|
if isinstance(value, (datetime, date)):
|
|
|
|
|
|
return value.isoformat().replace("+00:00", "Z")
|
2026-07-24 16:28:47 +08:00
|
|
|
|
if isinstance(value, Decimal):
|
|
|
|
|
|
return float(value)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
return value
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-07-24 11:27:51 +08:00
|
|
|
|
def _source_storage_descriptor(
|
|
|
|
|
|
payload: dict[str, Any],
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
file_id: str,
|
|
|
|
|
|
) -> tuple[str, dict[str, Any]]:
|
|
|
|
|
|
storage_object_id = str(
|
|
|
|
|
|
payload.get("storage_object_id")
|
|
|
|
|
|
or f"db://data-process/{task_id}/{file_id}/v1"
|
|
|
|
|
|
)
|
|
|
|
|
|
expected_local_prefix = f"local://data-process/{task_id}/{file_id}/v1/"
|
|
|
|
|
|
expected_database_reference = f"db://data-process/{task_id}/{file_id}/v1"
|
|
|
|
|
|
if storage_object_id.startswith(expected_local_prefix) and len(storage_object_id) > len(
|
|
|
|
|
|
expected_local_prefix
|
|
|
|
|
|
):
|
|
|
|
|
|
storage_backend = "local"
|
|
|
|
|
|
elif storage_object_id == expected_database_reference:
|
|
|
|
|
|
storage_backend = "database"
|
|
|
|
|
|
elif storage_object_id.startswith(("local://data-process/", "db://data-process/")):
|
|
|
|
|
|
raise DataProcessStoreError("source storage object owner mismatch")
|
|
|
|
|
|
else:
|
|
|
|
|
|
raise DataProcessStoreError("unsupported source storage object reference")
|
|
|
|
|
|
metadata = {
|
|
|
|
|
|
**(payload.get("metadata") or {}),
|
|
|
|
|
|
"storage_backend": storage_backend,
|
|
|
|
|
|
}
|
|
|
|
|
|
return storage_object_id, metadata
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-07-23 15:10:13 +08:00
|
|
|
|
def _decode_row(row: dict[str, Any] | None) -> dict[str, Any] | None:
|
|
|
|
|
|
if row is None:
|
|
|
|
|
|
return None
|
|
|
|
|
|
item = {key: _serialize_value(value) for key, value in row.items()}
|
|
|
|
|
|
for key, default in {
|
|
|
|
|
|
"config": {},
|
|
|
|
|
|
"metadata": {},
|
|
|
|
|
|
"quality_score": {},
|
|
|
|
|
|
"versions": [],
|
|
|
|
|
|
}.items():
|
|
|
|
|
|
if key in item:
|
|
|
|
|
|
item[key] = _json_value(item[key], default)
|
|
|
|
|
|
return item
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class DataProcessStore:
|
|
|
|
|
|
"""数据处理持久层。
|
|
|
|
|
|
|
|
|
|
|
|
构造函数不会连接数据库或执行迁移。部署方必须显式执行 002 SQL,
|
|
|
|
|
|
或在受控的管理命令中调用 :meth:`ensure_schema`,避免应用启动时
|
|
|
|
|
|
修改远程数据库。
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(self, database_url: str | None = None) -> None:
|
|
|
|
|
|
self.database_url = _database_url(database_url or get_settings().database_url)
|
|
|
|
|
|
|
|
|
|
|
|
@contextmanager
|
|
|
|
|
|
def connect(self) -> Iterator[psycopg.Connection[dict[str, Any]]]:
|
|
|
|
|
|
with psycopg.connect(self.database_url, row_factory=dict_row) as conn:
|
|
|
|
|
|
try:
|
|
|
|
|
|
yield conn
|
|
|
|
|
|
conn.commit()
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
conn.rollback()
|
|
|
|
|
|
raise
|
|
|
|
|
|
|
|
|
|
|
|
def ensure_schema(self) -> None:
|
|
|
|
|
|
"""显式安装数据处理表;API 路由和应用启动流程不会调用此方法。"""
|
|
|
|
|
|
schema_path = Path(__file__).resolve().parents[2] / "db" / "sql" / "002_data_process.sql"
|
|
|
|
|
|
sql = schema_path.read_text(encoding="utf-8")
|
2026-07-24 16:28:47 +08:00
|
|
|
|
with self.connect() as conn, conn.cursor() as cursor:
|
|
|
|
|
|
cursor.execute(sql)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
|
|
|
|
|
|
def list_tasks(
|
|
|
|
|
|
self,
|
|
|
|
|
|
*,
|
|
|
|
|
|
page: int = 1,
|
|
|
|
|
|
page_size: int = 20,
|
|
|
|
|
|
keyword: str | None = None,
|
|
|
|
|
|
status: str | None = None,
|
|
|
|
|
|
process_type: str | None = None,
|
|
|
|
|
|
tenant_id: str | None = None,
|
|
|
|
|
|
project_id: str | None = None,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
clauses = ["deleted_at IS NULL"]
|
|
|
|
|
|
params: list[Any] = []
|
|
|
|
|
|
if keyword:
|
|
|
|
|
|
clauses.append("(name ILIKE %s OR COALESCE(description, '') ILIKE %s)")
|
|
|
|
|
|
pattern = f"%{keyword.strip()}%"
|
|
|
|
|
|
params.extend([pattern, pattern])
|
|
|
|
|
|
if status:
|
|
|
|
|
|
clauses.append("status = %s")
|
|
|
|
|
|
params.append(status)
|
|
|
|
|
|
if process_type:
|
|
|
|
|
|
clauses.append("process_type = %s")
|
|
|
|
|
|
params.append(process_type)
|
|
|
|
|
|
if tenant_id:
|
|
|
|
|
|
clauses.append("tenant_id = %s")
|
|
|
|
|
|
params.append(tenant_id)
|
|
|
|
|
|
if project_id:
|
|
|
|
|
|
clauses.append("project_id = %s")
|
|
|
|
|
|
params.append(project_id)
|
|
|
|
|
|
where = " AND ".join(clauses)
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
total = conn.execute(
|
|
|
|
|
|
f"SELECT COUNT(*) AS count FROM data_process_tasks WHERE {where}", params
|
|
|
|
|
|
).fetchone()["count"]
|
|
|
|
|
|
rows = conn.execute(
|
|
|
|
|
|
f"""
|
|
|
|
|
|
SELECT * FROM data_process_tasks
|
|
|
|
|
|
WHERE {where}
|
|
|
|
|
|
ORDER BY created_at DESC, id DESC
|
|
|
|
|
|
LIMIT %s OFFSET %s
|
|
|
|
|
|
""",
|
|
|
|
|
|
[*params, page_size, (page - 1) * page_size],
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
return {
|
|
|
|
|
|
"items": [_decode_row(row) for row in rows],
|
|
|
|
|
|
"total": int(total),
|
|
|
|
|
|
"page": page,
|
|
|
|
|
|
"page_size": page_size,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def create_task(self, payload: dict[str, Any]) -> dict[str, Any]:
|
|
|
|
|
|
task_id = new_id("dpt")
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
try:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO data_process_tasks
|
|
|
|
|
|
(id, name, description, status, process_type, source_dataset_id, config,
|
|
|
|
|
|
progress, tenant_id, project_id, owner_id, created_by, updated_by,
|
|
|
|
|
|
created_at, updated_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, 'pending', %s, %s, %s, 0, %s, %s, %s, %s, %s, %s, %s)
|
|
|
|
|
|
RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
payload["name"],
|
|
|
|
|
|
payload.get("description") or "",
|
|
|
|
|
|
payload["process_type"],
|
|
|
|
|
|
payload.get("source_dataset_id"),
|
|
|
|
|
|
json_dumps(payload.get("config") or {}),
|
|
|
|
|
|
payload.get("tenant_id"),
|
|
|
|
|
|
payload.get("project_id"),
|
|
|
|
|
|
payload.get("owner_id"),
|
|
|
|
|
|
payload.get("created_by"),
|
|
|
|
|
|
payload.get("created_by"),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
except psycopg.errors.UniqueViolation as exc:
|
|
|
|
|
|
raise ConflictError("data process task name already exists") from exc
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
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:
|
2026-07-24 16:28:47 +08:00
|
|
|
|
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,
|
2026-07-25 17:04:14 +08:00
|
|
|
|
(SELECT COALESCE(json_agg(json_build_object(
|
|
|
|
|
|
'id', dataset.id,
|
|
|
|
|
|
'name', dataset.name,
|
|
|
|
|
|
'type', dataset.type,
|
|
|
|
|
|
'count', dataset.count,
|
|
|
|
|
|
'dataset_split', dataset.metadata::jsonb->>'dataset_split'
|
|
|
|
|
|
) ORDER BY CASE dataset.type
|
|
|
|
|
|
WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4 END), '[]'::json)
|
|
|
|
|
|
FROM datasets dataset
|
|
|
|
|
|
WHERE dataset.source_task_id=task.id AND dataset.source='task')
|
|
|
|
|
|
AS output_datasets,
|
2026-07-24 16:28:47 +08:00
|
|
|
|
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()
|
2026-07-23 15:10:13 +08:00
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("data process task not found")
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def _task_in_connection(
|
|
|
|
|
|
self,
|
|
|
|
|
|
conn: psycopg.Connection[dict[str, Any]],
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
*,
|
|
|
|
|
|
for_update: bool = False,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
lock = " FOR UPDATE" if for_update else ""
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}",
|
|
|
|
|
|
(task_id,),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("data process task not found")
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
@staticmethod
|
|
|
|
|
|
def _ensure_editable(task: dict[str, Any]) -> None:
|
|
|
|
|
|
if task["status"] not in EDITABLE_STATUSES:
|
|
|
|
|
|
raise InvalidStateError(f"task cannot be edited while status is {task['status']}")
|
|
|
|
|
|
if task.get("output_dataset_id"):
|
|
|
|
|
|
raise InvalidStateError("published task cannot be edited")
|
|
|
|
|
|
|
|
|
|
|
|
def update_task(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
|
|
|
|
|
allowed = {
|
|
|
|
|
|
"name",
|
|
|
|
|
|
"description",
|
|
|
|
|
|
"process_type",
|
|
|
|
|
|
"source_dataset_id",
|
|
|
|
|
|
}
|
|
|
|
|
|
values: dict[str, Any] = {key: value for key, value in payload.items() if key in allowed}
|
|
|
|
|
|
if payload.get("config") is not None:
|
|
|
|
|
|
values["config"] = json_dumps(payload["config"])
|
|
|
|
|
|
if not values:
|
|
|
|
|
|
return self.get_task(task_id)
|
|
|
|
|
|
try:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
invalidates_results = (
|
|
|
|
|
|
("config" in payload and payload.get("config") != task.get("config"))
|
|
|
|
|
|
or (
|
|
|
|
|
|
"process_type" in payload
|
|
|
|
|
|
and payload.get("process_type") != task.get("process_type")
|
|
|
|
|
|
)
|
|
|
|
|
|
or (
|
|
|
|
|
|
"source_dataset_id" in payload
|
|
|
|
|
|
and payload.get("source_dataset_id") != task.get("source_dataset_id")
|
|
|
|
|
|
)
|
|
|
|
|
|
)
|
|
|
|
|
|
if invalidates_results:
|
|
|
|
|
|
values.update(
|
|
|
|
|
|
{
|
|
|
|
|
|
"status": "pending",
|
|
|
|
|
|
"progress": 0,
|
|
|
|
|
|
"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,
|
|
|
|
|
|
}
|
|
|
|
|
|
)
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"DELETE FROM data_process_preview_items WHERE task_id=%s", (task_id,)
|
|
|
|
|
|
)
|
|
|
|
|
|
if (
|
|
|
|
|
|
"process_type" in payload
|
|
|
|
|
|
and payload.get("process_type") != task.get("process_type")
|
|
|
|
|
|
):
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_source_files
|
|
|
|
|
|
SET deleted_at=%s, updated_at=%s
|
|
|
|
|
|
WHERE task_id=%s AND deleted_at IS NULL
|
|
|
|
|
|
""",
|
|
|
|
|
|
(utcnow(), utcnow(), task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
values["input_count"] = 0
|
|
|
|
|
|
values["updated_at"] = utcnow()
|
|
|
|
|
|
assignments = ", ".join(f"{key}=%s" for key in values)
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
f"UPDATE data_process_tasks SET {assignments} WHERE id=%s RETURNING *",
|
|
|
|
|
|
[*values.values(), task_id],
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
except psycopg.errors.UniqueViolation as exc:
|
|
|
|
|
|
raise ConflictError("data process task name already exists") from exc
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def delete_task(self, task_id: str, *, deleted_by: str | None = None) -> None:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if task["status"] == "running":
|
|
|
|
|
|
raise InvalidStateError("running task must be stopped before deletion")
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET deleted_at=%s, deleted_by=%s, updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(now, deleted_by, now, task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def list_source_files(self, task_id: str) -> list[dict[str, Any]]:
|
|
|
|
|
|
self.get_task(task_id)
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
rows = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT id, task_id, storage_object_id, name, size_bytes, record_count,
|
|
|
|
|
|
file_format, checksum_sha256, version_no, content_preview, metadata,
|
|
|
|
|
|
tenant_id, project_id, created_by, created_at, updated_at
|
|
|
|
|
|
FROM data_process_source_files
|
|
|
|
|
|
WHERE task_id=%s AND deleted_at IS NULL
|
|
|
|
|
|
ORDER BY created_at, id
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id,),
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
return [_decode_row(row) or {} for row in rows]
|
|
|
|
|
|
|
|
|
|
|
|
def add_source_file(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
*,
|
|
|
|
|
|
name: str,
|
|
|
|
|
|
content: str,
|
|
|
|
|
|
raw_size: int,
|
|
|
|
|
|
checksum_sha256: str,
|
|
|
|
|
|
file_format: str,
|
|
|
|
|
|
record_count: int,
|
|
|
|
|
|
metadata: dict[str, Any] | None = None,
|
|
|
|
|
|
created_by: str | None = None,
|
2026-07-24 11:27:51 +08:00
|
|
|
|
source_file_id: str | None = None,
|
|
|
|
|
|
storage_object_id: str | None = None,
|
2026-07-23 15:10:13 +08:00
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
return self.add_source_files(
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
[
|
|
|
|
|
|
{
|
|
|
|
|
|
"name": name,
|
|
|
|
|
|
"content": content,
|
|
|
|
|
|
"raw_size": raw_size,
|
|
|
|
|
|
"checksum_sha256": checksum_sha256,
|
|
|
|
|
|
"file_format": file_format,
|
|
|
|
|
|
"record_count": record_count,
|
|
|
|
|
|
"metadata": metadata or {},
|
|
|
|
|
|
"created_by": created_by,
|
2026-07-24 11:27:51 +08:00
|
|
|
|
"id": source_file_id,
|
|
|
|
|
|
"storage_object_id": storage_object_id,
|
2026-07-23 15:10:13 +08:00
|
|
|
|
}
|
|
|
|
|
|
],
|
|
|
|
|
|
)[0]
|
|
|
|
|
|
|
|
|
|
|
|
def add_source_files(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
files: Sequence[dict[str, Any]],
|
|
|
|
|
|
) -> list[dict[str, Any]]:
|
|
|
|
|
|
"""在同一事务中登记一个上传批次,任一文件失败则全部回滚。"""
|
|
|
|
|
|
|
|
|
|
|
|
if not files:
|
|
|
|
|
|
raise DataProcessStoreError("at least one source file is required")
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
created: list[dict[str, Any]] = []
|
|
|
|
|
|
try:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
for payload in files:
|
2026-07-24 11:27:51 +08:00
|
|
|
|
file_id = str(payload.get("id") or new_id("dpsf"))
|
|
|
|
|
|
storage_object_id, metadata_payload = _source_storage_descriptor(
|
|
|
|
|
|
payload,
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
file_id,
|
|
|
|
|
|
)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO data_process_source_files
|
|
|
|
|
|
(id, task_id, storage_object_id, name, size_bytes, record_count,
|
|
|
|
|
|
file_format, checksum_sha256, version_no, content, content_preview,
|
|
|
|
|
|
metadata, tenant_id, project_id,
|
|
|
|
|
|
created_by, created_at, updated_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, 1, %s, %s, %s, %s,
|
|
|
|
|
|
%s, %s, %s, %s)
|
|
|
|
|
|
RETURNING id, task_id, storage_object_id, name, size_bytes, record_count,
|
|
|
|
|
|
file_format, checksum_sha256, version_no, content_preview, metadata,
|
|
|
|
|
|
tenant_id, project_id, created_by, created_at, updated_at
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
file_id,
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
storage_object_id,
|
|
|
|
|
|
payload["name"],
|
|
|
|
|
|
payload["raw_size"],
|
|
|
|
|
|
payload["record_count"],
|
|
|
|
|
|
payload["file_format"],
|
|
|
|
|
|
payload["checksum_sha256"],
|
|
|
|
|
|
payload["content"],
|
|
|
|
|
|
str(payload["content"])[:2000],
|
|
|
|
|
|
json_dumps(metadata_payload),
|
|
|
|
|
|
task.get("tenant_id"),
|
|
|
|
|
|
task.get("project_id"),
|
|
|
|
|
|
payload.get("created_by") or task.get("created_by"),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
created.append(_decode_row(row) or {})
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_preview_items WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='pending', progress=0, output_count=0, filtered_count=0,
|
|
|
|
|
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
|
|
|
|
|
generation_run_id=NULL, started_at=NULL, completed_at=NULL, input_count=(
|
|
|
|
|
|
SELECT COALESCE(SUM(record_count), 0)
|
|
|
|
|
|
FROM data_process_source_files
|
|
|
|
|
|
WHERE task_id=%s AND deleted_at IS NULL
|
|
|
|
|
|
), updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id, now, task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
except psycopg.errors.UniqueViolation as exc:
|
|
|
|
|
|
raise ConflictError(
|
|
|
|
|
|
"the same source file content is already attached to this task"
|
|
|
|
|
|
) from exc
|
|
|
|
|
|
return created
|
|
|
|
|
|
|
|
|
|
|
|
def get_source_file(
|
|
|
|
|
|
self, task_id: str, file_id: str, *, include_content: bool = True
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
# 先验证父任务仍然可见,避免软删除任务后通过已知文件 ID 读取正文。
|
|
|
|
|
|
self.get_task(task_id)
|
|
|
|
|
|
content_column = ", content" if include_content else ""
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
f"""
|
|
|
|
|
|
SELECT id, task_id, storage_object_id, name, size_bytes, record_count,
|
|
|
|
|
|
file_format, checksum_sha256, version_no, content_preview, metadata,
|
|
|
|
|
|
tenant_id, project_id, created_by, created_at, updated_at{content_column}
|
|
|
|
|
|
FROM data_process_source_files
|
|
|
|
|
|
WHERE id=%s AND task_id=%s AND deleted_at IS NULL
|
|
|
|
|
|
""",
|
|
|
|
|
|
(file_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("source file not found")
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def source_content_window(
|
|
|
|
|
|
self, task_id: str, file_id: str, offset: int, limit: int
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
source_file = self.get_source_file(task_id, file_id, include_content=True)
|
|
|
|
|
|
content = str(source_file.pop("content", ""))
|
|
|
|
|
|
window = content[offset : offset + limit]
|
|
|
|
|
|
return {
|
|
|
|
|
|
"file": source_file,
|
|
|
|
|
|
"content": window,
|
|
|
|
|
|
"offset": offset,
|
|
|
|
|
|
"limit": limit,
|
|
|
|
|
|
"total_chars": len(content),
|
|
|
|
|
|
"has_more": offset + len(window) < len(content),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def source_content_lines(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
file_id: str,
|
|
|
|
|
|
start_line: int,
|
|
|
|
|
|
line_count: int,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
source_file = self.get_source_file(task_id, file_id, include_content=True)
|
|
|
|
|
|
content = str(source_file.pop("content", ""))
|
|
|
|
|
|
lines = content.splitlines(keepends=True)
|
|
|
|
|
|
start_index = min(len(lines), start_line - 1)
|
|
|
|
|
|
selected = lines[start_index : start_index + line_count]
|
|
|
|
|
|
end_line = start_index + len(selected)
|
|
|
|
|
|
return {
|
|
|
|
|
|
"file": source_file,
|
|
|
|
|
|
"content": "".join(selected),
|
|
|
|
|
|
"start_line": start_line,
|
|
|
|
|
|
"end_line": end_line,
|
|
|
|
|
|
"line_count": len(selected),
|
|
|
|
|
|
"total_lines": len(lines),
|
|
|
|
|
|
"has_more": end_line < len(lines),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def delete_source_file(self, task_id: str, file_id: str) -> None:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_source_files
|
|
|
|
|
|
SET deleted_at=%s, updated_at=%s
|
|
|
|
|
|
WHERE id=%s AND task_id=%s AND deleted_at IS NULL
|
|
|
|
|
|
RETURNING id
|
|
|
|
|
|
""",
|
|
|
|
|
|
(utcnow(), utcnow(), file_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("source file not found")
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"DELETE FROM data_process_preview_items WHERE source_file_id=%s", (file_id,)
|
|
|
|
|
|
)
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='pending', progress=0, output_count=0, filtered_count=0,
|
|
|
|
|
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
|
|
|
|
|
input_count=(SELECT COALESCE(SUM(record_count), 0)
|
|
|
|
|
|
FROM data_process_source_files
|
|
|
|
|
|
WHERE task_id=%s AND deleted_at IS NULL),
|
|
|
|
|
|
updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id, utcnow(), task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def replace_preview_items(
|
2026-07-24 11:27:51 +08:00
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
items: Sequence[dict[str, Any]],
|
|
|
|
|
|
*,
|
|
|
|
|
|
source_file_ids: Sequence[str] | None = None,
|
2026-07-23 15:10:13 +08:00
|
|
|
|
) -> list[dict[str, Any]]:
|
2026-07-24 11:27:51 +08:00
|
|
|
|
selected_ids = (
|
|
|
|
|
|
list(dict.fromkeys(str(file_id) for file_id in source_file_ids))
|
|
|
|
|
|
if source_file_ids is not None
|
|
|
|
|
|
else None
|
|
|
|
|
|
)
|
|
|
|
|
|
if selected_ids is not None:
|
|
|
|
|
|
if not selected_ids or any(not file_id for file_id in selected_ids):
|
|
|
|
|
|
raise ValueError("source_file_ids must contain non-empty ids")
|
|
|
|
|
|
selected_set = set(selected_ids)
|
|
|
|
|
|
unexpected = {
|
|
|
|
|
|
str(item.get("source_file_id") or "")
|
|
|
|
|
|
for item in items
|
|
|
|
|
|
if str(item.get("source_file_id") or "") not in selected_set
|
|
|
|
|
|
}
|
|
|
|
|
|
if unexpected:
|
|
|
|
|
|
raise ValueError("preview items contain an unselected source file")
|
|
|
|
|
|
|
2026-07-23 15:10:13 +08:00
|
|
|
|
now = utcnow()
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
2026-07-24 11:27:51 +08:00
|
|
|
|
if selected_ids is None:
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"DELETE FROM data_process_preview_items WHERE task_id=%s", (task_id,)
|
|
|
|
|
|
)
|
|
|
|
|
|
else:
|
|
|
|
|
|
rows = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT id FROM data_process_source_files
|
|
|
|
|
|
WHERE task_id=%s AND deleted_at IS NULL AND id=ANY(%s)
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id, selected_ids),
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
found = {str(row["id"]) for row in rows}
|
|
|
|
|
|
missing = set(selected_ids) - found
|
|
|
|
|
|
if missing:
|
|
|
|
|
|
raise NotFoundError(
|
|
|
|
|
|
f"source files not found: {', '.join(sorted(missing))}"
|
|
|
|
|
|
)
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
DELETE FROM data_process_preview_items
|
|
|
|
|
|
WHERE task_id=%s AND source_file_id=ANY(%s)
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id, selected_ids),
|
|
|
|
|
|
)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
created: list[dict[str, Any]] = []
|
|
|
|
|
|
for item in items:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO data_process_preview_items
|
|
|
|
|
|
(id, task_id, source_file_id, original_content, edited_content,
|
|
|
|
|
|
source_start, source_end, source_start_line, source_end_line,
|
|
|
|
|
|
token_count, status, quality_score, created_at, updated_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
|
|
|
|
RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
item.get("id") or new_id("dpp"),
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
item.get("source_file_id"),
|
|
|
|
|
|
item.get("original_content") or "",
|
|
|
|
|
|
item.get("edited_content", item.get("original_content") or ""),
|
|
|
|
|
|
item.get("source_start"),
|
|
|
|
|
|
item.get("source_end"),
|
|
|
|
|
|
item.get("source_start_line"),
|
|
|
|
|
|
item.get("source_end_line"),
|
|
|
|
|
|
max(0, int(item.get("token_count") or 0)),
|
|
|
|
|
|
item.get("status") or "original",
|
|
|
|
|
|
json_dumps(item.get("quality_score") or {}),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
created.append(_decode_row(row) or {})
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='pending', progress=20, output_count=0, filtered_count=0,
|
|
|
|
|
|
duplicate_count=0, error_count=0, failure_reason=NULL, updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(now, task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
return created
|
|
|
|
|
|
|
|
|
|
|
|
def list_preview_items(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
*,
|
|
|
|
|
|
source_file_id: str | None = None,
|
|
|
|
|
|
page: int = 1,
|
|
|
|
|
|
page_size: int = 200,
|
|
|
|
|
|
keyword: str | None = None,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
self.get_task(task_id)
|
|
|
|
|
|
clauses = ["task_id=%s"]
|
|
|
|
|
|
params: list[Any] = [task_id]
|
|
|
|
|
|
if source_file_id:
|
|
|
|
|
|
clauses.append("source_file_id=%s")
|
|
|
|
|
|
params.append(source_file_id)
|
|
|
|
|
|
if keyword:
|
|
|
|
|
|
clauses.append("(original_content ILIKE %s OR edited_content ILIKE %s)")
|
|
|
|
|
|
pattern = f"%{keyword.strip()}%"
|
|
|
|
|
|
params.extend([pattern, pattern])
|
|
|
|
|
|
where = " AND ".join(clauses)
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
total = conn.execute(
|
|
|
|
|
|
f"SELECT COUNT(*) AS count FROM data_process_preview_items WHERE {where}", params
|
|
|
|
|
|
).fetchone()["count"]
|
|
|
|
|
|
rows = conn.execute(
|
|
|
|
|
|
f"""
|
|
|
|
|
|
SELECT * FROM data_process_preview_items
|
|
|
|
|
|
WHERE {where}
|
|
|
|
|
|
ORDER BY source_file_id NULLS LAST, source_start NULLS LAST, created_at, id
|
|
|
|
|
|
LIMIT %s OFFSET %s
|
|
|
|
|
|
""",
|
|
|
|
|
|
[*params, page_size, (page - 1) * page_size],
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
return {
|
|
|
|
|
|
"items": [_decode_row(row) for row in rows],
|
|
|
|
|
|
"total": int(total),
|
|
|
|
|
|
"page": page,
|
|
|
|
|
|
"page_size": page_size,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def get_preview_item(self, task_id: str, preview_id: str) -> dict[str, Any]:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"SELECT * FROM data_process_preview_items WHERE id=%s AND task_id=%s",
|
|
|
|
|
|
(preview_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("preview item not found")
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def create_preview_item(self, task_id: str, item: dict[str, Any]) -> dict[str, Any]:
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
if item.get("source_file_id"):
|
|
|
|
|
|
source = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT id FROM data_process_source_files
|
|
|
|
|
|
WHERE id=%s AND task_id=%s AND deleted_at IS NULL
|
|
|
|
|
|
""",
|
|
|
|
|
|
(item["source_file_id"], task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not source:
|
|
|
|
|
|
raise NotFoundError("source file not found")
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO data_process_preview_items
|
|
|
|
|
|
(id, task_id, source_file_id, original_content, edited_content,
|
|
|
|
|
|
source_start, source_end, source_start_line, source_end_line,
|
|
|
|
|
|
token_count, status, quality_score, created_at, updated_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
|
|
|
|
RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
new_id("dpp"),
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
item.get("source_file_id"),
|
|
|
|
|
|
item.get("original_content") or "",
|
|
|
|
|
|
item.get("edited_content") or "",
|
|
|
|
|
|
item.get("source_start"),
|
|
|
|
|
|
item.get("source_end"),
|
|
|
|
|
|
item.get("source_start_line"),
|
|
|
|
|
|
item.get("source_end_line"),
|
|
|
|
|
|
max(0, int(item.get("token_count") or 0)),
|
|
|
|
|
|
item.get("status") or "manual",
|
|
|
|
|
|
json_dumps(item.get("quality_score") or {}),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
self._invalidate_results(conn, task_id, now)
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def update_preview_item(
|
|
|
|
|
|
self, task_id: str, preview_id: str, payload: dict[str, Any]
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
existing = conn.execute(
|
|
|
|
|
|
"SELECT * FROM data_process_preview_items WHERE id=%s AND task_id=%s",
|
|
|
|
|
|
(preview_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not existing:
|
|
|
|
|
|
raise NotFoundError("preview item not found")
|
|
|
|
|
|
expected_updated_at = payload.get("expected_updated_at")
|
|
|
|
|
|
current_updated_at = _serialize_value(existing.get("updated_at"))
|
|
|
|
|
|
if expected_updated_at and expected_updated_at != current_updated_at:
|
|
|
|
|
|
raise ConflictError("preview item was modified by another request")
|
|
|
|
|
|
edited = payload["edited_content"]
|
|
|
|
|
|
status = payload.get("status")
|
|
|
|
|
|
if not status:
|
|
|
|
|
|
if not edited.strip():
|
|
|
|
|
|
status = "invalid"
|
|
|
|
|
|
elif edited == existing["original_content"]:
|
|
|
|
|
|
status = "original"
|
|
|
|
|
|
else:
|
|
|
|
|
|
status = "modified"
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_preview_items
|
|
|
|
|
|
SET edited_content=%s, token_count=%s, status=%s, quality_score=%s,
|
|
|
|
|
|
updated_at=%s
|
|
|
|
|
|
WHERE id=%s AND task_id=%s
|
|
|
|
|
|
RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
edited,
|
|
|
|
|
|
estimate_token_count(edited),
|
|
|
|
|
|
status,
|
|
|
|
|
|
json_dumps(payload.get("quality_score") or {}),
|
|
|
|
|
|
now,
|
|
|
|
|
|
preview_id,
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
self._invalidate_results(conn, task_id, now)
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def delete_preview_item(self, task_id: str, preview_id: str) -> None:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
self._ensure_editable(task)
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"DELETE FROM data_process_preview_items WHERE id=%s AND task_id=%s RETURNING id",
|
|
|
|
|
|
(preview_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("preview item not found")
|
|
|
|
|
|
self._invalidate_results(conn, task_id, utcnow())
|
|
|
|
|
|
|
|
|
|
|
|
def _invalidate_results(
|
|
|
|
|
|
self,
|
|
|
|
|
|
conn: psycopg.Connection[dict[str, Any]],
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
now: str,
|
|
|
|
|
|
) -> None:
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='pending', progress=20, output_count=0, filtered_count=0,
|
|
|
|
|
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
|
|
|
|
|
generation_run_id=NULL, updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(now, task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def start_generation(self, task_id: str, *, replace_existing: bool = True) -> dict[str, Any]:
|
|
|
|
|
|
if not replace_existing:
|
|
|
|
|
|
raise DataProcessStoreError("incremental generation is not supported")
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if task.get("output_dataset_id"):
|
|
|
|
|
|
raise InvalidStateError("published task cannot be regenerated")
|
|
|
|
|
|
if task["status"] == "running":
|
|
|
|
|
|
raise ConflictError("data process task is already running")
|
|
|
|
|
|
preview_count = conn.execute(
|
|
|
|
|
|
"SELECT COUNT(*) AS count FROM data_process_preview_items WHERE task_id=%s",
|
|
|
|
|
|
(task_id,),
|
|
|
|
|
|
).fetchone()["count"]
|
|
|
|
|
|
if not preview_count:
|
|
|
|
|
|
raise InvalidStateError("preview must be built before generation")
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
generation_run_id = new_id("dprun")
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='running', progress=30, failure_reason=NULL, started_at=%s,
|
|
|
|
|
|
completed_at=NULL, filtered_count=0, duplicate_count=0, error_count=0,
|
|
|
|
|
|
generation_run_id=%s, updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(now, generation_run_id, now, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def stop_task(self, task_id: str) -> dict[str, Any]:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if task["status"] != "running":
|
|
|
|
|
|
raise InvalidStateError("only a running task can be stopped")
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='stopped', failure_reason=NULL, generation_run_id=NULL,
|
|
|
|
|
|
updated_at=%s
|
|
|
|
|
|
WHERE id=%s RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(now, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def generation_is_running(self, task_id: str, generation_run_id: str) -> bool:
|
|
|
|
|
|
task = self.get_task(task_id)
|
|
|
|
|
|
return (
|
|
|
|
|
|
task["status"] == "running"
|
|
|
|
|
|
and task.get("generation_run_id") == generation_run_id
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def update_generation_progress(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
generation_run_id: str,
|
|
|
|
|
|
processed_count: int,
|
|
|
|
|
|
total_count: int,
|
|
|
|
|
|
) -> bool:
|
|
|
|
|
|
ratio = processed_count / max(1, total_count)
|
|
|
|
|
|
progress = min(95.0, 30.0 + ratio * 65.0)
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET progress=%s, updated_at=%s
|
|
|
|
|
|
WHERE id=%s AND status='running' AND generation_run_id=%s
|
|
|
|
|
|
RETURNING id
|
|
|
|
|
|
""",
|
|
|
|
|
|
(progress, utcnow(), task_id, generation_run_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
return row is not None
|
|
|
|
|
|
|
|
|
|
|
|
def complete_generation(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
results: Sequence[dict[str, Any]],
|
|
|
|
|
|
*,
|
|
|
|
|
|
generation_run_id: str,
|
|
|
|
|
|
filtered_count: int = 0,
|
|
|
|
|
|
duplicate_count: int = 0,
|
|
|
|
|
|
error_count: int = 0,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if (
|
|
|
|
|
|
task["status"] != "running"
|
|
|
|
|
|
or task.get("generation_run_id") != generation_run_id
|
|
|
|
|
|
):
|
|
|
|
|
|
return task
|
|
|
|
|
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
|
|
|
|
|
for result in results:
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO data_process_results
|
|
|
|
|
|
(id, task_id, preview_item_id, instruction, input, output,
|
|
|
|
|
|
original_instruction, original_input, original_output, status, error,
|
|
|
|
|
|
split, quality_score, created_at, updated_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
result.get("id") or new_id("dpr"),
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
result.get("preview_item_id"),
|
|
|
|
|
|
result.get("instruction") or "",
|
|
|
|
|
|
result.get("input") or "",
|
|
|
|
|
|
result.get("output") or "",
|
|
|
|
|
|
result.get("original_instruction", result.get("instruction") or ""),
|
|
|
|
|
|
result.get("original_input", result.get("input") or ""),
|
|
|
|
|
|
result.get("original_output", result.get("output") or ""),
|
|
|
|
|
|
result.get("status") or "valid",
|
|
|
|
|
|
result.get("error"),
|
|
|
|
|
|
result.get("split"),
|
|
|
|
|
|
json_dumps(result.get("quality_score") or {}),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
)
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='completed', progress=100, output_count=%s, filtered_count=%s,
|
|
|
|
|
|
duplicate_count=%s, error_count=%s, failure_reason=NULL,
|
|
|
|
|
|
completed_at=%s, generation_run_id=NULL, updated_at=%s
|
|
|
|
|
|
WHERE id=%s AND generation_run_id=%s RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
len(results),
|
|
|
|
|
|
filtered_count,
|
|
|
|
|
|
duplicate_count,
|
|
|
|
|
|
error_count,
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
generation_run_id,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def mark_failed(
|
|
|
|
|
|
self, task_id: str, reason: str, *, generation_run_id: str
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if (
|
|
|
|
|
|
task["status"] != "running"
|
|
|
|
|
|
or task.get("generation_run_id") != generation_run_id
|
|
|
|
|
|
):
|
|
|
|
|
|
return task
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET status='failed', failure_reason=%s, completed_at=%s,
|
|
|
|
|
|
generation_run_id=NULL, updated_at=%s
|
|
|
|
|
|
WHERE id=%s AND generation_run_id=%s RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(reason[:4000], now, now, task_id, generation_run_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def progress(self, task_id: str) -> dict[str, Any]:
|
|
|
|
|
|
task = self.get_task(task_id)
|
|
|
|
|
|
return {
|
|
|
|
|
|
"task_id": task["id"],
|
|
|
|
|
|
"status": task["status"],
|
|
|
|
|
|
"progress": float(task.get("progress") or 0),
|
|
|
|
|
|
"input_count": int(task.get("input_count") or 0),
|
|
|
|
|
|
"output_count": int(task.get("output_count") or 0),
|
|
|
|
|
|
"filtered_count": int(task.get("filtered_count") or 0),
|
|
|
|
|
|
"duplicate_count": int(task.get("duplicate_count") or 0),
|
|
|
|
|
|
"error_count": int(task.get("error_count") or 0),
|
|
|
|
|
|
"failure_reason": task.get("failure_reason"),
|
|
|
|
|
|
"started_at": task.get("started_at"),
|
|
|
|
|
|
"completed_at": task.get("completed_at"),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def list_results(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
*,
|
|
|
|
|
|
page: int = 1,
|
|
|
|
|
|
page_size: int = 100,
|
|
|
|
|
|
status: str | None = None,
|
|
|
|
|
|
split: str | None = None,
|
|
|
|
|
|
keyword: str | None = None,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
self.get_task(task_id)
|
|
|
|
|
|
clauses = ["task_id=%s"]
|
|
|
|
|
|
params: list[Any] = [task_id]
|
|
|
|
|
|
if status:
|
|
|
|
|
|
clauses.append("status=%s")
|
|
|
|
|
|
params.append(status)
|
|
|
|
|
|
if split:
|
|
|
|
|
|
clauses.append("split=%s")
|
|
|
|
|
|
params.append(split)
|
|
|
|
|
|
if keyword:
|
|
|
|
|
|
clauses.append("(instruction ILIKE %s OR input ILIKE %s OR output ILIKE %s)")
|
|
|
|
|
|
pattern = f"%{keyword.strip()}%"
|
|
|
|
|
|
params.extend([pattern, pattern, pattern])
|
|
|
|
|
|
where = " AND ".join(clauses)
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
total = conn.execute(
|
|
|
|
|
|
f"SELECT COUNT(*) AS count FROM data_process_results WHERE {where}", params
|
|
|
|
|
|
).fetchone()["count"]
|
|
|
|
|
|
rows = conn.execute(
|
|
|
|
|
|
f"""
|
|
|
|
|
|
SELECT * FROM data_process_results WHERE {where}
|
|
|
|
|
|
ORDER BY created_at, id LIMIT %s OFFSET %s
|
|
|
|
|
|
""",
|
|
|
|
|
|
[*params, page_size, (page - 1) * page_size],
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
return {
|
|
|
|
|
|
"items": [_decode_row(row) for row in rows],
|
|
|
|
|
|
"total": int(total),
|
|
|
|
|
|
"page": page,
|
|
|
|
|
|
"page_size": page_size,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
def get_result(self, task_id: str, result_id: str) -> dict[str, Any]:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"SELECT * FROM data_process_results WHERE id=%s AND task_id=%s",
|
|
|
|
|
|
(result_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("data process result not found")
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def update_result(
|
|
|
|
|
|
self, task_id: str, result_id: str, payload: dict[str, Any]
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
allowed = {"instruction", "input", "output", "quality_score"}
|
|
|
|
|
|
values = {key: value for key, value in payload.items() if key in allowed}
|
|
|
|
|
|
if "quality_score" in values:
|
|
|
|
|
|
values["quality_score"] = json_dumps(values["quality_score"])
|
|
|
|
|
|
if not values:
|
|
|
|
|
|
raise DataProcessStoreError("no result fields supplied")
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if task["status"] == "running":
|
|
|
|
|
|
raise InvalidStateError("results cannot be edited while generation is running")
|
|
|
|
|
|
if task.get("output_dataset_id"):
|
|
|
|
|
|
raise InvalidStateError("published results cannot be edited")
|
|
|
|
|
|
current = conn.execute(
|
|
|
|
|
|
"SELECT * FROM data_process_results WHERE id=%s AND task_id=%s",
|
|
|
|
|
|
(result_id, task_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not current:
|
|
|
|
|
|
raise NotFoundError("data process result not found")
|
|
|
|
|
|
expected_updated_at = payload.get("expected_updated_at")
|
|
|
|
|
|
current_updated_at = _serialize_value(current.get("updated_at"))
|
|
|
|
|
|
if expected_updated_at and expected_updated_at != current_updated_at:
|
|
|
|
|
|
raise ConflictError("data process result was modified by another request")
|
|
|
|
|
|
merged = {**current, **values}
|
|
|
|
|
|
quality = payload.get("quality_score") or {}
|
|
|
|
|
|
hard_valid = bool(
|
|
|
|
|
|
str(merged.get("instruction") or "").strip()
|
|
|
|
|
|
and str(merged.get("output") or "").strip()
|
|
|
|
|
|
)
|
|
|
|
|
|
quality_valid = bool(quality.get("is_valid", hard_valid))
|
|
|
|
|
|
changed = any(
|
|
|
|
|
|
str(merged.get(field) or "")
|
|
|
|
|
|
!= str(merged.get(f"original_{field}") or "")
|
|
|
|
|
|
for field in ("instruction", "input", "output")
|
|
|
|
|
|
)
|
|
|
|
|
|
status = "invalid" if not hard_valid or not quality_valid else (
|
|
|
|
|
|
"modified" if changed else "valid"
|
|
|
|
|
|
)
|
|
|
|
|
|
values["status"] = status
|
|
|
|
|
|
flags = quality.get("flags") if isinstance(quality, dict) else None
|
|
|
|
|
|
values["error"] = ", ".join(str(flag) for flag in flags or []) or (
|
|
|
|
|
|
"quality validation failed" if status == "invalid" else None
|
|
|
|
|
|
)
|
|
|
|
|
|
values["updated_at"] = utcnow()
|
|
|
|
|
|
assignments = ", ".join(f"{key}=%s" for key in values)
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
f"""UPDATE data_process_results SET {assignments}
|
|
|
|
|
|
WHERE id=%s AND task_id=%s RETURNING *""",
|
|
|
|
|
|
[*values.values(), result_id, task_id],
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET error_count=(
|
|
|
|
|
|
SELECT COUNT(*) FROM data_process_results
|
|
|
|
|
|
WHERE task_id=%s AND status='invalid'
|
|
|
|
|
|
), updated_at=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id, utcnow(), task_id),
|
|
|
|
|
|
)
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def get_generation_model(self, model_id: str) -> dict[str, Any]:
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT id, name, type, purpose, model_source, description, path,
|
|
|
|
|
|
api_url, api_key, online_model_name, create_time
|
|
|
|
|
|
FROM models WHERE id=%s
|
|
|
|
|
|
""",
|
|
|
|
|
|
(model_id,),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
if not row:
|
|
|
|
|
|
raise NotFoundError("generation model not found")
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def save_generation_model_snapshot(
|
|
|
|
|
|
self,
|
|
|
|
|
|
task_id: str,
|
|
|
|
|
|
model_snapshot: dict[str, Any],
|
|
|
|
|
|
*,
|
|
|
|
|
|
generation_run_id: str,
|
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
|
# API 密钥仅用于本次调用,绝不能进入任务配置、详情响应或审计快照。
|
|
|
|
|
|
safe_snapshot = {
|
|
|
|
|
|
key: value for key, value in model_snapshot.items() if key != "api_key"
|
|
|
|
|
|
}
|
|
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if (
|
|
|
|
|
|
task["status"] != "running"
|
|
|
|
|
|
or task.get("generation_run_id") != generation_run_id
|
|
|
|
|
|
):
|
|
|
|
|
|
raise InvalidStateError("generation run is no longer active")
|
|
|
|
|
|
config = dict(task.get("config") or {})
|
|
|
|
|
|
config["generation_model_snapshot"] = safe_snapshot
|
|
|
|
|
|
row = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks SET config=%s, updated_at=%s
|
|
|
|
|
|
WHERE id=%s AND generation_run_id=%s RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(json_dumps(config), utcnow(), task_id, generation_run_id),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
return _decode_row(row) or {}
|
|
|
|
|
|
|
|
|
|
|
|
def publish(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
2026-07-25 17:04:14 +08:00
|
|
|
|
"""按精确配额发布训练、验证、测试三个独立数据集。"""
|
2026-07-23 15:10:13 +08:00
|
|
|
|
with self.connect() as conn:
|
|
|
|
|
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
|
|
|
|
|
if task["status"] != "completed":
|
|
|
|
|
|
raise InvalidStateError("only a completed task can be published")
|
|
|
|
|
|
rows = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT * FROM data_process_results
|
|
|
|
|
|
WHERE task_id=%s ORDER BY created_at, id
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id,),
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
if not rows:
|
|
|
|
|
|
raise InvalidStateError("task has no results to publish")
|
|
|
|
|
|
invalid_count = sum(
|
|
|
|
|
|
1
|
|
|
|
|
|
for row in rows
|
|
|
|
|
|
if row["status"] == "invalid"
|
|
|
|
|
|
or not str(row.get("instruction") or "").strip()
|
|
|
|
|
|
or not str(row.get("output") or "").strip()
|
|
|
|
|
|
)
|
|
|
|
|
|
if invalid_count:
|
|
|
|
|
|
raise InvalidStateError(f"task contains {invalid_count} invalid results")
|
|
|
|
|
|
|
|
|
|
|
|
now = utcnow()
|
|
|
|
|
|
requested_split = payload.get("split") or {
|
|
|
|
|
|
"train": 80,
|
|
|
|
|
|
"validation": 10,
|
|
|
|
|
|
"test": 10,
|
|
|
|
|
|
}
|
2026-07-24 20:43:47 +08:00
|
|
|
|
assignments = stable_split_assignments(
|
|
|
|
|
|
[str(row["id"]) for row in rows],
|
|
|
|
|
|
requested_split,
|
|
|
|
|
|
seed=task_id,
|
|
|
|
|
|
)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
records = [
|
|
|
|
|
|
{
|
|
|
|
|
|
"instruction": row["instruction"],
|
|
|
|
|
|
"input": row["input"],
|
|
|
|
|
|
"output": row["output"],
|
2026-07-24 20:43:47 +08:00
|
|
|
|
"split": assignment,
|
2026-07-23 15:10:13 +08:00
|
|
|
|
}
|
2026-07-24 20:43:47 +08:00
|
|
|
|
for row, assignment in zip(rows, assignments, strict=True)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
]
|
2026-07-24 20:43:47 +08:00
|
|
|
|
split_order = ("train", "validation", "test")
|
|
|
|
|
|
split_counts = {
|
|
|
|
|
|
split_name: assignments.count(split_name) for split_name in split_order
|
|
|
|
|
|
}
|
2026-07-25 17:04:14 +08:00
|
|
|
|
split_specs: list[dict[str, Any]] = []
|
2026-07-24 20:43:47 +08:00
|
|
|
|
for split_name in split_order:
|
|
|
|
|
|
split_records = [
|
|
|
|
|
|
(source_row, record)
|
|
|
|
|
|
for source_row, record in zip(rows, records, strict=True)
|
|
|
|
|
|
if record["split"] == split_name
|
|
|
|
|
|
]
|
|
|
|
|
|
file_id = new_id("dfile")
|
|
|
|
|
|
version_id = new_id("dfv")
|
|
|
|
|
|
content = "".join(
|
|
|
|
|
|
json_dumps(record) + "\n" for _, record in split_records
|
|
|
|
|
|
)
|
|
|
|
|
|
raw = content.encode("utf-8")
|
2026-07-25 17:04:14 +08:00
|
|
|
|
split_specs.append(
|
2026-07-24 20:43:47 +08:00
|
|
|
|
{
|
|
|
|
|
|
"split": split_name,
|
|
|
|
|
|
"records": split_records,
|
|
|
|
|
|
"file_id": file_id,
|
|
|
|
|
|
"version_id": version_id,
|
|
|
|
|
|
"content": content,
|
|
|
|
|
|
"raw": raw,
|
|
|
|
|
|
"checksum": hashlib.sha256(raw).hexdigest(),
|
|
|
|
|
|
"storage_object_id": (
|
|
|
|
|
|
f"db://data-process/{task_id}/{file_id}/v1"
|
|
|
|
|
|
),
|
|
|
|
|
|
}
|
|
|
|
|
|
)
|
2026-07-23 15:10:13 +08:00
|
|
|
|
source_result_ids = [row["id"] for row in rows]
|
2026-07-25 17:04:14 +08:00
|
|
|
|
common_metadata = {
|
2026-07-23 15:10:13 +08:00
|
|
|
|
"source": "data_process",
|
|
|
|
|
|
"storage_backend": "database",
|
|
|
|
|
|
"source_task_id": task_id,
|
|
|
|
|
|
"source_file_ids": [item["id"] for item in self._source_ids(conn, task_id)],
|
|
|
|
|
|
"source_result_ids": source_result_ids,
|
|
|
|
|
|
"format": payload.get("format") or "alpaca_jsonl",
|
2026-07-24 20:43:47 +08:00
|
|
|
|
"split": requested_split,
|
2026-07-23 15:10:13 +08:00
|
|
|
|
}
|
2026-07-25 17:04:14 +08:00
|
|
|
|
|
|
|
|
|
|
existing_datasets = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT * FROM datasets
|
|
|
|
|
|
WHERE source_task_id=%s AND source='task'
|
|
|
|
|
|
ORDER BY created_at, id
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id,),
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
existing_by_split: dict[str, dict[str, Any]] = {}
|
|
|
|
|
|
primary_existing = None
|
|
|
|
|
|
for existing in existing_datasets:
|
|
|
|
|
|
existing_metadata = _json_value(existing.get("metadata"), {})
|
|
|
|
|
|
existing_split = str(existing_metadata.get("dataset_split") or "")
|
|
|
|
|
|
if existing_split in split_order:
|
|
|
|
|
|
existing_by_split[existing_split] = existing
|
|
|
|
|
|
if str(existing["id"]) == str(task.get("output_dataset_id") or ""):
|
|
|
|
|
|
primary_existing = existing
|
|
|
|
|
|
if primary_existing and "train" not in existing_by_split:
|
|
|
|
|
|
# 兼容旧版“一个数据集包含三个文件”的发布物,原数据集复用为训练集。
|
|
|
|
|
|
existing_by_split["train"] = primary_existing
|
|
|
|
|
|
|
|
|
|
|
|
existing_group_metadata = _json_value(
|
|
|
|
|
|
(primary_existing or {}).get("metadata"), {}
|
|
|
|
|
|
)
|
|
|
|
|
|
base_dataset_name = str(
|
|
|
|
|
|
existing_group_metadata.get("base_dataset_name")
|
|
|
|
|
|
or payload["dataset_name"]
|
|
|
|
|
|
).strip()
|
|
|
|
|
|
for suffix in ("-训练集", "-验证集", "-测试集"):
|
|
|
|
|
|
if base_dataset_name.endswith(suffix):
|
|
|
|
|
|
base_dataset_name = base_dataset_name[: -len(suffix)].rstrip()
|
|
|
|
|
|
break
|
|
|
|
|
|
split_group_id = str(
|
|
|
|
|
|
existing_group_metadata.get("split_group_id")
|
|
|
|
|
|
or f"dsg_{hashlib.sha256(task_id.encode()).hexdigest()[:20]}"
|
|
|
|
|
|
)
|
|
|
|
|
|
dataset_ids = {
|
|
|
|
|
|
spec["split"]: str(existing_by_split[spec["split"]]["id"])
|
|
|
|
|
|
if spec["split"] in existing_by_split
|
|
|
|
|
|
else new_id("dataset")
|
|
|
|
|
|
for spec in split_specs
|
|
|
|
|
|
}
|
|
|
|
|
|
created_any = any(
|
|
|
|
|
|
spec["split"] not in existing_by_split for spec in split_specs
|
|
|
|
|
|
)
|
|
|
|
|
|
split_labels = {
|
|
|
|
|
|
"train": "训练集",
|
|
|
|
|
|
"validation": "验证集",
|
|
|
|
|
|
"test": "测试集",
|
|
|
|
|
|
}
|
|
|
|
|
|
dataset_types = {"train": "train", "validation": "val", "test": "test"}
|
|
|
|
|
|
published_datasets: list[dict[str, Any]] = []
|
2026-07-23 15:10:13 +08:00
|
|
|
|
try:
|
2026-07-25 17:04:14 +08:00
|
|
|
|
for spec in split_specs:
|
|
|
|
|
|
split_name = str(spec["split"])
|
|
|
|
|
|
dataset_id = dataset_ids[split_name]
|
|
|
|
|
|
existing_dataset = existing_by_split.get(split_name)
|
|
|
|
|
|
dataset_metadata = {
|
|
|
|
|
|
**common_metadata,
|
|
|
|
|
|
"base_dataset_name": base_dataset_name,
|
|
|
|
|
|
"dataset_split": split_name,
|
|
|
|
|
|
"split_group_id": split_group_id,
|
|
|
|
|
|
"split_dataset_ids": dataset_ids,
|
|
|
|
|
|
"split_counts": {
|
|
|
|
|
|
name: split_counts[name] if name == split_name else 0
|
|
|
|
|
|
for name in split_order
|
|
|
|
|
|
},
|
|
|
|
|
|
}
|
|
|
|
|
|
dataset_name = f"{base_dataset_name}-{split_labels[split_name]}"
|
|
|
|
|
|
if existing_dataset:
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"DELETE FROM dataset_records WHERE dataset_id=%s", (dataset_id,)
|
|
|
|
|
|
)
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""DELETE FROM dataset_file_versions
|
|
|
|
|
|
WHERE dataset_file_id IN
|
|
|
|
|
|
(SELECT id FROM dataset_files WHERE dataset_id=%s)""",
|
|
|
|
|
|
(dataset_id,),
|
|
|
|
|
|
)
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"DELETE FROM dataset_files WHERE dataset_id=%s", (dataset_id,)
|
|
|
|
|
|
)
|
|
|
|
|
|
dataset = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE datasets
|
|
|
|
|
|
SET name=%s, type=%s, storage_type=%s, size=%s, size_bytes=%s,
|
|
|
|
|
|
count=%s, record_count=%s, description=%s, metadata=%s,
|
|
|
|
|
|
updated_at=%s
|
|
|
|
|
|
WHERE id=%s RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
dataset_name,
|
|
|
|
|
|
dataset_types[split_name],
|
|
|
|
|
|
payload.get("storage_type") or "local",
|
|
|
|
|
|
f"{len(spec['raw'])} B",
|
|
|
|
|
|
len(spec["raw"]),
|
|
|
|
|
|
len(spec["records"]),
|
|
|
|
|
|
len(spec["records"]),
|
|
|
|
|
|
payload.get("description") or task.get("description") or "",
|
|
|
|
|
|
json_dumps(dataset_metadata),
|
|
|
|
|
|
now,
|
|
|
|
|
|
dataset_id,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
|
|
|
|
|
else:
|
|
|
|
|
|
dataset = conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO datasets
|
|
|
|
|
|
(id, name, type, storage_type, source, task_id, source_task_id,
|
|
|
|
|
|
size, size_bytes, count, record_count, description, metadata,
|
|
|
|
|
|
tenant_id, project_id, owner_id, created_by, create_time,
|
|
|
|
|
|
created_at, updated_at)
|
2026-07-25 18:19:37 +08:00
|
|
|
|
VALUES (
|
|
|
|
|
|
%s, %s, %s, %s, 'task', %s, %s,
|
|
|
|
|
|
%s, %s, %s, %s, %s, %s,
|
|
|
|
|
|
%s, %s, %s, %s, %s, %s, %s
|
|
|
|
|
|
)
|
2026-07-25 17:04:14 +08:00
|
|
|
|
RETURNING *
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
dataset_id,
|
|
|
|
|
|
dataset_name,
|
|
|
|
|
|
dataset_types[split_name],
|
|
|
|
|
|
payload.get("storage_type") or "local",
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
f"{len(spec['raw'])} B",
|
|
|
|
|
|
len(spec["raw"]),
|
|
|
|
|
|
len(spec["records"]),
|
|
|
|
|
|
len(spec["records"]),
|
|
|
|
|
|
payload.get("description") or task.get("description") or "",
|
|
|
|
|
|
json_dumps(dataset_metadata),
|
|
|
|
|
|
task.get("tenant_id"),
|
|
|
|
|
|
task.get("project_id"),
|
|
|
|
|
|
task.get("owner_id"),
|
|
|
|
|
|
payload.get("created_by") or task.get("created_by"),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
).fetchone()
|
2026-07-24 20:43:47 +08:00
|
|
|
|
|
2026-07-25 17:04:14 +08:00
|
|
|
|
file_metadata = {**dataset_metadata, "file_split": split_name}
|
2026-07-24 20:43:47 +08:00
|
|
|
|
version = {
|
|
|
|
|
|
"id": spec["version_id"],
|
|
|
|
|
|
"version_no": 1,
|
|
|
|
|
|
"version": 1,
|
2026-07-25 17:04:14 +08:00
|
|
|
|
"description": f"data process {split_name} publish",
|
2026-07-24 20:43:47 +08:00
|
|
|
|
"checksum_sha256": spec["checksum"],
|
|
|
|
|
|
"size_bytes": len(spec["raw"]),
|
|
|
|
|
|
"record_count": len(spec["records"]),
|
|
|
|
|
|
"created_at": now,
|
|
|
|
|
|
"create_time": now,
|
|
|
|
|
|
"source_task_id": task_id,
|
|
|
|
|
|
"storage_object_id": spec["storage_object_id"],
|
|
|
|
|
|
}
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO dataset_files
|
|
|
|
|
|
(id, dataset_id, name, storage_object_id, size, content,
|
|
|
|
|
|
active_version_id, versions, create_time, current_version_id,
|
|
|
|
|
|
size_bytes, record_count, file_format, checksum_sha256, version_no,
|
|
|
|
|
|
source_task_id, tenant_id, project_id, created_by, metadata,
|
|
|
|
|
|
created_at, updated_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
|
|
|
|
|
%s, %s, 1, %s, %s, %s, %s, %s, %s, %s)
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
spec["file_id"],
|
|
|
|
|
|
dataset_id,
|
2026-07-25 17:04:14 +08:00
|
|
|
|
f"{base_dataset_name}.{split_name}.jsonl",
|
2026-07-24 20:43:47 +08:00
|
|
|
|
spec["storage_object_id"],
|
|
|
|
|
|
f"{len(spec['raw'])} B",
|
|
|
|
|
|
spec["content"],
|
|
|
|
|
|
spec["version_id"],
|
|
|
|
|
|
json_dumps([version]),
|
|
|
|
|
|
now,
|
|
|
|
|
|
spec["version_id"],
|
|
|
|
|
|
len(spec["raw"]),
|
|
|
|
|
|
len(spec["records"]),
|
|
|
|
|
|
"jsonl",
|
|
|
|
|
|
spec["checksum"],
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
task.get("tenant_id"),
|
|
|
|
|
|
task.get("project_id"),
|
|
|
|
|
|
payload.get("created_by") or task.get("created_by"),
|
|
|
|
|
|
json_dumps(file_metadata),
|
|
|
|
|
|
now,
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
)
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO dataset_file_versions
|
|
|
|
|
|
(id, dataset_file_id, version_no, storage_object_id, content_preview,
|
|
|
|
|
|
description, size_bytes, record_count, checksum_sha256,
|
|
|
|
|
|
source_task_id, metadata, created_by, created_at)
|
|
|
|
|
|
VALUES (%s, %s, 1, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
spec["version_id"],
|
|
|
|
|
|
spec["file_id"],
|
|
|
|
|
|
spec["storage_object_id"],
|
|
|
|
|
|
spec["content"][:2000],
|
2026-07-25 17:04:14 +08:00
|
|
|
|
f"data process {split_name} publish",
|
2026-07-24 20:43:47 +08:00
|
|
|
|
len(spec["raw"]),
|
|
|
|
|
|
len(spec["records"]),
|
|
|
|
|
|
spec["checksum"],
|
2026-07-23 15:10:13 +08:00
|
|
|
|
task_id,
|
2026-07-24 20:43:47 +08:00
|
|
|
|
json_dumps(file_metadata),
|
|
|
|
|
|
payload.get("created_by") or task.get("created_by"),
|
2026-07-23 15:10:13 +08:00
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
)
|
2026-07-24 20:43:47 +08:00
|
|
|
|
for line_number, (source_row, record) in enumerate(
|
|
|
|
|
|
spec["records"], start=1
|
|
|
|
|
|
):
|
|
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
INSERT INTO dataset_records
|
|
|
|
|
|
(id, dataset_id, dataset_file_id, version_id, line_no, split,
|
|
|
|
|
|
instruction, input, output, raw, status, source_task_id,
|
|
|
|
|
|
source_result_id, preview_item_id, created_at)
|
|
|
|
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
|
|
|
|
|
%s, %s, %s, %s)
|
|
|
|
|
|
""",
|
|
|
|
|
|
(
|
|
|
|
|
|
new_id("drec"),
|
|
|
|
|
|
dataset_id,
|
|
|
|
|
|
spec["file_id"],
|
|
|
|
|
|
spec["version_id"],
|
|
|
|
|
|
line_number,
|
|
|
|
|
|
record["split"],
|
|
|
|
|
|
record["instruction"],
|
|
|
|
|
|
record["input"],
|
|
|
|
|
|
record["output"],
|
|
|
|
|
|
json_dumps(
|
|
|
|
|
|
{
|
|
|
|
|
|
**record,
|
|
|
|
|
|
"source_task_id": task_id,
|
|
|
|
|
|
"source_result_id": source_row["id"],
|
|
|
|
|
|
"preview_item_id": source_row.get("preview_item_id"),
|
|
|
|
|
|
}
|
|
|
|
|
|
),
|
|
|
|
|
|
source_row["status"],
|
|
|
|
|
|
task_id,
|
|
|
|
|
|
source_row["id"],
|
|
|
|
|
|
source_row.get("preview_item_id"),
|
|
|
|
|
|
now,
|
|
|
|
|
|
),
|
|
|
|
|
|
)
|
2026-07-25 17:04:14 +08:00
|
|
|
|
published_datasets.append(_decode_row(dataset) or {})
|
|
|
|
|
|
|
2026-07-23 15:10:13 +08:00
|
|
|
|
except psycopg.errors.UniqueViolation as exc:
|
|
|
|
|
|
raise ConflictError("dataset name already exists") from exc
|
2026-07-25 17:04:14 +08:00
|
|
|
|
train_dataset_id = dataset_ids.get("train")
|
|
|
|
|
|
if not train_dataset_id:
|
|
|
|
|
|
raise InvalidStateError("published split does not contain training data")
|
2026-07-23 15:10:13 +08:00
|
|
|
|
conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
UPDATE data_process_tasks
|
|
|
|
|
|
SET output_dataset_id=%s, updated_at=%s, updated_by=%s
|
|
|
|
|
|
WHERE id=%s
|
|
|
|
|
|
""",
|
2026-07-25 17:04:14 +08:00
|
|
|
|
(train_dataset_id, now, payload.get("created_by"), task_id),
|
2026-07-23 15:10:13 +08:00
|
|
|
|
)
|
2026-07-25 17:04:14 +08:00
|
|
|
|
train_dataset = next(
|
|
|
|
|
|
item
|
|
|
|
|
|
for item in published_datasets
|
|
|
|
|
|
if _json_value(item.get("metadata"), {}).get("dataset_split") == "train"
|
|
|
|
|
|
)
|
2026-07-24 20:43:47 +08:00
|
|
|
|
return {
|
2026-07-25 17:04:14 +08:00
|
|
|
|
"dataset": train_dataset,
|
|
|
|
|
|
"datasets": published_datasets,
|
|
|
|
|
|
"output_datasets": published_datasets,
|
|
|
|
|
|
"created": created_any,
|
2026-07-24 20:43:47 +08:00
|
|
|
|
"split_counts": split_counts,
|
|
|
|
|
|
}
|
2026-07-23 15:10:13 +08:00
|
|
|
|
|
|
|
|
|
|
@staticmethod
|
|
|
|
|
|
def _source_ids(
|
|
|
|
|
|
conn: psycopg.Connection[dict[str, Any]], task_id: str
|
|
|
|
|
|
) -> list[dict[str, Any]]:
|
|
|
|
|
|
return conn.execute(
|
|
|
|
|
|
"""
|
|
|
|
|
|
SELECT id FROM data_process_source_files
|
|
|
|
|
|
WHERE task_id=%s AND deleted_at IS NULL ORDER BY created_at, id
|
|
|
|
|
|
""",
|
|
|
|
|
|
(task_id,),
|
|
|
|
|
|
).fetchall()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@lru_cache
|
|
|
|
|
|
def get_data_process_store() -> DataProcessStore:
|
|
|
|
|
|
return DataProcessStore()
|