diff --git a/studio/backend/core/data_recipe/__init__.py b/studio/backend/core/data_recipe/__init__.py new file mode 100644 index 0000000000..0665a8e4c2 --- /dev/null +++ b/studio/backend/core/data_recipe/__init__.py @@ -0,0 +1,7 @@ +""" +Data Recipe core (DataDesigner wrapper + job runner). +""" + +from .jobs import JobManager, get_job_manager + +__all__ = ["JobManager", "get_job_manager"] diff --git a/studio/backend/core/data_recipe/jobs/__init__.py b/studio/backend/core/data_recipe/jobs/__init__.py new file mode 100644 index 0000000000..bac519ee3a --- /dev/null +++ b/studio/backend/core/data_recipe/jobs/__init__.py @@ -0,0 +1,4 @@ +from .manager import JobManager, get_job_manager + +__all__ = ["JobManager", "get_job_manager"] + diff --git a/studio/backend/core/data_recipe/jobs/manager.py b/studio/backend/core/data_recipe/jobs/manager.py new file mode 100644 index 0000000000..eae789ed50 --- /dev/null +++ b/studio/backend/core/data_recipe/jobs/manager.py @@ -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 + diff --git a/studio/backend/core/data_recipe/jobs/parse.py b/studio/backend/core/data_recipe/jobs/parse.py new file mode 100644 index 0000000000..612b491f41 --- /dev/null +++ b/studio/backend/core/data_recipe/jobs/parse.py @@ -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\\d+) records across (?P\\d+) columns" +) +_RE_COLCFG = re.compile(r"model config for column '(?P[^']+)'") +_RE_PROCESSING_COL = re.compile(r"Processing .* column '(?P[^']+)'") +_RE_PROGRESS = re.compile( + r"progress: (?P\\d+)/(?P\\d+) \\((?P\\d+)%\\) complete, " + r"(?P\\d+) ok, (?P\\d+) failed, (?P[0-9.]+) rec/s, eta (?P[0-9.]+)s" +) +_RE_BATCH = re.compile(r"Processing batch (?P\\d+) of (?P\\d+)") +_RE_USAGE_MODEL = re.compile(r"model:\\s*(?P.+)$") +_RE_USAGE_TOKENS = re.compile( + r"tokens:\\s*input=(?P\\d+),\\s*output=(?P\\d+),\\s*total=(?P\\d+),\\s*tps=(?P[0-9.]+)" +) +_RE_USAGE_REQUESTS = re.compile( + r"requests:\\s*success=(?P\\d+),\\s*failed=(?P\\d+),\\s*total=(?P\\d+),\\s*rpm=(?P[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)} + diff --git a/studio/backend/core/data_recipe/jobs/types.py b/studio/backend/core/data_recipe/jobs/types.py new file mode 100644 index 0000000000..8e203f2add --- /dev/null +++ b/studio/backend/core/data_recipe/jobs/types.py @@ -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 + diff --git a/studio/backend/core/data_recipe/jobs/worker.py b/studio/backend/core/data_recipe/jobs/worker.py new file mode 100644 index 0000000000..1fb54454ab --- /dev/null +++ b/studio/backend/core/data_recipe/jobs/worker.py @@ -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), + } + ) + diff --git a/studio/backend/core/data_recipe/service.py b/studio/backend/core/data_recipe/service.py new file mode 100644 index 0000000000..736a40e967 --- /dev/null +++ b/studio/backend/core/data_recipe/service.py @@ -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 diff --git a/studio/backend/main.py b/studio/backend/main.py index e957c668e3..6b5242eb28 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -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 - diff --git a/studio/backend/models/__init__.py b/studio/backend/models/__init__.py index 584fb9f7c2..8a66cd4c06 100644 --- a/studio/backend/models/__init__.py +++ b/studio/backend/models/__init__.py @@ -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", ] diff --git a/studio/backend/models/data_recipe.py b/studio/backend/models/data_recipe.py new file mode 100644 index 0000000000..a906586552 --- /dev/null +++ b/studio/backend/models/data_recipe.py @@ -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 + diff --git a/studio/backend/routes/__init__.py b/studio/backend/routes/__init__.py index 5a16125a64..84bd513682 100644 --- a/studio/backend/routes/__init__.py +++ b/studio/backend/routes/__init__.py @@ -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"] \ No newline at end of file +__all__ = [ + "training_router", + "models_router", + "inference_router", + "datasets_router", + "auth_router", + "data_recipe_router", +] diff --git a/studio/backend/routes/data_recipe.py b/studio/backend/routes/data_recipe.py new file mode 100644 index 0000000000..a7bf901b0d --- /dev/null +++ b/studio/backend/routes/data_recipe.py @@ -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") +