feat: add Data Recipe core functionality with job manager, API routes, and validation services

This commit is contained in:
Shine1i 2026-02-15 13:43:46 +01:00
commit 85653237ea
12 changed files with 1030 additions and 3 deletions

View file

@ -0,0 +1,7 @@
"""
Data Recipe core (DataDesigner wrapper + job runner).
"""
from .jobs import JobManager, get_job_manager
__all__ = ["JobManager", "get_job_manager"]

View file

@ -0,0 +1,4 @@
from .manager import JobManager, get_job_manager
__all__ = ["JobManager", "get_job_manager"]

View file

@ -0,0 +1,293 @@
from __future__ import annotations
import asyncio
import json
import queue
import threading
import time
import uuid
from collections import deque
from dataclasses import dataclass
from typing import Any
import multiprocessing as mp
from .parse import apply_update, coerce_event, parse_log_message
from .types import Job
from .worker import run_job_process
_CTX = mp.get_context("spawn")
@dataclass
class Subscription:
replay: list[dict]
_q: queue.Queue
_next_id: int = 0
async def next_event(self, *, timeout_sec: float) -> dict | None:
"""Wait for next event (SSE), w/ timeout so we can check disconnects."""
try:
return await asyncio.to_thread(self._q.get, True, timeout_sec)
except queue.Empty:
return None
def format_sse(self, event: dict) -> bytes:
"""Turn event dict into SSE bytes (id/event/data)."""
self._next_id += 1
body = json.dumps(event, separators=(",", ":"), ensure_ascii=False)
event_type = event.get("type") or "message"
return (
f"id: {self._next_id}\n"
f"event: {event_type}\n"
f"data: {body}\n\n"
).encode("utf-8")
class JobManager:
def __init__(self) -> None:
"""Single-job runner (in-mem). Simple on purpose, not a whole platform."""
self._lock = threading.Lock()
self._job: Job | None = None
self._proc: mp.Process | None = None
self._mp_q: Any | None = None
self._events: deque[dict] = deque(maxlen=5000)
self._subs: list[queue.Queue] = []
self._pump_thread: threading.Thread | None = None
def start(self, *, recipe: dict, run: dict) -> str:
"""Spawn the job subprocess (one at a time, no cap)."""
with self._lock:
if self._proc is not None and self._proc.is_alive():
raise RuntimeError("job already running")
job_id = uuid.uuid4().hex
self._job = Job(job_id=job_id, status="pending", started_at=time.time())
self._events.clear()
mp_q = _CTX.Queue()
proc = _CTX.Process(
target=run_job_process,
kwargs={"event_queue": mp_q, "recipe": recipe, "run": run},
daemon=True,
)
proc.start()
self._mp_q = mp_q
self._proc = proc
self._pump_thread = threading.Thread(target=self._pump_loop, daemon=True)
self._pump_thread.start()
self._emit({"type": "job.enqueued", "ts": time.time(), "job_id": job_id})
return job_id
def cancel(self, job_id: str) -> bool:
"""Hard stop. We terminate the subprocess. Quick + reliable."""
with self._lock:
if self._job is None or self._job.job_id != job_id:
return False
if self._proc is None or not self._proc.is_alive():
return True
self._job.status = "cancelling"
self._emit({"type": "job.cancelling", "ts": time.time(), "job_id": job_id})
try:
self._proc.terminate()
except Exception:
pass
return True
def get_status(self, job_id: str) -> dict | None:
"""UI-friendly snapshot. Poll this if you don't want SSE."""
with self._lock:
if self._job is None or self._job.job_id != job_id:
return None
job = self._job
return {
"job_id": job.job_id,
"status": job.status,
"stage": job.stage,
"current_column": job.current_column,
"batch": {"idx": job.batch.idx, "total": job.batch.total},
"progress": {
"done": job.progress.done,
"total": job.progress.total,
"percent": job.progress.percent,
"eta_sec": job.progress.eta_sec,
"rate": job.progress.rate,
"ok": job.progress.ok,
"failed": job.progress.failed,
},
"model_usage": {
name: {
"model": usage.model,
"tokens": {
"input": usage.input_tokens,
"output": usage.output_tokens,
"total": usage.total_tokens,
"tps": usage.tps,
},
"requests": {
"success": usage.requests_success,
"failed": usage.requests_failed,
"total": usage.requests_total,
"rpm": usage.rpm,
},
}
for name, usage in job.model_usage.items()
},
"rows": job.rows,
"cols": job.cols,
"error": job.error,
"started_at": job.started_at,
"finished_at": job.finished_at,
}
def get_current_status(self) -> dict | None:
"""Single-job convenience (last/current)."""
job_id = self.get_current_job_id()
if job_id is None:
return None
return self.get_status(job_id)
def get_current_job_id(self) -> str | None:
"""Return current job_id (or None)."""
with self._lock:
return None if self._job is None else self._job.job_id
def get_analysis(self, job_id: str) -> dict | None:
"""Final profiling output (only after job completes)."""
with self._lock:
if self._job is None or self._job.job_id != job_id:
return None
return self._job.analysis
def subscribe(self, job_id: str) -> Subscription | None:
"""SSE subscribe: get replay buffer + live events stream."""
with self._lock:
if self._job is None or self._job.job_id != job_id:
return None
q: queue.Queue = queue.Queue(maxsize=2000)
self._subs.append(q)
return Subscription(replay=list(self._events), _q=q)
def unsubscribe(self, sub: Subscription) -> None:
"""Drop SSE subscriber (client disconnected)."""
with self._lock:
self._subs = [q for q in self._subs if q is not sub._q]
def _emit(self, event: dict) -> None:
"""Broadcast event to replay buffer + all subscribers."""
self._events.append(event)
stale: list[queue.Queue] = []
for q in self._subs:
try:
q.put_nowait(event)
except Exception:
stale.append(q)
if stale:
self._subs = [q for q in self._subs if q not in stale]
def _snapshot(self) -> tuple[Job, mp.Process, Any] | None:
"""Grab pointers for the pump loop (avoid holding lock too long)."""
with self._lock:
if self._job is None or self._proc is None or self._mp_q is None:
return None
return self._job, self._proc, self._mp_q
@staticmethod
def _read_queue_with_timeout(q: Any, *, timeout_sec: float) -> dict | None:
"""Try read 1 event from mp queue. Timeout = pump stays responsive."""
try:
return coerce_event(q.get(timeout=timeout_sec))
except queue.Empty:
return None
except Exception:
return None
@staticmethod
def _drain_queue(q: Any) -> list[dict]:
"""Drain mp queue fast (used on process exit)."""
events: list[dict] = []
while True:
try:
events.append(coerce_event(q.get_nowait()))
except queue.Empty:
return events
except Exception:
return events
def _pump_loop(self) -> None:
"""Background thread: consumes worker events + updates job snapshot."""
while True:
snap = self._snapshot()
if snap is None:
return
job, proc, mp_q = snap
event = self._read_queue_with_timeout(mp_q, timeout_sec=0.25)
if event is not None:
self._handle_event(job, event)
continue
if proc.is_alive():
continue
for e in self._drain_queue(mp_q):
self._handle_event(job, e)
with self._lock:
if self._job and self._job.status in {"pending", "active", "cancelling"}:
if self._job.status == "cancelling":
self._job.status = "cancelled"
else:
self._job.status = "error"
self._job.error = self._job.error or "process exited"
self._job.finished_at = time.time()
self._emit(
{
"type": f"job.{self._job.status}",
"ts": time.time(),
"job_id": self._job.job_id,
}
)
return
def _handle_event(self, job: Job, event: dict) -> None:
"""Apply event -> job state + forward to SSE."""
et = event.get("type")
msg = event.get("message") if et == "log" else None
with self._lock:
if self._job is None or self._job.job_id != job.job_id:
return
if et == "job.started":
self._job.status = "active"
if et == "job.completed":
self._job.status = "completed"
self._job.finished_at = time.time()
self._job.analysis = event.get("analysis")
self._job.artifact_path = event.get("artifact_path")
if et == "job.error":
self._job.status = "error"
self._job.finished_at = time.time()
self._job.error = event.get("error") or "error"
if msg:
upd = parse_log_message(msg)
if upd:
apply_update(self._job, upd)
self._emit(event)
_JOB_MANAGER: JobManager | None = None
def get_job_manager() -> JobManager:
"""Singleton JobManager (we only run 1 job anyway)."""
global _JOB_MANAGER
if _JOB_MANAGER is None:
_JOB_MANAGER = JobManager()
return _JOB_MANAGER

View file

@ -0,0 +1,200 @@
from __future__ import annotations
import re
from dataclasses import dataclass
from typing import Any
from .types import Job, ModelUsage, Progress
@dataclass(frozen=True)
class ParsedUpdate:
stage: str | None = None
current_column: str | None = None
progress: Progress | None = None
rows: int | None = None
cols: int | None = None
batch_idx: int | None = None
batch_total: int | None = None
usage_model: str | None = None
usage_input_tokens: int | None = None
usage_output_tokens: int | None = None
usage_total_tokens: int | None = None
usage_tps: float | None = None
usage_requests_success: int | None = None
usage_requests_failed: int | None = None
usage_requests_total: int | None = None
usage_rpm: float | None = None
usage_section_start: bool | None = None
# welp, best effort to parse the logs and convert them to structured information so we can access it and read it properly on the client
# i couldnt find what datadesigner do for progress tracking besides the logs so il probably raise a pr in their repo to add some sort of progress monitoring
_RE_SAMPLERS = re.compile(
r"Preparing samplers to generate (?P<rows>\\d+) records across (?P<cols>\\d+) columns"
)
_RE_COLCFG = re.compile(r"model config for column '(?P<col>[^']+)'")
_RE_PROCESSING_COL = re.compile(r"Processing .* column '(?P<col>[^']+)'")
_RE_PROGRESS = re.compile(
r"progress: (?P<done>\\d+)/(?P<total>\\d+) \\((?P<pct>\\d+)%\\) complete, "
r"(?P<ok>\\d+) ok, (?P<failed>\\d+) failed, (?P<rate>[0-9.]+) rec/s, eta (?P<eta>[0-9.]+)s"
)
_RE_BATCH = re.compile(r"Processing batch (?P<idx>\\d+) of (?P<total>\\d+)")
_RE_USAGE_MODEL = re.compile(r"model:\\s*(?P<model>.+)$")
_RE_USAGE_TOKENS = re.compile(
r"tokens:\\s*input=(?P<input>\\d+),\\s*output=(?P<output>\\d+),\\s*total=(?P<total>\\d+),\\s*tps=(?P<tps>[0-9.]+)"
)
_RE_USAGE_REQUESTS = re.compile(
r"requests:\\s*success=(?P<success>\\d+),\\s*failed=(?P<failed>\\d+),\\s*total=(?P<total>\\d+),\\s*rpm=(?P<rpm>[0-9.]+)"
)
def parse_log_message(msg: str) -> ParsedUpdate | None:
m = _RE_SAMPLERS.search(msg)
if m:
return ParsedUpdate(
stage="sampling",
rows=int(m.group("rows")),
cols=int(m.group("cols")),
)
if "Sorting column configs into a Directed Acyclic Graph" in msg:
return ParsedUpdate(stage="dag")
if "Running health checks for models" in msg:
return ParsedUpdate(stage="healthcheck")
if "Preview generation in progress" in msg:
return ParsedUpdate(stage="preview")
if "Creating Data Designer dataset" in msg:
return ParsedUpdate(stage="create")
if "Measuring dataset column statistics" in msg:
return ParsedUpdate(stage="profiling")
m = _RE_COLCFG.search(msg)
if m:
col = m.group("col")
return ParsedUpdate(stage="column_config", current_column=col)
m = _RE_PROCESSING_COL.search(msg)
if m:
col = m.group("col")
return ParsedUpdate(stage="generating", current_column=col)
m = _RE_PROGRESS.search(msg)
if m:
p = Progress(
done=int(m.group("done")),
total=int(m.group("total")),
percent=float(m.group("pct")),
ok=int(m.group("ok")),
failed=int(m.group("failed")),
rate=float(m.group("rate")),
eta_sec=float(m.group("eta")),
)
return ParsedUpdate(stage="generating", progress=p)
m = _RE_BATCH.search(msg)
if m:
return ParsedUpdate(
stage="batch",
batch_idx=int(m.group("idx")),
batch_total=int(m.group("total")),
)
if "Model usage summary" in msg:
return ParsedUpdate(usage_section_start=True)
m = _RE_USAGE_MODEL.search(msg)
if m and " |-- model:" in msg:
return ParsedUpdate(usage_model=str(m.group("model")).strip())
m = _RE_USAGE_TOKENS.search(msg)
if m:
return ParsedUpdate(
usage_input_tokens=int(m.group("input")),
usage_output_tokens=int(m.group("output")),
usage_total_tokens=int(m.group("total")),
usage_tps=float(m.group("tps")),
)
m = _RE_USAGE_REQUESTS.search(msg)
if m:
return ParsedUpdate(
usage_requests_success=int(m.group("success")),
usage_requests_failed=int(m.group("failed")),
usage_requests_total=int(m.group("total")),
usage_rpm=float(m.group("rpm")),
)
return None
def apply_update(job: Job, update: ParsedUpdate) -> None:
if update.stage is not None:
job.stage = update.stage
if update.current_column is not None:
job.current_column = update.current_column
if update.rows is not None:
job.rows = update.rows
if update.cols is not None:
job.cols = update.cols
if update.progress is not None:
job.progress = update.progress
if update.batch_idx is not None:
job.batch.idx = update.batch_idx
if update.batch_total is not None:
job.batch.total = update.batch_total
if update.stage in {
"profiling",
"generating",
"sampling",
"healthcheck",
"dag",
"create",
"preview",
}:
# usage summary is a short block; reset once we move into the next stage.
job._in_usage_summary = False
if update.usage_section_start is not None:
job._in_usage_summary = update.usage_section_start
if update.usage_section_start:
job._current_usage_model = None
if not job._in_usage_summary:
return
if update.usage_model is not None:
name = update.usage_model.strip().strip("'").strip('"')
job._current_usage_model = name
if name not in job.model_usage:
job.model_usage[name] = ModelUsage(model=name)
if job._current_usage_model is None:
return
usage = job.model_usage.get(job._current_usage_model)
if usage is None:
return
if update.usage_input_tokens is not None:
usage.input_tokens = update.usage_input_tokens
if update.usage_output_tokens is not None:
usage.output_tokens = update.usage_output_tokens
if update.usage_total_tokens is not None:
usage.total_tokens = update.usage_total_tokens
if update.usage_tps is not None:
usage.tps = update.usage_tps
if update.usage_requests_success is not None:
usage.requests_success = update.usage_requests_success
if update.usage_requests_failed is not None:
usage.requests_failed = update.usage_requests_failed
if update.usage_requests_total is not None:
usage.requests_total = update.usage_requests_total
if update.usage_rpm is not None:
usage.rpm = update.usage_rpm
def coerce_event(obj: Any) -> dict:
# worker sends dict already
return obj if isinstance(obj, dict) else {"type": "log", "message": str(obj)}

View file

@ -0,0 +1,67 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
JobStatus = Literal[
"created",
"pending",
"active",
"cancelling",
"cancelled",
"error",
"completed",
]
@dataclass
class Progress:
done: int | None = None
total: int | None = None
percent: float | None = None
eta_sec: float | None = None
rate: float | None = None
ok: int | None = None
failed: int | None = None
@dataclass
class BatchProgress:
idx: int | None = None
total: int | None = None
@dataclass
class ModelUsage:
model: str
input_tokens: int | None = None
output_tokens: int | None = None
total_tokens: int | None = None
tps: float | None = None
requests_success: int | None = None
requests_failed: int | None = None
requests_total: int | None = None
rpm: float | None = None
@dataclass
class Job:
job_id: str
status: JobStatus = "created"
stage: str | None = None
current_column: str | None = None
progress: Progress = field(default_factory=Progress)
batch: BatchProgress = field(default_factory=BatchProgress)
rows: int | None = None
cols: int | None = None
error: str | None = None
started_at: float | None = None
finished_at: float | None = None
analysis: dict[str, Any] | None = None
artifact_path: str | None = None
model_usage: dict[str, ModelUsage] = field(default_factory=dict)
_current_usage_model: str | None = None
_in_usage_summary: bool = False

View file

@ -0,0 +1,89 @@
from __future__ import annotations
import logging
import time
import traceback
from typing import Any
from ..service import build_config_builder, create_data_designer
class _QueueLogHandler(logging.Handler):
def __init__(self, event_queue):
super().__init__()
self._q = event_queue
def emit(self, record: logging.LogRecord) -> None:
try:
event = {
"type": "log",
"ts": record.created,
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
}
self._q.put(event)
except Exception:
pass
def run_job_process(
*,
event_queue,
recipe: dict[str, Any],
run: dict[str, Any],
) -> None:
"""
Subprocess entrypoint.
Sends events to `event_queue`.
"""
handler = _QueueLogHandler(event_queue)
handler.setLevel(logging.INFO)
# Attach once to root. data_designer.* loggers should propagate to root.
root = logging.getLogger()
root.addHandler(handler)
root.setLevel(logging.INFO)
dd = logging.getLogger("data_designer")
dd.setLevel(logging.INFO)
dd.propagate = True
event_queue.put({"type": "job.started", "ts": time.time()})
try:
# lazy import so backend can boot even if data-designer isn't installed yet
from data_designer.config.run_config import RunConfig
rows = int(run.get("rows") or 1000)
dataset_name = str(run.get("dataset_name") or "dataset")
run_config_raw = run.get("run_config") or {}
builder = build_config_builder(recipe)
designer = create_data_designer(recipe)
if run_config_raw:
designer.set_run_config(RunConfig.model_validate(run_config_raw))
results = designer.create(builder, num_records=rows, dataset_name=dataset_name)
analysis = results.load_analysis().model_dump(mode="json")
artifact_path = str(results.artifact_storage.base_dataset_path)
event_queue.put(
{
"type": "job.completed",
"ts": time.time(),
"analysis": analysis,
"artifact_path": artifact_path,
}
)
except Exception as exc:
event_queue.put(
{
"type": "job.error",
"ts": time.time(),
"error": str(exc),
"stack": traceback.format_exc(limit=20),
}
)

View file

@ -0,0 +1,147 @@
from __future__ import annotations
import os
from typing import Any
def _to_jsonable(value: Any) -> Any:
# pydantic/fastapi can't serialize numpy arrays/scalars.
try:
import numpy as np # type: ignore
except Exception: # pragma: no cover
np = None # type: ignore
if np is not None:
if isinstance(value, np.ndarray):
return value.tolist()
if isinstance(value, np.generic):
return value.item()
if isinstance(value, dict):
return {str(k): _to_jsonable(v) for k, v in value.items()}
if isinstance(value, (list, tuple, set)):
return [_to_jsonable(v) for v in value]
# pandas Timestamp/date-like
if hasattr(value, "isoformat") and callable(value.isoformat):
try:
return value.isoformat()
except Exception:
pass
return value
def build_model_providers(recipe: dict[str, Any]):
from data_designer.config.default_model_settings import get_default_providers
from data_designer.config.models import ModelProvider
providers: list[ModelProvider] = []
for provider in recipe.get("model_providers", []):
api_key = provider.get("api_key")
api_key_env = provider.get("api_key_env")
if not api_key and api_key_env:
api_key = os.getenv(api_key_env)
providers.append(
ModelProvider(
name=provider["name"],
endpoint=provider["endpoint"],
provider_type=provider.get("provider_type", "openai"),
api_key=api_key,
extra_headers=provider.get("extra_headers"),
extra_body=provider.get("extra_body"),
)
)
# DataDesigner currently expects at least one provider even if they only use static samplers,
# but it's fine it gives a warning only.
return providers or get_default_providers()
def build_mcp_providers(
recipe: dict[str, Any],
) -> list:
from data_designer.config.mcp import LocalStdioMCPProvider, MCPProvider
providers: list[MCPProvider | LocalStdioMCPProvider] = []
for provider in recipe.get("mcp_providers", []):
if not isinstance(provider, dict):
continue
provider_type = provider.get("provider_type")
if provider_type == "stdio":
env = provider.get("env")
if not isinstance(env, dict):
env = {}
args = provider.get("args")
if not isinstance(args, list):
args = []
providers.append(
LocalStdioMCPProvider(
name=str(provider.get("name", "")),
command=str(provider.get("command", "")),
args=[str(value) for value in args],
env={str(key): str(value) for key, value in env.items()},
)
)
continue
if provider_type in {"sse", "streamable_http"}:
api_key = provider.get("api_key")
api_key_env = provider.get("api_key_env")
if not api_key and api_key_env:
api_key = os.getenv(str(api_key_env))
providers.append(
MCPProvider(
name=str(provider.get("name", "")),
endpoint=str(provider.get("endpoint", "")),
api_key=str(api_key) if api_key else None,
)
)
return providers
def build_config_builder(recipe: dict[str, Any]):
from data_designer.config import DataDesignerConfigBuilder
recipe_core = {
key: value
for key, value in recipe.items()
if key not in {"model_providers", "mcp_providers"}
}
return DataDesignerConfigBuilder.from_config({"data_designer": recipe_core})
def create_data_designer(recipe: dict[str, Any]):
from data_designer.interface.data_designer import DataDesigner
return DataDesigner(
model_providers=build_model_providers(recipe),
mcp_providers=build_mcp_providers(recipe),
)
def validate_recipe(recipe: dict[str, Any]) -> None:
builder = build_config_builder(recipe)
designer = create_data_designer(recipe)
designer.validate(builder)
def preview_recipe(
recipe: dict[str, Any],
num_records: int,
) -> tuple[list[dict[str, Any]], dict[str, Any] | None]:
builder = build_config_builder(recipe)
designer = create_data_designer(recipe)
results = designer.preview(builder, num_records=num_records)
dataset: list[dict[str, Any]] = []
if results.dataset is not None:
raw_rows = results.dataset.to_dict(orient="records")
dataset = [_to_jsonable(row) for row in raw_rows]
artifacts = (
None
if results.processor_artifacts is None
else _to_jsonable(results.processor_artifacts)
)
return dataset, artifacts

View file

@ -13,7 +13,14 @@ from pathlib import Path
from datetime import datetime
# Import routers
from routes import training_router, models_router, inference_router, datasets_router, auth_router
from routes import (
training_router,
models_router,
inference_router,
datasets_router,
auth_router,
data_recipe_router,
)
from auth import storage
from utils.hardware import detect_hardware
import utils.hardware.hardware as _hw_module
@ -67,6 +74,7 @@ app.include_router(training_router, prefix="/api/train", tags=["training"])
app.include_router(models_router, prefix="/api/models", tags=["models"])
app.include_router(inference_router, prefix="/api/inference", tags=["inference"])
app.include_router(datasets_router, prefix="/api/datasets", tags=["datasets"])
app.include_router(data_recipe_router, prefix="/api/data-recipe", tags=["data-recipe"])
# ============ Health and System Endpoints ============
@ -143,4 +151,3 @@ def setup_frontend(app: FastAPI, build_path: Path):
return True
return False

View file

@ -38,6 +38,13 @@ from .responses import (
LoRABaseModelResponse,
VisionCheckResponse,
)
from .data_recipe import (
RecipePayload,
PreviewResponse,
ValidateError,
ValidateResponse,
JobCreateResponse,
)
__all__ = [
# Training schemas
@ -71,4 +78,10 @@ __all__ = [
"TrainingMetricsResponse",
"LoRABaseModelResponse",
"VisionCheckResponse",
# Data recipe
"RecipePayload",
"PreviewResponse",
"ValidateError",
"ValidateResponse",
"JobCreateResponse",
]

View file

@ -0,0 +1,37 @@
"""
Pydantic schemas for Data Recipe (DataDesigner) API.
"""
from __future__ import annotations
from typing import Any
from pydantic import BaseModel, Field
class RecipePayload(BaseModel):
recipe: dict[str, Any] = Field(default_factory=dict)
run: dict[str, Any] | None = None
ui: dict[str, Any] | None = None
class PreviewResponse(BaseModel):
dataset: list[dict[str, Any]] = Field(default_factory=list)
processor_artifacts: dict[str, Any] | None = None
class ValidateError(BaseModel):
message: str
path: str | None = None
code: str | None = None
class ValidateResponse(BaseModel):
valid: bool
errors: list[ValidateError] = Field(default_factory=list)
raw_detail: str | None = None
class JobCreateResponse(BaseModel):
job_id: str

View file

@ -7,5 +7,13 @@ from routes.models import router as models_router
from routes.inference import router as inference_router
from routes.datasets import router as datasets_router
from routes.auth import router as auth_router
from routes.data_recipe import router as data_recipe_router
__all__ = ["training_router", "models_router", "inference_router", "datasets_router", "auth_router"]
__all__ = [
"training_router",
"models_router",
"inference_router",
"datasets_router",
"auth_router",
"data_recipe_router",
]

View file

@ -0,0 +1,155 @@
"""
Data Recipe routes (DataDesigner runner).
"""
from __future__ import annotations
import sys
from pathlib import Path
from typing import Any
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import JSONResponse, StreamingResponse
# same thing as other files do
backend_path = Path(__file__).parent.parent.parent
if str(backend_path) not in sys.path:
sys.path.insert(0, str(backend_path))
from core.data_recipe.jobs import get_job_manager
from core.data_recipe.service import preview_recipe, validate_recipe
from models.data_recipe import JobCreateResponse, PreviewResponse, RecipePayload, ValidateError, ValidateResponse
router = APIRouter()
@router.post("/validate", response_model=ValidateResponse)
def validate(payload: RecipePayload) -> ValidateResponse:
recipe = payload.recipe
if not recipe.get("columns"):
return ValidateResponse(
valid=False,
errors=[ValidateError(message="Recipe must include columns.")],
)
try:
validate_recipe(recipe)
except RuntimeError as exc:
raise HTTPException(status_code=503, detail=str(exc)) from exc
except Exception as exc:
detail = str(exc).strip() or "Validation failed."
return ValidateResponse(
valid=False,
errors=[ValidateError(message=detail)],
raw_detail=detail,
)
return ValidateResponse(valid=True)
@router.post("/preview", response_model=PreviewResponse)
def preview(payload: RecipePayload) -> PreviewResponse:
recipe = payload.recipe
if not recipe.get("columns"):
raise HTTPException(status_code=400, detail="Recipe must include columns.")
run = payload.run or {}
num_records = int(run.get("rows") or 5)
try:
dataset, artifacts = preview_recipe(recipe, num_records)
except RuntimeError as exc:
raise HTTPException(status_code=503, detail=str(exc)) from exc
except Exception as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return PreviewResponse(dataset=dataset, processor_artifacts=artifacts)
@router.post("/jobs", response_class=JSONResponse, response_model=JobCreateResponse)
def create_job(payload: RecipePayload):
recipe = payload.recipe
if not recipe.get("columns"):
raise HTTPException(status_code=400, detail="Recipe must include columns.")
run: dict[str, Any] = payload.run or {}
run_config_raw = run.get("run_config")
if run_config_raw is not None:
try:
from data_designer.config.run_config import RunConfig
RunConfig.model_validate(run_config_raw)
except Exception as exc:
raise HTTPException(status_code=400, detail=f"invalid run_config: {exc}") from exc
mgr = get_job_manager()
try:
job_id = mgr.start(recipe=recipe, run=run)
except RuntimeError as exc:
raise HTTPException(status_code=409, detail=str(exc)) from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return {"job_id": job_id}
@router.get("/jobs/{job_id}/status")
def job_status(job_id: str):
mgr = get_job_manager()
state = mgr.get_status(job_id)
if state is None:
raise HTTPException(status_code=404, detail="job not found")
return state
@router.get("/jobs/current")
def current_job():
mgr = get_job_manager()
state = mgr.get_current_status()
if state is None:
raise HTTPException(status_code=404, detail="no job")
return state
@router.post("/jobs/{job_id}/cancel")
def cancel_job(job_id: str):
mgr = get_job_manager()
ok = mgr.cancel(job_id)
if not ok:
raise HTTPException(status_code=404, detail="job not found")
return mgr.get_status(job_id)
@router.get("/jobs/{job_id}/analysis")
def job_analysis(job_id: str):
mgr = get_job_manager()
analysis = mgr.get_analysis(job_id)
if analysis is None:
raise HTTPException(status_code=404, detail="analysis not ready")
return analysis
@router.get("/jobs/{job_id}/events")
async def job_events(request: Request, job_id: str):
mgr = get_job_manager()
sub = mgr.subscribe(job_id)
if sub is None:
raise HTTPException(status_code=404, detail="job not found")
async def gen():
try:
for event in sub.replay:
yield sub.format_sse(event)
while True:
if await request.is_disconnected():
break
event = await sub.next_event(timeout_sec=1.0)
if event is None:
continue
yield sub.format_sse(event)
finally:
mgr.unsubscribe(sub)
return StreamingResponse(gen(), media_type="text/event-stream")