diff --git a/frontend/src/api.ts b/frontend/src/api.ts index 3c2e8ee..3f9a8b6 100644 --- a/frontend/src/api.ts +++ b/frontend/src/api.ts @@ -72,7 +72,8 @@ export type HistoryEntry = { details?: string | null; }; -export type ExpressionDiagnostic = { +export type FormError = { + path: string[]; message: string; start_line: number; start_column: number; @@ -80,17 +81,17 @@ export type ExpressionDiagnostic = { end_column: number; }; -export async function validateExpression( - expression: string, -): Promise { - const r = await fetch("/api/expressions/validate", { +export async function validateTaskForm( + taskId: string, + formData: Record, +): Promise { + const r = await fetch(`/api/tasks/${taskId}/validate`, { method: "POST", headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ expression }), + body: JSON.stringify(formData), }); if (!r.ok) throw new Error(`validate: ${r.status}`); - const body = (await r.json()) as { diagnostics: ExpressionDiagnostic[] }; - return body.diagnostics; + return r.json(); } export async function fetchHistory(): Promise { diff --git a/frontend/src/components/ExpressionWidget.test.tsx b/frontend/src/components/ExpressionWidget.test.tsx index af3c0e1..d953c57 100644 --- a/frontend/src/components/ExpressionWidget.test.tsx +++ b/frontend/src/components/ExpressionWidget.test.tsx @@ -3,14 +3,6 @@ import type { WidgetProps } from "@rjsf/utils"; import { describe, expect, it, vi } from "vitest"; import { ExpressionWidget, findToken } from "./ExpressionWidget"; -vi.mock("../api", async () => { - const actual = await vi.importActual("../api"); - return { - ...actual, - validateExpression: vi.fn().mockResolvedValue([]), - }; -}); - const tokens = [ { label: "col", @@ -35,6 +27,22 @@ function renderWidget(onChange = vi.fn(), value = "") { onChange, options: { tokens }, schema: { type: "string" }, + registry: { + formContext: { + expressionErrors: { + root_expression: [ + { + path: ["expression"], + message: "unknown name 'foo'", + start_line: 1, + start_column: 1, + end_line: 1, + end_column: 4, + }, + ], + }, + }, + }, } as unknown as WidgetProps; return render(); } diff --git a/frontend/src/components/ExpressionWidget.tsx b/frontend/src/components/ExpressionWidget.tsx index 17d2de6..34e93d2 100644 --- a/frontend/src/components/ExpressionWidget.tsx +++ b/frontend/src/components/ExpressionWidget.tsx @@ -4,7 +4,7 @@ import { useTheme } from "@mui/material/styles"; import type { WidgetProps } from "@rjsf/utils"; import type { editor, Position } from "monaco-editor"; import { useEffect, useRef } from "react"; -import { type ExpressionDiagnostic, validateExpression } from "../api"; +import type { FormError } from "../api"; const LANGUAGE_ID = "hyperleda-expression"; const MARKER_OWNER = "hyperleda-expression"; @@ -16,6 +16,10 @@ export type ExpressionToken = { detail: string; }; +export type ExpressionFormContext = { + expressionErrors?: Record; +}; + let languageRegistered = false; let currentTokens: ExpressionToken[] = []; @@ -47,20 +51,29 @@ export function findToken( return tokens.find((token) => token.label === word); } -function toMarkers( - monaco: Monaco, - diagnostics: ExpressionDiagnostic[], -): editor.IMarkerData[] { - return diagnostics.map((diagnostic) => ({ +function toMarkers(monaco: Monaco, errors: FormError[]): editor.IMarkerData[] { + return errors.map((error) => ({ severity: monaco.MarkerSeverity.Error, - message: diagnostic.message, - startLineNumber: diagnostic.start_line, - startColumn: diagnostic.start_column, - endLineNumber: diagnostic.end_line, - endColumn: diagnostic.end_column, + message: error.message, + startLineNumber: error.start_line, + startColumn: error.start_column, + endLineNumber: error.end_line, + endColumn: error.end_column, })); } +function readFieldErrors(formContext: unknown, fieldId: string): FormError[] { + if (typeof formContext !== "object" || formContext === null) { + return []; + } + const errors = (formContext as ExpressionFormContext).expressionErrors; + if (!errors || typeof errors !== "object") { + return []; + } + const fieldErrors = errors[fieldId]; + return Array.isArray(fieldErrors) ? fieldErrors : []; +} + function registerExpressionLanguage( monaco: Monaco, tokens: ExpressionToken[], @@ -148,54 +161,26 @@ export function ExpressionWidget(props: WidgetProps) { const isDark = theme.palette.mode === "dark"; const editorRef = useRef(null); const monacoRef = useRef(null); - const requestIdRef = useRef(0); - const timeoutRef = useRef(null); + const formContext = + props.registry?.formContext ?? props.formContext ?? undefined; + const fieldErrors = readFieldErrors(formContext, props.id); - useEffect( - () => () => { - if (timeoutRef.current !== null) { - window.clearTimeout(timeoutRef.current); - } - requestIdRef.current += 1; - }, - [], - ); - - function applyMarkers(diagnostics: ExpressionDiagnostic[]) { + useEffect(() => { const editorInstance = editorRef.current; const monaco = monacoRef.current; const model = editorInstance?.getModel(); if (!editorInstance || !monaco || !model) { return; } + const context = + props.registry?.formContext ?? props.formContext ?? undefined; + const errors = readFieldErrors(context, props.id); monaco.editor.setModelMarkers( model, MARKER_OWNER, - toMarkers(monaco, diagnostics), + toMarkers(monaco, errors), ); - } - - function scheduleValidation(text: string) { - if (timeoutRef.current !== null) { - window.clearTimeout(timeoutRef.current); - } - timeoutRef.current = window.setTimeout(() => { - const requestId = ++requestIdRef.current; - validateExpression(text) - .then((diagnostics) => { - if (requestId !== requestIdRef.current) { - return; - } - applyMarkers(diagnostics); - }) - .catch(() => { - if (requestId !== requestIdRef.current) { - return; - } - applyMarkers([]); - }); - }, 200); - } + }, [props.registry?.formContext, props.formContext, props.id]); function handleBeforeMount(monaco: Monaco) { registerExpressionLanguage(monaco, tokens); @@ -213,8 +198,14 @@ export function ExpressionWidget(props: WidgetProps) { e.stopPropagation(); } }); - applyMarkers([]); - scheduleValidation(value.replace(/\r?\n/g, "")); + const model = editorInstance.getModel(); + if (model) { + monaco.editor.setModelMarkers( + model, + MARKER_OWNER, + toMarkers(monaco, fieldErrors), + ); + } } return ( @@ -241,7 +232,6 @@ export function ExpressionWidget(props: WidgetProps) { onChange={(next) => { const singleLine = (next ?? "").replace(/\r?\n/g, ""); props.onChange(singleLine); - scheduleValidation(singleLine); }} beforeMount={handleBeforeMount} onMount={handleMount} diff --git a/frontend/src/components/TaskPage.test.tsx b/frontend/src/components/TaskPage.test.tsx index ad9cced..30f122d 100644 --- a/frontend/src/components/TaskPage.test.tsx +++ b/frontend/src/components/TaskPage.test.tsx @@ -43,6 +43,7 @@ const fakeTaskSchema = { vi.mock("../api", () => ({ fetchTaskSchema: vi.fn(() => Promise.resolve(fakeTaskSchema)), submitTask: vi.fn(), + validateTaskForm: vi.fn(() => Promise.resolve([])), })); function renderTaskPage(taskId: string, formData?: Record) { diff --git a/frontend/src/components/TaskPage.tsx b/frontend/src/components/TaskPage.tsx index c251036..12a0f7c 100644 --- a/frontend/src/components/TaskPage.tsx +++ b/frontend/src/components/TaskPage.tsx @@ -1,4 +1,4 @@ -import { useEffect, useState } from "react"; +import { useEffect, useRef, useState } from "react"; import { useLocation, useParams } from "react-router-dom"; import Form from "@rjsf/mui"; import type { RegistryWidgetsType } from "@rjsf/utils"; @@ -10,7 +10,12 @@ import CircularProgress from "@mui/material/CircularProgress"; import IconButton from "@mui/material/IconButton"; import Popover from "@mui/material/Popover"; import Typography from "@mui/material/Typography"; -import { fetchTaskSchema, submitTask } from "../api"; +import { + type FormError, + fetchTaskSchema, + submitTask, + validateTaskForm, +} from "../api"; import { extractUiSchema } from "../extractUiSchema"; import { ExpressionWidget } from "./ExpressionWidget"; import { FoldableObjectFieldTemplate } from "./FoldableObjectFieldTemplate"; @@ -21,6 +26,19 @@ const widgets: RegistryWidgetsType = { expression: ExpressionWidget, }; +function pathToFieldId(path: string[]): string { + return ["root", ...path].join("_"); +} + +function toExpressionErrors(errors: FormError[]): Record { + const byId: Record = {}; + for (const error of errors) { + const fieldId = pathToFieldId(error.path); + (byId[fieldId] ??= []).push(error); + } + return byId; +} + export function TaskPage() { const { taskId } = useParams<{ taskId: string }>(); const location = useLocation(); @@ -34,13 +52,17 @@ export function TaskPage() { const [loadError, setLoadError] = useState(null); const [submitError, setSubmitError] = useState(null); const [runId, setRunId] = useState(null); - const [prefillData, setPrefillData] = useState>({}); + const [formData, setFormData] = useState>({}); + const [formErrors, setFormErrors] = useState([]); + const requestIdRef = useRef(0); + const timeoutRef = useRef(null); useEffect(() => { const state = location.state as { formData?: Record; } | null; - setPrefillData(state?.formData ?? {}); + setFormData(state?.formData ?? {}); + setFormErrors([]); }, [location.state, taskId]); useEffect(() => { @@ -72,6 +94,37 @@ export function TaskPage() { }; }, [taskId]); + useEffect(() => { + if (!taskId || !schema) { + return () => {}; + } + if (timeoutRef.current !== null) { + window.clearTimeout(timeoutRef.current); + } + timeoutRef.current = window.setTimeout(() => { + const requestId = ++requestIdRef.current; + validateTaskForm(taskId, formData) + .then((errors) => { + if (requestId !== requestIdRef.current) { + return; + } + setFormErrors(errors); + }) + .catch(() => { + if (requestId !== requestIdRef.current) { + return; + } + setFormErrors([]); + }); + }, 200); + return () => { + if (timeoutRef.current !== null) { + window.clearTimeout(timeoutRef.current); + } + requestIdRef.current += 1; + }; + }, [taskId, schema, formData]); + if (!taskId) return null; if (runId) { @@ -157,16 +210,20 @@ export function TaskPage() {
{ + formContext={{ expressionErrors: toExpressionErrors(formErrors) }} + onChange={({ formData: next }) => { + setFormData((next ?? {}) as Record); + }} + onSubmit={async ({ formData: submittedData }) => { setSubmitError(null); try { const submitted = await submitTask( taskId, - formData as Record, + submittedData as Record, ); setRunId(submitted.run_id); } catch (e: unknown) { diff --git a/tests/test_formula_validate.py b/tests/test_formula_validate.py index c1695a1..9c7ad43 100644 --- a/tests/test_formula_validate.py +++ b/tests/test_formula_validate.py @@ -1,6 +1,8 @@ from dataclasses import asdict -from uploader.app.lib.formula import validate_expression +from pydantic import BaseModel, ValidationError + +from uploader.app.lib.formula import ExpressionStr, expression_form_errors, validate_expression def test_validate_expression_accepts_valid() -> None: @@ -37,3 +39,23 @@ def test_validate_expression_reports_unknown_names() -> None: def test_validate_expression_ignores_names_inside_strings() -> None: assert validate_expression('col("foo_bar")') == [] + + +def test_expression_str_puts_diagnostics_in_validation_error_ctx() -> None: + class SampleForm(BaseModel): + name: str + expression: ExpressionStr + + try: + SampleForm.model_validate({"expression": 'foo + col("a")'}) + except ValidationError as e: + errors = expression_form_errors(e) + else: + raise AssertionError("expected ValidationError") + + assert len(errors) == 1 + assert errors[0].path == ("expression",) + assert "foo" in errors[0].message + assert errors[0].start_line == 1 + assert errors[0].start_column == 1 + assert errors[0].end_column == 4 diff --git a/tests/test_server_integration.py b/tests/test_server_integration.py index 461de89..bcecde8 100644 --- a/tests/test_server_integration.py +++ b/tests/test_server_integration.py @@ -10,6 +10,7 @@ import uploader.app.report as report import uploader.history as history import uploader.tasks as tasks +from uploader.app.lib.formula import ExpressionStr from uploader.cli import app @@ -152,16 +153,55 @@ def cancellable_handler(form: FakeTaskForm, emit: Callable[[report.Event], None] assert fake_entry["message"] == "Task was cancelled by user." -def test_validate_expression_endpoint() -> None: +def test_validate_task_form_endpoint(isolated_task_state: None) -> None: + class FakeExpressionForm(BaseModel): + name: str + expression: ExpressionStr + + def noop_handler(form: FakeExpressionForm, emit: Callable[[report.Event], None]) -> None: + emit(report.DoneEvent(message=f"Completed {form.name}")) + + tasks.register_task( + tasks.TaskDefinition( + id="fake-expression-task", + title="Fake Expression Task", + description="Task used for form validation testing.", + form_model=FakeExpressionForm, + handler=noop_handler, + group="Tests", + ), + ) + client = TestClient(app) - ok = client.post("/api/expressions/validate", json={"expression": 'to_deg(col("RAJ2000"))'}) + + unknown = client.post("/api/tasks/does-not-exist/validate", json={}) + assert unknown.status_code == 404 + + ok = client.post( + "/api/tasks/fake-expression-task/validate", + json={"expression": 'to_deg(col("RAJ2000"))'}, + ) assert ok.status_code == 200 - assert ok.json() == {"diagnostics": []} + assert ok.json() == [] - bad = client.post("/api/expressions/validate", json={"expression": 'col("a)'}) + bad = client.post( + "/api/tasks/fake-expression-task/validate", + json={"expression": 'col("a)'}, + ) assert bad.status_code == 200 - diagnostics = bad.json()["diagnostics"] - assert len(diagnostics) == 1 - assert "string" in diagnostics[0]["message"].lower() - assert diagnostics[0]["start_line"] == 1 - assert diagnostics[0]["start_column"] >= 1 + errors = bad.json() + assert len(errors) == 1 + assert errors[0]["path"] == ["expression"] + assert "string" in errors[0]["message"].lower() + assert errors[0]["start_line"] == 1 + assert errors[0]["start_column"] >= 1 + assert "end_line" in errors[0] + assert "end_column" in errors[0] + + submit_bad = client.post( + "/api/tasks/fake-expression-task/submit", + json={"name": "alpha", "expression": 'col("a)'}, + ) + assert submit_bad.status_code == 422 + detail = submit_bad.json()["detail"] + assert any(item.get("type") == "expression" for item in detail) diff --git a/uploader/app/lib/formula/__init__.py b/uploader/app/lib/formula/__init__.py index 1ab0c6a..6cd810b 100644 --- a/uploader/app/lib/formula/__init__.py +++ b/uploader/app/lib/formula/__init__.py @@ -9,7 +9,13 @@ expression_syntax_help, expression_tokens, ) -from uploader.app.lib.formula.validate import ExpressionDiagnostic, validate_expression +from uploader.app.lib.formula.validate import ( + ExpressionDiagnostic, + ExpressionStr, + FormError, + expression_form_errors, + validate_expression, +) from uploader.app.lib.formula.values import TextValue, Value, column_quantity __all__ = [ @@ -17,11 +23,14 @@ "ExpressionDiagnostic", "ExpressionError", "ExpressionEvaluationError", + "ExpressionStr", "ExpressionSyntaxError", + "FormError", "TextValue", "Value", "column_quantity", "evaluate", + "expression_form_errors", "expression_json_schema_extra", "expression_syntax_help", "expression_tokens", diff --git a/uploader/app/lib/formula/validate.py b/uploader/app/lib/formula/validate.py index 58d8f6f..f15df99 100644 --- a/uploader/app/lib/formula/validate.py +++ b/uploader/app/lib/formula/validate.py @@ -1,6 +1,9 @@ import ast -from dataclasses import dataclass -from typing import final +from dataclasses import asdict, dataclass +from typing import Annotated, final + +from pydantic import AfterValidator, ValidationError +from pydantic_core import PydanticCustomError from uploader.app.lib.formula.namespace import COL_FUNCTION, FUNCTIONS, NAMED_CONSTANTS @@ -15,6 +18,31 @@ class ExpressionDiagnostic: end_column: int +@final +@dataclass(frozen=True) +class FormError: + path: tuple[str, ...] + message: str + start_line: int + start_column: int + end_line: int + end_column: int + + +def _validate_expression_str(source: str) -> str: + diagnostics = validate_expression(source) + if not diagnostics: + return source + raise PydanticCustomError( + "expression", + "invalid expression", + {"diagnostics": [asdict(item) for item in diagnostics]}, + ) + + +ExpressionStr = Annotated[str, AfterValidator(_validate_expression_str)] + + def _from_syntax_error(error: SyntaxError, source: str) -> ExpressionDiagnostic: start_line = error.lineno if error.lineno and error.lineno > 0 else 1 start_column = error.offset if error.offset and error.offset > 0 else 1 @@ -112,3 +140,29 @@ def validate_expression(source: str) -> list[ExpressionDiagnostic]: return [] _, _, diagnostics = diagnose_expression(source) return diagnostics + + +def expression_form_errors(error: ValidationError) -> list[FormError]: + result: list[FormError] = [] + for item in error.errors(): + if item["type"] != "expression": + continue + ctx = item.get("ctx") or {} + diagnostics = ctx.get("diagnostics") + if not isinstance(diagnostics, list): + continue + path = tuple(str(part) for part in item["loc"]) + for diagnostic in diagnostics: + if not isinstance(diagnostic, dict): + continue + result.append( + FormError( + path=path, + message=str(diagnostic["message"]), + start_line=int(diagnostic["start_line"]), + start_column=int(diagnostic["start_column"]), + end_line=int(diagnostic["end_line"]), + end_column=int(diagnostic["end_column"]), + ), + ) + return result diff --git a/uploader/cli.py b/uploader/cli.py index bf24feb..32cc3d5 100644 --- a/uploader/cli.py +++ b/uploader/cli.py @@ -10,9 +10,9 @@ from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, StreamingResponse -from pydantic import BaseModel, ValidationError +from pydantic import ValidationError -from uploader.app.lib.formula import validate_expression +from uploader.app.lib.formula import expression_form_errors from uploader.history import load_history from uploader.task_registry import register_all_tasks from uploader.tasks import TASKS, cancel_run, get_run, start_task @@ -51,15 +51,6 @@ def list_history() -> list[dict[str, object]]: return [entry.model_dump() for entry in load_history()] -class ValidateExpressionRequest(BaseModel): - expression: str - - -@app.post("/api/expressions/validate") -def validate_expression_endpoint(body: ValidateExpressionRequest) -> dict[str, object]: - return {"diagnostics": [asdict(item) for item in validate_expression(body.expression)]} - - @app.get("/api/tasks/{task_id}/schema") def task_schema(task_id: str) -> dict[str, object]: if task_id not in TASKS: @@ -75,6 +66,17 @@ def task_schema(task_id: str) -> dict[str, object]: } +@app.post("/api/tasks/{task_id}/validate") +def validate_task(task_id: str, body: dict[str, object]) -> list[dict[str, object]]: + if task_id not in TASKS: + raise HTTPException(status_code=404, detail="Unknown task") + try: + TASKS[task_id].form_model.model_validate(body) + except ValidationError as e: + return [asdict(item) for item in expression_form_errors(e)] + return [] + + @app.post("/api/tasks/{task_id}/submit") def submit_task(task_id: str, body: dict[str, object]) -> dict[str, str]: if task_id not in TASKS: diff --git a/uploader/forms/structured_catalog.py b/uploader/forms/structured_catalog.py index 4eecba4..72d1dfc 100644 --- a/uploader/forms/structured_catalog.py +++ b/uploader/forms/structured_catalog.py @@ -9,7 +9,7 @@ from uploader.app import log from uploader.app.catalogs import fetch_catalogs from uploader.app.endpoints import db_dsn_map, env_map -from uploader.app.lib.formula import expression_json_schema_extra, expression_syntax_help +from uploader.app.lib.formula import ExpressionStr, expression_json_schema_extra, expression_syntax_help from uploader.app.storage import PgStorage from uploader.app.structured.generic import upload_catalog_columns from uploader.clients.gen.client import adminapi @@ -66,12 +66,12 @@ def build_catalog_form(schema: CatalogSchema) -> type[BaseModel]: extra = expression_json_schema_extra() if _field_required(field): field_definitions[field.name] = ( - str, + ExpressionStr, Field(..., title=field.name, description=description, json_schema_extra=extra), ) else: field_definitions[field.name] = ( - str, + ExpressionStr, Field(default="", title=field.name, description=description, json_schema_extra=extra), ) field_definitions["write"] = ( diff --git a/uploader/forms/structured_designation.py b/uploader/forms/structured_designation.py index 94f879a..47d2507 100644 --- a/uploader/forms/structured_designation.py +++ b/uploader/forms/structured_designation.py @@ -7,7 +7,7 @@ import uploader.app.report as report from uploader.app.endpoints import db_dsn_map, env_map -from uploader.app.lib.formula import expression_json_schema_extra +from uploader.app.lib.formula import ExpressionStr, expression_json_schema_extra from uploader.app.storage import PgStorage from uploader.app.structured.designations import upload_designations as run_upload_designations from uploader.clients.gen.client import adminapi @@ -31,7 +31,7 @@ class StructuredDesignationAdvancedSettings(BaseModel): class StructuredDesignationForm(BaseModel): table_name: str = Field(..., title="Name of the table") - expression: str = Field( + expression: ExpressionStr = Field( ..., title="Designation expression", description=(