feat: add Data Recipe core functionality with job manager, API routes, and validation services
This commit is contained in:
parent
8bd7cfaab1
commit
85653237ea
12 changed files with 1030 additions and 3 deletions
7
studio/backend/core/data_recipe/__init__.py
Normal file
7
studio/backend/core/data_recipe/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
Data Recipe core (DataDesigner wrapper + job runner).
|
||||
"""
|
||||
|
||||
from .jobs import JobManager, get_job_manager
|
||||
|
||||
__all__ = ["JobManager", "get_job_manager"]
|
||||
4
studio/backend/core/data_recipe/jobs/__init__.py
Normal file
4
studio/backend/core/data_recipe/jobs/__init__.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
from .manager import JobManager, get_job_manager
|
||||
|
||||
__all__ = ["JobManager", "get_job_manager"]
|
||||
|
||||
293
studio/backend/core/data_recipe/jobs/manager.py
Normal file
293
studio/backend/core/data_recipe/jobs/manager.py
Normal 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
|
||||
|
||||
200
studio/backend/core/data_recipe/jobs/parse.py
Normal file
200
studio/backend/core/data_recipe/jobs/parse.py
Normal 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)}
|
||||
|
||||
67
studio/backend/core/data_recipe/jobs/types.py
Normal file
67
studio/backend/core/data_recipe/jobs/types.py
Normal 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
|
||||
|
||||
89
studio/backend/core/data_recipe/jobs/worker.py
Normal file
89
studio/backend/core/data_recipe/jobs/worker.py
Normal 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),
|
||||
}
|
||||
)
|
||||
|
||||
147
studio/backend/core/data_recipe/service.py
Normal file
147
studio/backend/core/data_recipe/service.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
37
studio/backend/models/data_recipe.py
Normal file
37
studio/backend/models/data_recipe.py
Normal 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
|
||||
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
155
studio/backend/routes/data_recipe.py
Normal file
155
studio/backend/routes/data_recipe.py
Normal 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")
|
||||
|
||||
Loading…
Add table
Add a link
Reference in a new issue