This commit is contained in:
wangjiming
2026-08-18 14:49:34 +08:00
10 changed files with 313 additions and 36 deletions

View File

@@ -5,12 +5,15 @@ import asyncio
import hashlib
import uuid
import time
from io import BytesIO
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any
from urllib.parse import quote
from zipfile import ZIP_DEFLATED, ZipFile
from fastapi import APIRouter, BackgroundTasks, Body, Depends, File, HTTPException, Query, Request, UploadFile
from fastapi.responses import PlainTextResponse, StreamingResponse
from fastapi.responses import PlainTextResponse, Response, StreamingResponse
import httpx
@@ -365,9 +368,16 @@ async def _fine_tune_preflight_payload(
store: Any,
payload: dict[str, Any],
validate: bool = True,
sync_resources: bool = True,
) -> dict[str, Any]:
node, job_payload = store.prepare_compute_job_payload_from_payload(payload)
return await _fine_tune_preflight_with_job_payload(node, job_payload, validate=validate, sync_resources=False, store=store)
return await _fine_tune_preflight_with_job_payload(
node,
job_payload,
validate=validate,
sync_resources=sync_resources,
store=store,
)
async def _fine_tune_preflight_with_job_payload(
@@ -1147,13 +1157,42 @@ async def _sync_training_dataset_to_compute_node(
) -> list[dict[str, Any]]:
if get_settings().minio_enabled:
files = store.training_dataset_files(dataset_id)
objects = store.storage_objects_for_resource("dataset", dataset_id)
object_by_name = {Path(str(item.get("file_name") or item.get("object_key") or "")).name: item for item in objects}
object_by_resource_name: dict[tuple[str, str], dict[str, Any]] = {}
resource_ids = {str(dataset_id)} | {
str(item.get("dataset_id"))
for item in files
if item.get("dataset_id")
}
for resource_id in resource_ids:
for obj in store.storage_objects_for_resource("dataset", resource_id):
file_name = Path(str(obj.get("file_name") or obj.get("object_key") or "")).name
object_by_resource_name[(resource_id, file_name)] = obj
results: list[dict[str, Any]] = []
client = ComputeNodeClient(node["api_base_url"])
for item in files:
target_name = Path(str(item.get("name") or f"{item['id']}.jsonl")).name
obj = object_by_name.get(target_name)
item_dataset_id = str(item.get("dataset_id") or dataset_id)
obj = object_by_resource_name.get((item_dataset_id, target_name))
if not obj and item.get("content"):
# 兼容 MinIO 接入前已经发布的数据处理数据集:
# 预检时用数据库正文补建对象,避免要求用户重新处理数据集。
raw = str(item.get("content") or "").encode("utf-8")
version_id = str(item.get("active_version_id") or item["id"])
object_key = f"datasets/{item_dataset_id}/versions/{version_id}/{target_name}"
uploaded = get_object_storage().put_bytes(object_key, raw, "application/jsonl")
obj = store.create_storage_object({
"resource_type": "dataset",
"resource_id": item_dataset_id,
"version_id": version_id,
"bucket": uploaded["bucket"],
"object_key": object_key,
"file_name": target_name,
"content_type": "application/jsonl",
"byte_size": len(raw),
"checksum_sha256": hashlib.sha256(raw).hexdigest(),
"status": "available",
})
store.link_dataset_file_storage_object(str(item["id"]), obj["id"])
if not obj:
raise RuntimeError(f"dataset file is not available in MinIO: {target_name}")
url = get_object_storage().presigned_get(obj["object_key"])
@@ -1165,6 +1204,12 @@ async def _sync_training_dataset_to_compute_node(
"byte_size": obj.get("byte_size") or 0,
"relative_path": f"datasets/{dataset_id}/{target_name}",
})
store.upsert_resource_replica(
node["id"],
"dataset",
dataset_id,
str(result.get("local_path") or ""),
)
results.append({**result, "file_id": item.get("id"), "name": target_name, "node_id": node["id"]})
return results
if not dataset_id:
@@ -1233,7 +1278,7 @@ async def upload_dataset_files(
if get_settings().minio_enabled:
object_key = f"datasets/{dataset_id}/versions/{created_file.get('active_version_id') or created_file['id']}/{Path(created_file['name']).name}"
uploaded = get_object_storage().put_bytes(object_key, raw, file.content_type or "application/octet-stream")
get_platform_store().create_storage_object({
storage_object = get_platform_store().create_storage_object({
"resource_type": "dataset", "resource_id": dataset_id,
"version_id": created_file.get("active_version_id") or created_file["id"],
"bucket": uploaded["bucket"], "object_key": object_key,
@@ -1241,6 +1286,7 @@ async def upload_dataset_files(
"byte_size": len(raw), "checksum_sha256": hashlib.sha256(raw).hexdigest(),
"status": "available",
})
store.link_dataset_file_storage_object(created_file["id"], storage_object["id"])
if sync_to_compute:
for file_id, file_name, raw in pending_sync:
compute_sync.extend(
@@ -1256,10 +1302,60 @@ async def upload_dataset_files(
@router.get("/dataset-manage/download/{dataset_id}")
async def download_dataset(dataset_id: str) -> PlainTextResponse:
dataset = get_platform_store().dataset(dataset_id)
content = "\n".join([f"{file['name']}" for file in dataset.get("files", [])])
return PlainTextResponse(content, media_type="text/plain")
async def download_dataset(dataset_id: str, current_user: dict = Depends(get_current_user)) -> Response:
store = get_platform_store()
try:
dataset = store.dataset(dataset_id)
except KeyError:
raise fail(404, "dataset not found")
if not has_resource_access("dataset", dataset_id, current_user, "read"):
raise fail(403, "no permission to access this dataset")
files = []
for item in dataset.get("files", []):
if item.get("deleted_at"):
continue
try:
full_file = store.dataset_file(str(item["id"]))
except KeyError:
continue
files.append({**item, "content": full_file.get("content") or ""})
if not files:
raise fail(404, "dataset has no downloadable files")
if len(files) == 1:
item = files[0]
filename = Path(str(item.get("name") or f"{dataset_id}.jsonl")).name
encoded_name = quote(filename, safe="")
return Response(
content=str(item.get("content") or "").encode("utf-8"),
media_type="application/octet-stream",
headers={
"Content-Disposition": f"attachment; filename=dataset-file; filename*=UTF-8''{encoded_name}",
},
)
archive = BytesIO()
with ZipFile(archive, "w", compression=ZIP_DEFLATED) as bundle:
used_names: set[str] = set()
for index, item in enumerate(files, start=1):
filename = Path(str(item.get("name") or f"file-{index}.jsonl")).name
unique_name = filename
if unique_name in used_names:
stem = Path(filename).stem
suffix = Path(filename).suffix
unique_name = f"{stem}-{index}{suffix}"
used_names.add(unique_name)
bundle.writestr(unique_name, str(item.get("content") or ""))
archive.seek(0)
archive_name = quote(f"{dataset.get('name') or dataset_id}.zip", safe="")
return StreamingResponse(
archive,
media_type="application/zip",
headers={
"Content-Disposition": f"attachment; filename=dataset.zip; filename*=UTF-8''{archive_name}",
},
)
@router.get("/dataset-manage/download/{dataset_id}/{file_id}")
@@ -1289,8 +1385,11 @@ async def dataset_list(current_user: dict = Depends(get_current_user)) -> dict[s
@op_log(module=OpModule.DATASET, action=OpAction.CREATE, target_type="dataset", target_name_param="name")
async def create_dataset(payload: dict[str, Any] = Body(...), current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
payload.setdefault("created_by", current_user.get("id"))
dataset = get_platform_store().create_dataset(payload)
return ok({"id": dataset["id"]})
try:
dataset = get_platform_store().create_dataset(payload)
return ok({"id": dataset["id"]})
except ValueError as exc:
raise fail(409, str(exc))
@router.get("/dataset-manage/{dataset_id}")
@@ -1423,7 +1522,7 @@ async def start_fine_tune(
@router.post("/fine-tune/preflight")
async def fine_tune_create_preflight(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
try:
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=True))
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=True, sync_resources=True))
except RuntimeError as exc:
return ok({"valid": False, "errors": [str(exc)], "warnings": [], "diagnostics": _training_diagnostics([str(exc)])})
except Exception as exc: # noqa: BLE001 - expose compute validation errors to training create page
@@ -1433,7 +1532,7 @@ async def fine_tune_create_preflight(payload: dict[str, Any] = Body(...)) -> dic
@router.post("/fine-tune/command-preview")
async def fine_tune_create_command_preview(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
try:
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=False))
return ok(await _fine_tune_preflight_payload(get_platform_store(), payload, validate=False, sync_resources=False))
except RuntimeError as exc:
return ok({"valid": False, "errors": [str(exc)], "warnings": [], "diagnostics": _training_diagnostics([str(exc)])})
except Exception as exc: # noqa: BLE001