Skip to content
Merged

Dev #64

Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions backend/openmlr/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,13 +109,20 @@ async def lifespan(app: FastAPI):
from .auth.router import router as auth_router
from .routes.agent import router as agent_router
from .routes.compute import router as compute_router
from .routes.datasets import router as datasets_router
from .routes.eval import router as eval_router
from .routes.experiments import router as experiments_router
from .routes.figures import router as figures_router
from .routes.health import router as health_router
from .routes.keys import router as keys_router
from .routes.mcp import router as mcp_router
from .routes.models import router as models_router
from .routes.projects import router as projects_router
from .routes.reproducibility import router as reproducibility_router
from .routes.research import router as research_router
from .routes.review import router as review_router
from .routes.settings import router as settings_router
from .routes.sweeps import router as sweeps_router
from .routes.terminal import router as terminal_router

app.include_router(auth_router)
Expand All @@ -126,8 +133,15 @@ async def lifespan(app: FastAPI):
app.include_router(compute_router)
app.include_router(mcp_router)
app.include_router(projects_router)
app.include_router(research_router)
app.include_router(review_router)
app.include_router(eval_router)
app.include_router(experiments_router)
app.include_router(datasets_router)
app.include_router(sweeps_router)
app.include_router(models_router)
app.include_router(figures_router)
app.include_router(reproducibility_router)
app.include_router(terminal_router)


Expand Down
201 changes: 201 additions & 0 deletions backend/openmlr/routes/datasets.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,201 @@
"""Dataset management, profiling, validation, and split REST API routes.

Provides endpoints for inspecting data files, analyzing column statistics, validating schemas,
and generating partitioned splits for ML training workflows.
"""

from __future__ import annotations

import logging
import os
import re
from pathlib import Path
from typing import Any

from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, Field

from ..db.models import User
from ..dependencies import get_current_user_optional
from ..services.dataset_profiler import DatasetProfiler

router = APIRouter(prefix="/api/datasets", tags=["datasets"])
logger = logging.getLogger(__name__)

DATASET_PATH_DESC = "Path to the dataset file"
SAFE_PATH_PATTERN = re.compile(r"^[a-zA-Z0-9_\-./ ]+$")


def _safe_dataset_path(path_str: str) -> Path:
"""Safely validate and resolve a dataset file path, mitigating path injection risks."""
clean = os.path.normpath(str(path_str).strip())
if not clean or not SAFE_PATH_PATTERN.match(clean) or ".." in clean or clean.startswith(("/etc", "/proc", "/sys")):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or disallowed dataset path parameter",
)
resolved = Path(clean).resolve()
if not resolved.exists() or not resolved.is_file():
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Dataset file not found at '{path_str}'",
)
return resolved


def _safe_output_dir(dir_str: str) -> Path:
"""Safely validate and resolve an output directory path."""
clean = os.path.normpath(str(dir_str).strip())
if not clean or not SAFE_PATH_PATTERN.match(clean) or ".." in clean or clean.startswith(("/etc", "/proc", "/sys")):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or disallowed output directory path parameter",
)
return Path(clean).resolve()


class ProfileDatasetRequest(BaseModel):
"""Payload for profiling a dataset file."""

path: str = Field(..., min_length=1, description=DATASET_PATH_DESC)
sample_size: int = Field(default=5000, ge=10, le=100000, description="Max rows to sample")


class InspectSamplesRequest(BaseModel):
"""Payload for sampling rows from a dataset."""

path: str = Field(..., min_length=1, description=DATASET_PATH_DESC)
n: int = Field(default=5, ge=1, le=100, description="Number of rows to return")
offset: int = Field(default=0, ge=0, description="Row offset")
strategy: str = Field(default="head", description="Sampling strategy: head, random, stratified")
label_column: str | None = Field(default=None, description="Column for stratification")


class ValidateDatasetRequest(BaseModel):
"""Payload for validating a dataset."""

path: str = Field(..., min_length=1, description=DATASET_PATH_DESC)
expected_columns: list[str] | None = Field(default=None, description="Required column names")
max_null_pct: float = Field(default=20.0, ge=0.0, le=100.0, description="Max allowed null percentage")
max_token_length: int | None = Field(default=None, ge=1, description="Max allowed text tokens")


class SplitDatasetRequest(BaseModel):
"""Payload for splitting a dataset into train/val/test partitions."""

path: str = Field(..., min_length=1, description="Path to the source dataset file")
output_dir: str = Field(..., min_length=1, description="Target directory for output partitions")
train_ratio: float = Field(default=0.8, gt=0.0, lt=1.0, description="Train ratio")
val_ratio: float = Field(default=0.1, ge=0.0, lt=1.0, description="Validation ratio")
test_ratio: float = Field(default=0.1, ge=0.0, lt=1.0, description="Test ratio")
stratify_column: str | None = Field(default=None, description="Column to stratify on")
seed: int = Field(default=42, description="Random seed")


@router.post("/profile")
async def profile_dataset(
req: ProfileDatasetRequest,
user: User | None = Depends(get_current_user_optional),
) -> dict[str, Any]:
"""Compute comprehensive statistical profile and diagnostics for a dataset file."""
path = _safe_dataset_path(req.path)

try:
profile = DatasetProfiler.profile(path, sample_size=req.sample_size)
return {
"success": True,
"profile": profile.to_dict(),
}
except Exception as e:
logger.exception("Failed profiling dataset '%s': %s", req.path, e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Error profiling dataset: {e}",
)


@router.post("/inspect")
async def inspect_samples(
req: InspectSamplesRequest,
user: User | None = Depends(get_current_user_optional),
) -> dict[str, Any]:
"""Retrieve sample records from a dataset."""
path = _safe_dataset_path(req.path)

try:
samples = DatasetProfiler.sample_records(
path,
n=req.n,
offset=req.offset,
strategy=req.strategy,
label_column=req.label_column,
)
return {
"success": True,
"total_sampled": len(samples),
"samples": samples,
}
except Exception as e:
logger.exception("Failed inspecting dataset '%s': %s", req.path, e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Error inspecting dataset: {e}",
)


@router.post("/validate")
async def validate_dataset(
req: ValidateDatasetRequest,
user: User | None = Depends(get_current_user_optional),
) -> dict[str, Any]:
"""Validate dataset structure, required schema, and constraints."""
path = _safe_dataset_path(req.path)

try:
result = DatasetProfiler.validate_dataset(
path,
expected_columns=req.expected_columns,
max_null_pct=req.max_null_pct,
max_token_length=req.max_token_length,
)
return {
"success": True,
"validation": result,
}
except Exception as e:
logger.exception("Failed validating dataset '%s': %s", req.path, e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Error validating dataset: {e}",
)


@router.post("/split")
async def split_dataset(
req: SplitDatasetRequest,
user: User | None = Depends(get_current_user_optional),
) -> dict[str, Any]:
"""Partition a dataset file into train/val/test splits."""
path = _safe_dataset_path(req.path)
out_dir = _safe_output_dir(req.output_dir)

try:
manifest = DatasetProfiler.split_dataset(
path,
output_dir=out_dir,
train_ratio=req.train_ratio,
val_ratio=req.val_ratio,
test_ratio=req.test_ratio,
stratify_column=req.stratify_column,
seed=req.seed,
)
return {
"success": True,
"manifest": manifest,
}
except Exception as e:
logger.exception("Failed splitting dataset '%s': %s", req.path, e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Error splitting dataset: {e}",
)
Loading
Loading