365 lines
12 KiB
Python
365 lines
12 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""Seed inspect endpoints for data recipe."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
import binascii
|
|
from itertools import islice
|
|
from pathlib import Path
|
|
from typing import Any
|
|
from uuid import uuid4
|
|
|
|
from fastapi import APIRouter, HTTPException
|
|
from data_designer_unstructured_seed.chunking import (
|
|
build_unstructured_preview_rows,
|
|
resolve_chunking,
|
|
)
|
|
from core.data_recipe.jsonable import to_preview_jsonable
|
|
from utils.paths import ensure_dir, seed_uploads_root
|
|
|
|
from models.data_recipe import (
|
|
SeedInspectRequest,
|
|
SeedInspectResponse,
|
|
SeedInspectUploadRequest,
|
|
)
|
|
|
|
router = APIRouter()
|
|
|
|
DATA_EXTS = (".parquet", ".jsonl", ".json", ".csv")
|
|
DEFAULT_SPLIT = "train"
|
|
LOCAL_UPLOAD_EXTS = {".csv", ".json", ".jsonl"}
|
|
UNSTRUCTURED_UPLOAD_EXTS = {".txt", ".md"}
|
|
SEED_UPLOAD_DIR = seed_uploads_root()
|
|
|
|
|
|
def _serialize_preview_value(value: Any) -> Any:
|
|
return to_preview_jsonable(value)
|
|
|
|
|
|
def _serialize_preview_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
return [
|
|
{str(key): _serialize_preview_value(value) for key, value in row.items()}
|
|
for row in rows
|
|
]
|
|
|
|
|
|
def _normalize_optional_text(value: str | None) -> str | None:
|
|
if value is None:
|
|
return None
|
|
trimmed = value.strip()
|
|
return trimmed if trimmed else None
|
|
|
|
|
|
def _list_hf_data_files(*, dataset_name: str, token: str | None) -> list[str]:
|
|
try:
|
|
from huggingface_hub import HfApi
|
|
from huggingface_hub.utils import HfHubHTTPError
|
|
except ImportError:
|
|
return []
|
|
try:
|
|
api = HfApi()
|
|
repo_files = api.list_repo_files(dataset_name, repo_type = "dataset", token = token)
|
|
return [file for file in repo_files if file.lower().endswith(DATA_EXTS)]
|
|
except (HfHubHTTPError, OSError, ValueError):
|
|
return []
|
|
|
|
|
|
def _select_best_file(data_files: list[str], split: str = DEFAULT_SPLIT) -> str | None:
|
|
if not data_files:
|
|
return None
|
|
split_lower = split.lower()
|
|
|
|
def score(path: str) -> tuple[int, int]:
|
|
name = path.lower()
|
|
if f"/{split_lower}/" in name:
|
|
return (0, len(path))
|
|
if (
|
|
f"_{split_lower}." in name
|
|
or f"-{split_lower}." in name
|
|
or f"/{split_lower}." in name
|
|
or f"/{split_lower}_" in name
|
|
or f"/{split_lower}-" in name
|
|
):
|
|
return (1, len(path))
|
|
return (2, len(path))
|
|
|
|
return sorted(data_files, key = score)[0]
|
|
|
|
|
|
def _resolve_seed_hf_path(
|
|
dataset_name: str, data_files: list[str], split: str = DEFAULT_SPLIT
|
|
) -> str | None:
|
|
selected = _select_best_file(data_files, split)
|
|
if not selected:
|
|
return None
|
|
|
|
ext = Path(selected).suffix.lower()
|
|
if ext not in DATA_EXTS:
|
|
return f"datasets/{dataset_name}/{selected}"
|
|
|
|
parent = Path(selected).parent.as_posix()
|
|
if not parent or parent == ".":
|
|
return f"datasets/{dataset_name}/**/*{ext}"
|
|
return f"datasets/{dataset_name}/{parent}/**/*{ext}"
|
|
|
|
|
|
def _build_stream_load_kwargs(
|
|
*,
|
|
dataset_name: str,
|
|
split: str,
|
|
subset: str | None,
|
|
token: str | None,
|
|
data_file: str | None = None,
|
|
) -> dict[str, Any]:
|
|
kwargs: dict[str, Any] = {
|
|
"path": dataset_name,
|
|
"split": split,
|
|
"streaming": True,
|
|
}
|
|
if data_file:
|
|
kwargs["data_files"] = [data_file]
|
|
if subset:
|
|
kwargs["name"] = subset
|
|
if token:
|
|
kwargs["token"] = token
|
|
return kwargs
|
|
|
|
|
|
def _load_preview_rows(
|
|
*,
|
|
load_dataset_fn,
|
|
load_kwargs: dict[str, Any],
|
|
preview_size: int,
|
|
) -> list[dict[str, Any]]:
|
|
streamed_ds = load_dataset_fn(**load_kwargs)
|
|
return [row for row in islice(streamed_ds, preview_size)]
|
|
|
|
|
|
def _extract_columns(rows: list[dict[str, Any]]) -> list[str]:
|
|
columns_seen: dict[str, None] = {}
|
|
for row in rows:
|
|
for key in row.keys():
|
|
columns_seen[str(key)] = None
|
|
return list(columns_seen.keys())
|
|
|
|
|
|
def _sanitize_filename(filename: str) -> str:
|
|
name = Path(filename).name.strip().replace("\x00", "")
|
|
if not name:
|
|
return "seed_upload"
|
|
return name
|
|
|
|
|
|
def _decode_base64_payload(content_base64: str) -> bytes:
|
|
raw = content_base64.strip()
|
|
if "," in raw and raw.lower().startswith("data:"):
|
|
raw = raw.split(",", 1)[1]
|
|
try:
|
|
return base64.b64decode(raw, validate = True)
|
|
except binascii.Error as exc:
|
|
raise HTTPException(status_code = 400, detail = "invalid base64 payload") from exc
|
|
|
|
|
|
def _read_preview_rows_from_local_file(
|
|
path: Path, preview_size: int
|
|
) -> list[dict[str, Any]]:
|
|
try:
|
|
import pandas as pd
|
|
except ImportError as exc:
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"seed inspect dependencies unavailable: {exc}"
|
|
) from exc
|
|
|
|
ext = path.suffix.lower()
|
|
try:
|
|
if ext == ".csv":
|
|
df = pd.read_csv(path, nrows = preview_size)
|
|
elif ext == ".jsonl":
|
|
df = pd.read_json(path, lines = True).head(preview_size)
|
|
elif ext == ".json":
|
|
try:
|
|
df = pd.read_json(path).head(preview_size)
|
|
except ValueError:
|
|
df = pd.read_json(path, lines = True).head(preview_size)
|
|
else:
|
|
raise HTTPException(status_code = 422, detail = f"unsupported file type: {ext}")
|
|
except HTTPException:
|
|
raise
|
|
except (ValueError, OSError) as exc:
|
|
raise HTTPException(
|
|
status_code = 422, detail = f"seed inspect failed: {exc}"
|
|
) from exc
|
|
|
|
rows = df.to_dict(orient = "records")
|
|
return _serialize_preview_rows(rows)
|
|
|
|
|
|
def _read_preview_rows_from_unstructured_file(
|
|
*,
|
|
path: Path,
|
|
preview_size: int,
|
|
chunk_size: int | None,
|
|
chunk_overlap: int | None,
|
|
) -> list[dict[str, Any]]:
|
|
size, overlap = resolve_chunking(chunk_size, chunk_overlap)
|
|
try:
|
|
rows = build_unstructured_preview_rows(
|
|
source_path = path,
|
|
preview_size = preview_size,
|
|
chunk_size = size,
|
|
chunk_overlap = overlap,
|
|
)
|
|
except (FileNotFoundError, RuntimeError, ValueError, OSError) as exc:
|
|
raise HTTPException(
|
|
status_code = 422, detail = f"seed inspect failed: {exc}"
|
|
) from exc
|
|
return _serialize_preview_rows(rows)
|
|
|
|
|
|
@router.post("/seed/inspect", response_model = SeedInspectResponse)
|
|
def inspect_seed_dataset(payload: SeedInspectRequest) -> SeedInspectResponse:
|
|
dataset_name = payload.dataset_name.strip()
|
|
if not dataset_name or dataset_name.count("/") < 1:
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = "dataset_name must be a Hugging Face repo id like org/repo",
|
|
)
|
|
|
|
try:
|
|
from datasets import load_dataset
|
|
except ImportError as exc:
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"seed inspect dependencies unavailable: {exc}"
|
|
) from exc
|
|
|
|
split = _normalize_optional_text(payload.split) or DEFAULT_SPLIT
|
|
subset = _normalize_optional_text(payload.subset)
|
|
token = _normalize_optional_text(payload.hf_token)
|
|
preview_size = int(payload.preview_size)
|
|
|
|
preview_rows: list[dict[str, Any]] = []
|
|
data_files = _list_hf_data_files(dataset_name = dataset_name, token = token)
|
|
|
|
selected_file = _select_best_file(data_files, split)
|
|
if selected_file:
|
|
try:
|
|
single_file_kwargs = _build_stream_load_kwargs(
|
|
dataset_name = dataset_name,
|
|
split = split,
|
|
subset = subset,
|
|
token = token,
|
|
data_file = selected_file,
|
|
)
|
|
preview_rows = _load_preview_rows(
|
|
load_dataset_fn = load_dataset,
|
|
load_kwargs = single_file_kwargs,
|
|
preview_size = preview_size,
|
|
)
|
|
except (ValueError, OSError, RuntimeError):
|
|
preview_rows = []
|
|
|
|
if not preview_rows:
|
|
try:
|
|
split_kwargs = _build_stream_load_kwargs(
|
|
dataset_name = dataset_name,
|
|
split = split,
|
|
subset = subset,
|
|
token = token,
|
|
)
|
|
preview_rows = _load_preview_rows(
|
|
load_dataset_fn = load_dataset,
|
|
load_kwargs = split_kwargs,
|
|
preview_size = preview_size,
|
|
)
|
|
except (ValueError, OSError, RuntimeError) as exc:
|
|
raise HTTPException(
|
|
status_code = 422, detail = f"seed inspect failed: {exc}"
|
|
) from exc
|
|
|
|
if not preview_rows:
|
|
raise HTTPException(
|
|
status_code = 422, detail = "dataset appears empty or unreadable"
|
|
)
|
|
preview_rows = _serialize_preview_rows(preview_rows)
|
|
columns = _extract_columns(preview_rows)
|
|
|
|
if not data_files:
|
|
resolved_path = f"datasets/{dataset_name}/**/*.parquet"
|
|
else:
|
|
resolved_path = _resolve_seed_hf_path(dataset_name, data_files, split)
|
|
if not resolved_path:
|
|
raise HTTPException(
|
|
status_code = 422, detail = "unable to resolve seed dataset path"
|
|
)
|
|
|
|
return SeedInspectResponse(
|
|
dataset_name = dataset_name,
|
|
resolved_path = resolved_path,
|
|
columns = columns,
|
|
preview_rows = preview_rows,
|
|
split = split,
|
|
subset = subset,
|
|
)
|
|
|
|
|
|
@router.post("/seed/inspect-upload", response_model = SeedInspectResponse)
|
|
def inspect_seed_upload(payload: SeedInspectUploadRequest) -> SeedInspectResponse:
|
|
seed_source_type = _normalize_optional_text(payload.seed_source_type) or "local"
|
|
filename = _sanitize_filename(payload.filename)
|
|
ext = Path(filename).suffix.lower()
|
|
if seed_source_type == "unstructured":
|
|
if ext not in UNSTRUCTURED_UPLOAD_EXTS:
|
|
allowed = ", ".join(sorted(UNSTRUCTURED_UPLOAD_EXTS))
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = f"unsupported file type: {ext}. allowed: {allowed}",
|
|
)
|
|
else:
|
|
if ext not in LOCAL_UPLOAD_EXTS:
|
|
allowed = ", ".join(sorted(LOCAL_UPLOAD_EXTS))
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = f"unsupported file type: {ext}. allowed: {allowed}",
|
|
)
|
|
|
|
file_bytes = _decode_base64_payload(payload.content_base64)
|
|
if not file_bytes:
|
|
raise HTTPException(status_code = 400, detail = "empty upload payload")
|
|
max_size_bytes = 50 * 1024 * 1024
|
|
if len(file_bytes) > max_size_bytes:
|
|
raise HTTPException(status_code = 413, detail = "file too large (max 50MB)")
|
|
|
|
ensure_dir(SEED_UPLOAD_DIR)
|
|
stored_name = f"{uuid4().hex}_{filename}"
|
|
stored_path = SEED_UPLOAD_DIR / stored_name
|
|
stored_path.write_bytes(file_bytes)
|
|
|
|
if seed_source_type == "unstructured":
|
|
preview_rows = _read_preview_rows_from_unstructured_file(
|
|
path = stored_path,
|
|
preview_size = int(payload.preview_size),
|
|
chunk_size = payload.unstructured_chunk_size,
|
|
chunk_overlap = payload.unstructured_chunk_overlap,
|
|
)
|
|
else:
|
|
preview_rows = _read_preview_rows_from_local_file(
|
|
stored_path,
|
|
int(payload.preview_size),
|
|
)
|
|
if not preview_rows:
|
|
raise HTTPException(
|
|
status_code = 422, detail = "dataset appears empty or unreadable"
|
|
)
|
|
columns = _extract_columns(preview_rows)
|
|
|
|
return SeedInspectResponse(
|
|
dataset_name = filename,
|
|
resolved_path = str(stored_path),
|
|
columns = columns,
|
|
preview_rows = preview_rows,
|
|
split = None,
|
|
subset = None,
|
|
)
|