diff --git a/README.md b/README.md index 05002f2..04784de 100644 --- a/README.md +++ b/README.md @@ -209,7 +209,7 @@ PRs welcome! Figaro is meant to be a readable, research-friendly FL platform. **Phase 1: Solidifying the Agentic Platform** - [x] **Interactive Agent Planning** — Multi-turn dialogue support for refining experiments, plus visual topology previews (Plan Preview) before execution. - [x] **Execution Transparency** — Real-time tracking of node-level status during execution and automated natural-language interpretation of results. -- [ ] **Strict Configuration Engine** — Implement strict Pydantic/JSON Schema validation to resolve historical inconsistencies between `config_schema` and underlying algorithms. +- [x] **Strict Configuration Engine** — Implement strict Pydantic/JSON Schema validation to resolve historical inconsistencies between `config_schema` and underlying algorithms. - [ ] **Advanced Experiment Tracking** — Multi-dimensional search filtering (by metrics, hyperparameters, status) and configuration version control (diffing). **Phase 2: LLM & LoRA Federated Fine-Tuning** diff --git a/README.zh-CN.md b/README.zh-CN.md index 64e4396..67c527d 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -210,7 +210,7 @@ figaro/ **Phase 1:夯实 Agentic 实验平台** - [x] **Agent 交互体验升级** —— 支持多轮对话微调实验计划,提供实验执行前的Plan Preview。 - [x] **执行与分析透明化** —— 支持实验节点的实时状态追踪,以及 Agent 驱动的运行结果自动化图表解释。 -- [ ] **配置引擎重构** —— 引入基于 Pydantic/JSON Schema 的严格强校验,彻底修复 `config_schema` 与底层算法实现不一致的问题。 +- [x] **配置引擎重构** —— 引入基于 Pydantic/JSON Schema 的严格强校验,彻底修复 `config_schema` 与底层算法实现不一致的问题。 - [ ] **高阶实验管理** —— 支持按指标、超参等多维度搜索过滤实验历史,支持配置文件版本控制与 Diff 差异对比。 **Phase 2:LLM / LoRA 联邦微调支持** diff --git a/apps/backend/app/api/v1/endpoints/agent.py b/apps/backend/app/api/v1/endpoints/agent.py index 63c5888..29f6f1b 100644 --- a/apps/backend/app/api/v1/endpoints/agent.py +++ b/apps/backend/app/api/v1/endpoints/agent.py @@ -27,6 +27,7 @@ from app.services.agent import AgentExperimentRunService, AgentExperimentService, AgentOptimizationHistoryService, AgentRuntimeService, AgentState from app.services.agent.graph import FederatedAgentGraphBuilder from app.services.agent.objectives import AgentOptimizationObjective, resolve_objective +from app.services.agent.planning import build_schema_prompt_context, dumps_for_prompt from app.services.llm import LLMRegistry, LLMService from app.schemas.agent import AgentPlanPreviewRequest, AgentPlanPreviewResponse, ExperimentPlanPreview, AgentPlanReviseRequest import uuid @@ -49,6 +50,17 @@ async def list_agent_models() -> AgentModelsResponse: ) +@agent_router.get( + "/config/schema", + response_model=dict, + status_code=status.HTTP_200_OK, + summary="Get Agent config schema", +) +async def get_agent_config_schema() -> dict: + """Return the schema used by Agent planning and editing UI.""" + return AgentExperimentService.get_config_schema() + + def _build_experiments(history: list) -> list[AgentExperimentSummary]: """ Convert internal history snapshots into API experiment summaries. @@ -87,13 +99,55 @@ def _build_experiments(history: list) -> list[AgentExperimentSummary]: return experiments +def _record_value(record, key: str, default=None): + if hasattr(record, key): + return getattr(record, key) + if isinstance(record, dict): + return record.get(key, default) + return default + + +def _best_record(history: list): + best = None + best_score = None + for record in history: + score = _record_value(record, "score") + if score is None: + continue + try: + numeric_score = float(score) + except (TypeError, ValueError): + continue + if best_score is None or numeric_score > best_score: + best = record + best_score = numeric_score + return best + + +def _derive_best_payload(snapshot: dict) -> tuple[dict | None, dict | None]: + best_config = snapshot.get("best_config") + best_metrics = snapshot.get("best_metrics") + if best_config is not None and best_metrics is not None: + return best_config, best_metrics + + record = _best_record(snapshot.get("experiments", [])) + if record is None: + return best_config, best_metrics + if best_config is None: + best_config = _record_value(record, "config") + if best_metrics is None: + best_metrics = _record_value(record, "metrics") + return best_config, best_metrics + + def _build_progress_response(snapshot: dict) -> AgentOptimizeProgressResponse: """ Convert a runtime snapshot into the progress response schema. """ + best_config_payload, best_metrics_payload = _derive_best_payload(snapshot) best_metrics = ( - SimulationRunMetricsResponse.model_validate(snapshot["best_metrics"]) - if snapshot.get("best_metrics") is not None + SimulationRunMetricsResponse.model_validate(best_metrics_payload) + if best_metrics_payload is not None else None ) @@ -143,10 +197,11 @@ def _build_progress_response(snapshot: dict) -> AgentOptimizeProgressResponse: completed_iterations=snapshot.get("completed_iterations", 0), current_plan=current_plan, current_experiment=current_experiment, - best_config=snapshot.get("best_config"), + best_config=best_config_payload, best_metrics=best_metrics, experiments=_build_experiments(snapshot.get("experiments", [])), draft_experiments=snapshot.get("draft_experiments", []), + config_constraints=snapshot.get("config_constraints", {}), summary_text=snapshot.get("summary_text"), error_message=snapshot.get("error_message"), created_at=snapshot.get("created_at"), @@ -222,6 +277,7 @@ async def get_optimization_job( snapshot.setdefault("current_iteration", job.current_iteration) snapshot.setdefault("completed_iterations", job.completed_iterations) snapshot.setdefault("experiments", []) + snapshot.setdefault("config_constraints", {}) snapshot.setdefault("created_at", job.created_at) snapshot.setdefault("updated_at", job.updated_at) snapshot.setdefault("finished_at", job.finished_at) @@ -248,6 +304,7 @@ async def start_optimization( model_name=payload.model_name, objective=payload.objective, planned_experiments=payload.planned_experiments, + config_constraints=payload.config_constraints, ) return _build_progress_response(snapshot) @@ -304,6 +361,7 @@ async def optimize( job_name=payload.job_name, objective=payload.objective, resolved_objective=resolved_objective, + config_constraints=payload.config_constraints, ) # LangGraph executor expects a dict-like object; dataclass is fine. @@ -318,6 +376,12 @@ async def optimize( final_state = AgentState(**raw_state) # type: ignore[arg-type] experiments = _build_experiments(final_state.history) + best = _best_record(final_state.history) + best_metrics = ( + SimulationRunMetricsResponse.model_validate(_record_value(best, "metrics")) + if best is not None and _record_value(best, "metrics") is not None + else None + ) return AgentOptimizeResponse( goal=final_state.goal, @@ -326,8 +390,8 @@ async def optimize( objective=final_state.objective, resolved_objective=final_state.resolved_objective, iterations_executed=final_state.iteration, - best_config=None, - best_metrics=None, + best_config=_record_value(best, "config") if best is not None else None, + best_metrics=best_metrics, experiments=experiments, summary_text=final_state.summary, ) @@ -497,6 +561,7 @@ async def generate_plan_preview( max_iterations=10, objective=AgentOptimizationObjective.AUTO, resolved_objective=AgentOptimizationObjective.ACCURACY, + config_constraints=payload.config_constraints, ) # Call the graph's parse node directly to get the LLM result @@ -528,6 +593,9 @@ async def generate_plan_preview( "goal": payload.goal, "job_name": payload.job_name, "status": "pending_review", + "system_mode": payload.system_mode, + "model_name": payload.model_name, + "config_constraints": payload.config_constraints, "experiments": [], "draft_experiments": [p.model_dump() for p in previews] } @@ -537,7 +605,8 @@ async def generate_plan_preview( optimization_job_id=job.id, goal=payload.goal, experiments=previews, - system_mode=payload.system_mode + system_mode=payload.system_mode, + config_constraints=payload.config_constraints, ) @agent_router.post( @@ -560,6 +629,7 @@ async def revise_plan_preview( snapshot = job.snapshot_json current_experiments = snapshot.get("draft_experiments", []) + config_constraints = snapshot.get("config_constraints", {}) if not current_experiments: raise HTTPException(status_code=400, detail="No draft experiments found to revise.") @@ -567,12 +637,16 @@ async def revise_plan_preview( "You are an AI assistant managing a Federated Learning experiment plan. " "The user wants to modify the current experiment configuration based on their feedback. " "Update the JSON configuration to reflect their request. " + "Respect the structured config constraints and do not use disabled future options. " "Respond ONLY with a valid JSON array of experiment objects, matching the original schema. " "Do not include any other text." ) + schema_context = build_schema_prompt_context(AgentExperimentService.get_config_schema()) input_text = ( f"Current Plan (JSON):\n{json.dumps(current_experiments, indent=2)}\n\n" + f"Structured Config Constraints (JSON):\n{json.dumps(config_constraints, indent=2)}\n\n" + f"Schema Context (JSON):\n{dumps_for_prompt(schema_context)}\n\n" f"User Modification Request:\n{payload.instruction}\n\n" "Please output the updated JSON array of experiments:" ) @@ -608,5 +682,6 @@ async def revise_plan_preview( optimization_job_id=job.id, goal=snapshot.get("goal", ""), experiments=[ExperimentPlanPreview.model_validate(exp) for exp in updated_experiments], - system_mode=snapshot.get("system_mode", "simulation") - ) \ No newline at end of file + system_mode=snapshot.get("system_mode", "simulation"), + config_constraints=config_constraints, + ) diff --git a/apps/backend/app/schemas/agent.py b/apps/backend/app/schemas/agent.py index f11614f..3baec67 100644 --- a/apps/backend/app/schemas/agent.py +++ b/apps/backend/app/schemas/agent.py @@ -45,6 +45,10 @@ class AgentOptimizeRequest(BaseModel): default=None, description="Explicit list of experiment patches to run, bypassing LLM parse." ) + config_constraints: dict[str, Any] = Field( + default_factory=dict, + description="Structured config constraints selected by the user before planning.", + ) class AgentModelsResponse(BaseModel): @@ -148,6 +152,7 @@ class AgentOptimizeProgressResponse(BaseModel): best_metrics: SimulationRunMetricsResponse | None = None experiments: list[AgentExperimentSummary] = Field(default_factory=list) draft_experiments: list[dict[str, Any]] = Field(default_factory=list) + config_constraints: dict[str, Any] = Field(default_factory=dict) summary_text: str | None = None error_message: str | None = None created_at: datetime | None = None @@ -237,6 +242,7 @@ class AgentPlanPreviewRequest(BaseModel): job_name: str = Field(..., description="Name for the job group.") model_name: Optional[str] = None system_mode: str = "simulation" + config_constraints: dict[str, Any] = Field(default_factory=dict) class AgentPlanPreviewResponse(BaseModel): """Draft results returned to the frontend.""" @@ -244,7 +250,8 @@ class AgentPlanPreviewResponse(BaseModel): goal: str experiments: list[ExperimentPlanPreview] system_mode: str + config_constraints: dict[str, Any] = Field(default_factory=dict) class AgentPlanReviseRequest(BaseModel): """Payload for requesting a revision to an existing plan draft.""" - instruction: str = Field(..., description="Natural language feedback to revise the plan.") \ No newline at end of file + instruction: str = Field(..., description="Natural language feedback to revise the plan.") diff --git a/apps/backend/app/services/agent/experiment_service.py b/apps/backend/app/services/agent/experiment_service.py index 601e932..a6bbd8d 100644 --- a/apps/backend/app/services/agent/experiment_service.py +++ b/apps/backend/app/services/agent/experiment_service.py @@ -70,7 +70,7 @@ class AgentExperimentService: """Application service for agent experiment operations.""" EXPERIMENT_NAME_CONFLICT_MESSAGE = "Experiment name already exists" - _SCHEMA_META_KEYS = {"role", "depends_on", "hidden"} + _SCHEMA_META_KEYS = {"role", "depends_on", "hidden", "ui"} def __init__(self, session: AsyncSession): self.session = session diff --git a/apps/backend/app/services/agent/graph.py b/apps/backend/app/services/agent/graph.py index cdaabd8..f9d305f 100644 --- a/apps/backend/app/services/agent/graph.py +++ b/apps/backend/app/services/agent/graph.py @@ -26,7 +26,13 @@ from app.services.llm import LLMService from .capabilities import get_platform_capabilities -from .planning import build_initial_config, select_llm_model +from .planning import ( + build_initial_config, + build_schema_prompt_context, + collect_disabled_option_errors, + deep_merge_config, + select_llm_model, +) from .prompts import build_plan_prompt, build_plan_system_instructions from .state import AgentState, ExperimentPlan, ExperimentRecord from .summary import build_results_table, build_summary_text, get_last_global_accuracy @@ -74,17 +80,36 @@ async def _node_parse(self, state: AgentState) -> AgentState: state.phase = "parsing" await self._publish_progress(state) + experiment_service = AgentExperimentService(self._session) + schema = experiment_service.get_config_schema() + base_config = build_initial_config(schema) + constrained_base = experiment_service.normalize_simulation_config( + deep_merge_config(base_config, state.config_constraints) + ) + disabled_errors = collect_disabled_option_errors(constrained_base, schema) + if disabled_errors: + raise exceptions.BadRequestError("; ".join(disabled_errors)) + if getattr(state, "planned_experiments", None): experiments: list[ExperimentPlan] = [] for idx, exp in enumerate(state.planned_experiments): name = exp.get("name", f"exp-{idx + 1}") + raw_config = exp.get("config_patch", exp.get("config", {})) + if not isinstance(raw_config, dict): + raw_config = {} + normalized = experiment_service.normalize_simulation_config( + deep_merge_config(constrained_base, raw_config) + ) + disabled_errors = collect_disabled_option_errors(normalized, schema) + if disabled_errors: + raise exceptions.BadRequestError("; ".join(disabled_errors)) experiments.append( ExperimentPlan( iteration=idx + 1, iteration_goal=f"Run experiment: {name}", name=name, plan_summary=exp.get("plan_summary", ""), - config_patch=exp.get("config_patch", {}), + config_patch=normalized, ) ) state.experiments = experiments @@ -93,16 +118,16 @@ async def _node_parse(self, state: AgentState) -> AgentState: await self._publish_progress(state) return state - experiment_service = AgentExperimentService(self._session) - schema = experiment_service.get_config_schema() capabilities = get_platform_capabilities() - base_config = build_initial_config(schema) + schema_context = build_schema_prompt_context(schema) instructions = build_plan_system_instructions(state.goal) prompt = build_plan_prompt( state=state, capabilities=capabilities, - base_config=base_config, + base_config=constrained_base, + schema_context=schema_context, + config_constraints=state.config_constraints, ) logger.info("llm input instructions=%s", instructions) @@ -140,15 +165,17 @@ async def _node_parse(self, state: AgentState) -> AgentState: patch_config = exp.get("config", {}) if not isinstance(patch_config, dict): patch_config = {} - # Merge with base config and normalize - current_full_config = copy.deepcopy(base_config) - merged = self._deep_merge_dicts(current_full_config, patch_config) + # Merge with constrained base config and normalize. + merged = deep_merge_config(constrained_base, patch_config) try: normalized = experiment_service.normalize_simulation_config(merged) + disabled_errors = collect_disabled_option_errors(normalized, schema) + if disabled_errors: + raise exceptions.BadRequestError("; ".join(disabled_errors)) except Exception: # 如果合并出错,至少保证能跑,回退到基础配置 normalized = experiment_service.normalize_simulation_config( - copy.deepcopy(base_config) + copy.deepcopy(constrained_base) ) experiments.append( ExperimentPlan( @@ -165,7 +192,7 @@ async def _node_parse(self, state: AgentState) -> AgentState: # Fallback: if no experiments were parsed, run a single default if not experiments: normalized_base = experiment_service.normalize_simulation_config( - copy.deepcopy(base_config) + copy.deepcopy(constrained_base) ) experiments.append( ExperimentPlan( @@ -173,7 +200,7 @@ async def _node_parse(self, state: AgentState) -> AgentState: iteration_goal="Run default experiment (LLM parse failed)", name="default", plan_summary="Fallback: running single default configuration.", - config_patch={}, + config_patch=normalized_base, ) ) @@ -192,7 +219,9 @@ async def _node_run_sequential(self, state: AgentState) -> AgentState: experiment_service = AgentExperimentService(self._session) run_service = AgentExperimentRunService(self._session) schema = experiment_service.get_config_schema() - base_config = build_initial_config(schema) + base_config = experiment_service.normalize_simulation_config( + deep_merge_config(build_initial_config(schema), state.config_constraints) + ) max_wait_seconds = 3600 poll_interval = 2 @@ -203,13 +232,11 @@ async def _node_run_sequential(self, state: AgentState) -> AgentState: state.iteration = idx + 1 # -- Build config -- - merged = self._deep_merge_dicts(base_config, plan.config_patch) - try: - config = experiment_service.normalize_simulation_config(merged) - except exceptions.BadRequestError: - config = experiment_service.normalize_simulation_config( - copy.deepcopy(base_config) - ) + merged = deep_merge_config(base_config, plan.config_patch) + config = experiment_service.normalize_simulation_config(merged) + disabled_errors = collect_disabled_option_errors(config, schema) + if disabled_errors: + raise exceptions.BadRequestError("; ".join(disabled_errors)) state.current_config = config # -- Launch -- diff --git a/apps/backend/app/services/agent/planning.py b/apps/backend/app/services/agent/planning.py index 8ce3f89..c66570c 100644 --- a/apps/backend/app/services/agent/planning.py +++ b/apps/backend/app/services/agent/planning.py @@ -4,32 +4,181 @@ from __future__ import annotations +import copy +import json from typing import Any from app.core.config import settings +SCHEMA_META_KEYS = {"role", "depends_on", "hidden", "ui"} + + +def _is_field_definition(node: Any) -> bool: + if not isinstance(node, dict): + return False + if isinstance(node.get("type"), str): + return True + return "default" in node and not any( + key not in SCHEMA_META_KEYS and isinstance(value, dict) + for key, value in node.items() + ) + + +def _default_for_field(node: dict[str, Any]) -> Any: + if "default" in node: + return copy.deepcopy(node["default"]) + field_type = node.get("type") + if field_type == "bool": + return False + if field_type == "number": + return 0 + if field_type == "list_int": + return [] + if field_type == "select": + options = node.get("options") + if isinstance(options, list) and options: + return copy.deepcopy(options[0]) + return "" + return "" + def build_initial_config(schema: dict[str, Any]) -> dict[str, Any]: """ Build a default FL experiment config from the simulation schema. """ - federated = schema.get("federated") or {} - dataset = schema.get("dataset") or {} - model = schema.get("model") or {} + if not schema: + return { + "dataset": {"name": "CIFAR-10", "distribution": "non_iid", "alpha": 0.5}, + "model": {"name": "CNN"}, + "federated": { + "num_clients": 3, + "num_rounds": 10, + "clients_per_round": 2, + "local_epochs": 5, + "learning_rate": 0.01, + "aggregation": "fedavg", + "seed": 42, + }, + } + + output: dict[str, Any] = {} + for key, definition in schema.items(): + if key in SCHEMA_META_KEYS or not isinstance(definition, dict): + continue + if _is_field_definition(definition): + output[key] = _default_for_field(definition) + continue + output[key] = build_initial_config(definition) + return output + + +def deep_merge_config(base: dict[str, Any], override: dict[str, Any] | None) -> dict[str, Any]: + """Deep-merge config dictionaries while ignoring schema/internal metadata keys.""" + merged = copy.deepcopy(base) + if not isinstance(override, dict): + return merged + for key, value in override.items(): + if key.startswith("_"): + continue + base_value = merged.get(key) + if isinstance(base_value, dict) and isinstance(value, dict): + merged[key] = deep_merge_config(base_value, value) + else: + merged[key] = copy.deepcopy(value) + return merged + + +def _option_ui(definition: dict[str, Any]) -> dict[str, Any]: + ui = definition.get("ui") + if not isinstance(ui, dict): + return {} + option_ui = ui.get("options") + return option_ui if isinstance(option_ui, dict) else {} + + +def _is_disabled_option(definition: dict[str, Any], value: Any) -> bool: + metadata = _option_ui(definition).get(str(value)) + return isinstance(metadata, dict) and metadata.get("disabled") is True + + +def collect_disabled_option_errors( + config: dict[str, Any], + schema: dict[str, Any], + *, + path_prefix: str = "", +) -> list[str]: + """Return user-facing errors for disabled schema select options in config.""" + errors: list[str] = [] + for key, definition in schema.items(): + if key in SCHEMA_META_KEYS or not isinstance(definition, dict): + continue + full_path = f"{path_prefix}.{key}" if path_prefix else key + if _is_field_definition(definition): + if definition.get("type") == "select": + value = config.get(key) + if value is not None and _is_disabled_option(definition, value): + errors.append(f"{full_path}={value!r} is marked disabled in config_schema.yaml") + continue + child_config = config.get(key) + if isinstance(child_config, dict): + errors.extend(collect_disabled_option_errors(child_config, definition, path_prefix=full_path)) + return errors + + +def build_schema_prompt_context(schema: dict[str, Any]) -> dict[str, Any]: + """Build compact schema context for LLM planning prompts.""" + fields: list[dict[str, Any]] = [] + + def walk(node: dict[str, Any], prefix: str = "") -> None: + for key, definition in node.items(): + if key in SCHEMA_META_KEYS or not isinstance(definition, dict): + continue + path = f"{prefix}.{key}" if prefix else key + if _is_field_definition(definition): + ui = definition.get("ui") if isinstance(definition.get("ui"), dict) else {} + options = definition.get("options") if isinstance(definition.get("options"), list) else None + disabled_options = [ + option + for option in (options or []) + if _is_disabled_option(definition, option) + ] + executable_options = [ + option + for option in (options or []) + if option not in disabled_options + ] + field: dict[str, Any] = { + "path": path, + "type": definition.get("type"), + "default": _default_for_field(definition), + } + if options is not None: + field["options"] = options + field["executable_options"] = executable_options + field["disabled_options"] = disabled_options + if isinstance(ui, dict): + if ui.get("label"): + field["label"] = ui["label"] + if ui.get("prompt_hint"): + field["prompt_hint"] = ui["prompt_hint"] + if ui.get("featured") is True: + field["featured"] = True + fields.append(field) + continue + walk(definition, path) + + walk(schema) return { - "dataset": {"name": dataset.get("name", {}).get("default", "cifar10")}, - "model": {"name": model.get("name", {}).get("default", "cnn")}, - "federated": { - "num_clients": federated.get("num_clients", {}).get("default", 10), - "num_rounds": federated.get("num_rounds", {}).get("default", 20), - "clients_per_round": federated.get("clients_per_round", {}).get("default", 5), - "local_epochs": federated.get("local_epochs", {}).get("default", 5), - "learning_rate": federated.get("learning_rate", {}).get("default", 0.01), - "seed": federated.get("seed", {}).get("default", 42), - }, + "defaults": build_initial_config(schema), + "fields": fields, } +def dumps_for_prompt(value: Any) -> str: + """Stable JSON rendering for LLM prompt context.""" + return json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + + def select_llm_model(model_name: str | None) -> str: """Resolve the LLM model for planning requests.""" if model_name and model_name.strip(): diff --git a/apps/backend/app/services/agent/prompts/__init__.py b/apps/backend/app/services/agent/prompts/__init__.py index 0145912..73a3276 100644 --- a/apps/backend/app/services/agent/prompts/__init__.py +++ b/apps/backend/app/services/agent/prompts/__init__.py @@ -7,6 +7,7 @@ from pathlib import Path from typing import Any +from app.services.agent.planning import dumps_for_prompt from app.services.agent.state import AgentState _PLAN_INSTRUCTIONS_PATH = Path(__file__).with_name("plan_instructions.txt") @@ -33,13 +34,21 @@ def build_plan_prompt( state: AgentState, capabilities: dict[str, Any], base_config: dict[str, Any], + schema_context: dict[str, Any] | None = None, + config_constraints: dict[str, Any] | None = None, ) -> str: """Build the user prompt for the experiment-parsing LLM call.""" + constraints = config_constraints if isinstance(config_constraints, dict) else {} + schema_payload = schema_context if isinstance(schema_context, dict) else {} return ( f"Experiment request:\n{state.goal}\n\n" f"System mode: {state.system_mode}\n" - f"Capabilities: {capabilities}\n" - f"Default base config: {base_config}\n\n" + f"Runtime capabilities:\n{dumps_for_prompt(capabilities)}\n\n" + f"Schema context from config_schema.yaml:\n{dumps_for_prompt(schema_payload)}\n\n" + f"Default base config after applying user constraints:\n{dumps_for_prompt(base_config)}\n\n" + f"User-selected structured constraints:\n{dumps_for_prompt(constraints)}\n\n" "Parse the request into experiment configurations.\n" + "Every experiment must inherit the structured constraints unless the user's request explicitly changes them.\n" + "Use executable_options for select fields. Do not use disabled_options.\n" "Return ONLY a JSON object with plan_summary and experiments list, no commentary." ) diff --git a/apps/backend/app/services/agent/prompts/plan_instructions.txt b/apps/backend/app/services/agent/prompts/plan_instructions.txt index 954e736..316aaa1 100644 --- a/apps/backend/app/services/agent/prompts/plan_instructions.txt +++ b/apps/backend/app/services/agent/prompts/plan_instructions.txt @@ -1,43 +1,19 @@ You are a federated learning experiment planner. -Given a user's natural language experiment request, parse it into a list of experiment configurations. +Given a user's natural-language experiment request, parse it into a list of experiment configurations. Rules: -- If the user mentions multiple parameter values (e.g. "test alpha=0.1, 0.3, 0.5"), generate one experiment per value -- If the user mentions multiple dimensions (e.g. "test alpha=0.1,0.5 with 10,20 rounds"), generate the Cartesian product -- If the user doesn't specify something, use the defaults shown in the example below -- Each experiment's "config" must be a PARTIAL config dict containing ONLY the fields that differ from defaults or are explicitly requested -- `federated.seed` controls random seed for full reproducibility (model init, data split, client selection, training order). Use it as follows: - - DEFAULT: do NOT set seed unless the user explicitly asks. All experiments will share the same default seed, so runs that differ only in one swept parameter will have identical starting state, giving a clean comparison. - - If the user asks to repeat the same experiment multiple times / average over N runs / vary the seed / measure variance: generate N experiments with distinct seeds (e.g. 1, 2, 3). - - If the user asks to test specific seed values (e.g. "seed=1,2,3"): use those values directly. +- Treat config_schema.yaml as the source of truth for fields, defaults, selectable options, and disabled future options. +- If the user mentions multiple parameter values, generate one experiment per value. +- If the user mentions multiple dimensions, generate the Cartesian product. +- If the user does not specify a field, inherit the default base config and structured constraints from the prompt. +- Each experiment's "config" should contain the actual parameter values that make that experiment distinct. +- Do not use any option listed under disabled_options. +- `federated.seed` controls random seed for full reproducibility. Do not vary it unless the user asks for repeats, variance, or specific seeds. - Return ONLY a JSON object with: - "plan_summary": brief description of what will be tested - "experiments": list of objects, each with "name" (short identifier) and "config" (partial config dict) -Config structure (defaults shown — only include fields you want to change): -{ - "dataset": { - "name": "CIFAR-10", - "distribution": "non_iid", - "alpha": 0.5 - }, - "model": { - "name": "CNN", - "input_shape": [32, 32, 3], - "num_classes": 10 - }, - "federated": { - "num_clients": 10, - "num_rounds": 20, - "clients_per_round": 5, - "local_epochs": 5, - "learning_rate": 0.01, - "aggregation": "fedavg", - "seed": 42 - } -} - -Example — user says "test alpha=0.1, 0.3, 0.5": +Example: user says "test alpha=0.1, 0.3, 0.5": { "plan_summary": "Compare non-IID alpha values 0.1, 0.3, 0.5 on CIFAR-10", "experiments": [ @@ -47,25 +23,12 @@ Example — user says "test alpha=0.1, 0.3, 0.5": ] } -Example — user says "run the same experiment 3 times to measure variance": +Example: user says "run the same experiment 3 times to measure variance": { - "plan_summary": "Repeat default configuration 3 times with different seeds", + "plan_summary": "Repeat the selected configuration 3 times with different seeds", "experiments": [ {"name": "seed_1", "config": {"federated": {"seed": 1}}}, {"name": "seed_2", "config": {"federated": {"seed": 2}}}, {"name": "seed_3", "config": {"federated": {"seed": 3}}} ] } - -Example — user says "compare 10 and 20 clients with 10 and 20 rounds": -{ - "plan_summary": "Cartesian product: clients={10,20} x rounds={10,20}", - "experiments": [ - {"name": "10c_10r", "config": {"federated": {"num_clients": 10, "num_rounds": 10, "clients_per_round": 5}}}, - {"name": "10c_20r", "config": {"federated": {"num_clients": 10, "num_rounds": 20, "clients_per_round": 5}}}, - {"name": "20c_10r", "config": {"federated": {"num_clients": 20, "num_rounds": 10, "clients_per_round": 10}}}, - {"name": "20c_20r", "config": {"federated": {"num_clients": 20, "num_rounds": 20, "clients_per_round": 10}}} - ] -} - -IMPORTANT: Each experiment's "config" MUST contain the actual parameter values. Do NOT just put the values in the name — put them in the config dict. diff --git a/apps/backend/app/services/agent/runtime_service.py b/apps/backend/app/services/agent/runtime_service.py index 5a3335b..dca27e2 100644 --- a/apps/backend/app/services/agent/runtime_service.py +++ b/apps/backend/app/services/agent/runtime_service.py @@ -40,6 +40,7 @@ async def start_optimization( job_name: str | None, objective: AgentOptimizationObjective, planned_experiments: list[dict[str, Any]], + config_constraints: dict[str, Any] | None = None, ) -> dict[str, Any]: """ Create a background experiment task and return its initial snapshot. @@ -67,6 +68,7 @@ async def start_optimization( "best_config": None, "best_metrics": None, "experiments": [], + "config_constraints": copy.deepcopy(config_constraints or {}), "summary_text": None, "error_message": None, "created_at": now, @@ -100,6 +102,7 @@ async def start_optimization( objective=objective, resolved_objective=resolved_objective, planned_experiments=planned_experiments, + config_constraints=config_constraints or {}, ) ) task.add_done_callback(lambda finished_task, current_task_id=task_id: self._on_task_done(current_task_id, finished_task)) @@ -127,6 +130,7 @@ async def _run_task( objective: AgentOptimizationObjective, resolved_objective: AgentOptimizationObjective, planned_experiments: list[dict[str, Any]] | None = None, + config_constraints: dict[str, Any] | None = None, ) -> None: """Execute one experiment task in the background.""" logger.info("agent_task_started task_id=%s", task_id) @@ -142,6 +146,7 @@ async def _run_task( resolved_objective=resolved_objective, phase="parsing", planned_experiments=planned_experiments, + config_constraints=copy.deepcopy(config_constraints or {}), ) await self._update_from_state(task_id, initial_state, status="running") async with AsyncSessionLocal() as session: @@ -238,6 +243,7 @@ async def _mark_task_failed_if_stale(self, *, task_id: str, message: str) -> Non async def _update_from_state(self, task_id: str, state: AgentState, *, status: str) -> None: """Project AgentState into a serializable task snapshot.""" latest_record = state.history[-1] if state.history else None + best_record = self._select_best_record(state.history) current_experiment = self._build_current_experiment(state, latest_record) snapshot_to_persist = None @@ -260,9 +266,10 @@ async def _update_from_state(self, task_id: str, state: AgentState, *, status: s "completed_iterations": len(state.history), "current_plan": self._serialize_plan(state.current_plan), "current_experiment": current_experiment, - "best_config": None, - "best_metrics": None, + "best_config": copy.deepcopy(best_record.config) if best_record is not None else None, + "best_metrics": copy.deepcopy(best_record.metrics) if best_record is not None else None, "experiments": [self._serialize_record(record) for record in state.history], + "config_constraints": copy.deepcopy(state.config_constraints), "summary_text": state.summary, "error_message": state.error_message, "updated_at": utcnow(), @@ -274,6 +281,14 @@ async def _update_from_state(self, task_id: str, state: AgentState, *, status: s if snapshot_to_persist is not None: await self._persist_snapshot(snapshot_to_persist) + @staticmethod + def _select_best_record(records: list[ExperimentRecord]) -> ExperimentRecord | None: + """Return the highest-scoring experiment record, if any.""" + scored = [record for record in records if record.score is not None] + if not scored: + return None + return max(scored, key=lambda record: record.score or 0) + @staticmethod def _build_current_experiment( state: AgentState, diff --git a/apps/backend/app/services/agent/state.py b/apps/backend/app/services/agent/state.py index 911274b..6dc1e99 100644 --- a/apps/backend/app/services/agent/state.py +++ b/apps/backend/app/services/agent/state.py @@ -104,3 +104,4 @@ class AgentState: summary: str | None = None planned_experiments: list[dict[str, Any]] | None = None + config_constraints: dict[str, Any] = field(default_factory=dict) diff --git a/apps/backend/app/services/simulation/job_service.py b/apps/backend/app/services/simulation/job_service.py index a927632..59b4386 100644 --- a/apps/backend/app/services/simulation/job_service.py +++ b/apps/backend/app/services/simulation/job_service.py @@ -21,7 +21,7 @@ class SimulationJobService: """Application service for simulation job operations.""" JOB_NAME_CONFLICT_MESSAGE = "Job name already exists" - _SCHEMA_META_KEYS = {"role", "depends_on", "hidden"} + _SCHEMA_META_KEYS = {"role", "depends_on", "hidden", "ui"} def __init__(self, session: AsyncSession): diff --git a/apps/backend/tests/api/v1/endpoints/test_agent.py b/apps/backend/tests/api/v1/endpoints/test_agent.py index cefb45a..44d5b2b 100644 --- a/apps/backend/tests/api/v1/endpoints/test_agent.py +++ b/apps/backend/tests/api/v1/endpoints/test_agent.py @@ -17,6 +17,18 @@ def _metrics_payload(accuracy: float) -> dict: } +def test_agent_config_schema_endpoint_returns_ui_metadata(client): + response = client.get("/api/v1/agent/config/schema") + + assert response.status_code == 200 + payload = response.json() + assert payload["model"]["name"]["options"] == ["CNN", "LeNet", "ResNet"] + aggregation_ui = payload["federated"]["aggregation"]["ui"] + assert aggregation_ui["featured"] is True + assert aggregation_ui["options"]["fedprox"]["disabled"] is True + assert aggregation_ui["options"]["scaffold"]["badge"] == "experimental" + + def test_agent_optimize_endpoint_handles_agent_state_result(client, monkeypatch): import app.api.v1.endpoints.agent as agent_module @@ -83,6 +95,8 @@ def build(): assert payload["job_name"] == "opt-job-a" assert payload["iterations_executed"] == 2 assert payload["experiments"][0]["run_id"] == "run-1" + assert payload["best_config"] == {"federated": {"num_clients": 10}} + assert payload["best_metrics"]["global_results"]["global_accuracy"] == [0.8] assert payload["summary_text"] == "done" @@ -149,6 +163,8 @@ def build(): assert payload["resolved_objective"] == "accuracy" assert payload["iterations_executed"] == 1 assert payload["experiments"][0]["job_id"] == 11 + assert payload["best_config"] == {"federated": {"num_clients": 10}} + assert payload["best_metrics"]["global_results"]["global_accuracy"] == [0.88] assert payload["summary_text"] == "single-run" @@ -156,13 +172,14 @@ def test_agent_optimize_start_endpoint_returns_live_progress(client, monkeypatch import app.api.v1.endpoints.agent as agent_module # ADD `planned_experiments` here 👇 - async def _fake_start_optimization(*, goal, max_iterations, system_mode, model_name, job_name, objective, planned_experiments): + async def _fake_start_optimization(*, goal, max_iterations, system_mode, model_name, job_name, objective, planned_experiments, config_constraints): assert goal == "live optimize" assert max_iterations == 3 assert system_mode == "simulation" assert model_name == "gpt-live" assert job_name == "opt-job-live" assert planned_experiments is None # Optional: verify the default value is passed + assert config_constraints == {} return { "task_id": "task-1", "status": "running", @@ -296,6 +313,8 @@ async def _fake_get_task(task_id): assert payload["resolved_objective"] == "accuracy" assert payload["completed_iterations"] == 2 assert payload["experiments"][0]["job_id"] == 99 + assert payload["best_config"] == {"federated": {"num_rounds": 20}} + assert payload["best_metrics"]["global_results"]["global_accuracy"] == [0.95] def test_agent_optimization_jobs_history_endpoints(client, monkeypatch): diff --git a/apps/backend/tests/services/agent/test_planning.py b/apps/backend/tests/services/agent/test_planning.py index 252a945..0fa66db 100644 --- a/apps/backend/tests/services/agent/test_planning.py +++ b/apps/backend/tests/services/agent/test_planning.py @@ -1,5 +1,10 @@ from app.core.config import settings -from app.services.agent.planning import build_initial_config, select_llm_model +from app.services.agent.planning import ( + build_initial_config, + build_schema_prompt_context, + collect_disabled_option_errors, + select_llm_model, +) def test_select_llm_model_prefers_model_name(monkeypatch): @@ -38,7 +43,29 @@ def test_build_initial_config_uses_schema_defaults(): def test_build_initial_config_hardcoded_defaults_on_empty_schema(): cfg = build_initial_config({}) - assert cfg["dataset"]["name"] == "cifar10" - assert cfg["model"]["name"] == "cnn" - assert cfg["federated"]["num_clients"] == 10 - assert cfg["federated"]["num_rounds"] == 20 + assert cfg["dataset"]["name"] == "CIFAR-10" + assert cfg["model"]["name"] == "CNN" + assert cfg["federated"]["num_clients"] == 3 + assert cfg["federated"]["num_rounds"] == 10 + + +def test_schema_prompt_context_marks_disabled_options(): + schema = { + "federated": { + "aggregation": { + "type": "select", + "options": ["fedavg", "scaffold"], + "default": "fedavg", + "ui": {"options": {"scaffold": {"disabled": True}}}, + } + } + } + + context = build_schema_prompt_context(schema) + field = context["fields"][0] + assert field["path"] == "federated.aggregation" + assert field["executable_options"] == ["fedavg"] + assert field["disabled_options"] == ["scaffold"] + + errors = collect_disabled_option_errors({"federated": {"aggregation": "scaffold"}}, schema) + assert errors == ["federated.aggregation='scaffold' is marked disabled in config_schema.yaml"] diff --git a/apps/backend/tests/services/agent/test_prompts.py b/apps/backend/tests/services/agent/test_prompts.py index 4c9033d..7e22388 100644 --- a/apps/backend/tests/services/agent/test_prompts.py +++ b/apps/backend/tests/services/agent/test_prompts.py @@ -30,10 +30,24 @@ def test_build_plan_prompt_includes_history_and_context(): state=state, capabilities={"datasets": ["mnist"]}, base_config={"federated": {"num_clients": 3}}, + schema_context={ + "fields": [ + { + "path": "federated.aggregation", + "executable_options": ["fedavg"], + "disabled_options": ["scaffold"], + } + ] + }, + config_constraints={"model": {"name": "ResNet"}}, ) assert "Experiment request" in prompt assert "test alpha=0.1,0.3" in prompt - assert "Capabilities" in prompt + assert "Runtime capabilities" in prompt assert "Default base config" in prompt + assert "Schema context from config_schema.yaml" in prompt + assert "User-selected structured constraints" in prompt + assert "disabled_options" in prompt + assert "ResNet" in prompt assert "Return ONLY a JSON object" in prompt diff --git a/apps/backend/tests/services/agent/test_runtime_service.py b/apps/backend/tests/services/agent/test_runtime_service.py index b3b198b..da08d51 100644 --- a/apps/backend/tests/services/agent/test_runtime_service.py +++ b/apps/backend/tests/services/agent/test_runtime_service.py @@ -1,6 +1,7 @@ import asyncio from app.services.agent.runtime_service import AgentRuntimeService +from app.services.agent.state import AgentState, ExperimentRecord def test_runtime_service_marks_stale_queued_task_as_failed(monkeypatch): @@ -41,3 +42,54 @@ async def _explode(): assert persisted_snapshots[-1]["status"] == "failed" asyncio.run(_run()) + + +def test_runtime_service_persists_best_config_and_metrics(monkeypatch): + async def _run(): + service = AgentRuntimeService() + task_id = "task-best" + persisted_snapshots: list[dict] = [] + + async def _fake_persist_snapshot(snapshot): # noqa: ANN001 + persisted_snapshots.append(snapshot) + + monkeypatch.setattr(service, "_persist_snapshot", _fake_persist_snapshot) + + async with service._lock: + service._snapshots.clear() + service._tasks.clear() + service._snapshots[task_id] = { + "task_id": task_id, + "status": "running", + } + + state = AgentState(goal="compare configs") + state.history = [ + ExperimentRecord( + iteration=1, + run_id="run-low", + job_id=1, + name="low", + config={"model": {"name": "CNN"}}, + metrics={"global_results": {"global_accuracy": [0.7]}}, + score=0.7, + ), + ExperimentRecord( + iteration=2, + run_id="run-high", + job_id=2, + name="high", + config={"model": {"name": "ResNet"}}, + metrics={"global_results": {"global_accuracy": [0.9]}}, + score=0.9, + ), + ] + + await service._update_from_state(task_id, state, status="completed") + + snapshot = await service.get_task(task_id) + assert snapshot["best_config"] == {"model": {"name": "ResNet"}} + assert snapshot["best_metrics"] == {"global_results": {"global_accuracy": [0.9]}} + assert persisted_snapshots[-1]["best_config"] == {"model": {"name": "ResNet"}} + + asyncio.run(_run()) diff --git a/apps/frontend/src/api/agent.ts b/apps/frontend/src/api/agent.ts index 3a67e27..0e0cda6 100644 --- a/apps/frontend/src/api/agent.ts +++ b/apps/frontend/src/api/agent.ts @@ -11,6 +11,7 @@ export type AgentOptimizeRequest = { model_name?: string | null; job_name?: string | null; objective: AgentOptimizationObjective; + config_constraints?: Record; }; export type AgentConfigChange = { @@ -91,6 +92,7 @@ export type AgentOptimizeProgressResponse = { best_metrics: RunMetrics | null; experiments: AgentExperimentSummary[]; draft_experiments?: any[]; + config_constraints?: Record; summary_text: string | null; error_message: string | null; created_at: string | null; @@ -168,6 +170,7 @@ export interface AgentPlanPreviewRequest { job_name: string; model_name: string | null; system_mode: string; + config_constraints?: Record; }; export interface AgentPlanPreviewResponse { @@ -175,6 +178,7 @@ export interface AgentPlanPreviewResponse { goal: string; experiments: ExperimentPlanPreview[]; system_mode: string; + config_constraints?: Record; }; export interface AgentPlanReviseRequest { @@ -189,6 +193,11 @@ async function readJson(response: Response): Promise { } export const agentApi = { + async getConfigSchema(): Promise> { + const response = await fetch(`${baseUrl}/api/v1/agent/config/schema`); + return readJson>(response); + }, + async listModels(): Promise { const response = await fetch(`${baseUrl}/api/v1/agent/models`); return readJson(response); @@ -294,4 +303,4 @@ export const agentApi = { } return response.json(); }, -}; \ No newline at end of file +}; diff --git a/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx b/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx index 887a3b8..ef8e760 100644 --- a/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx +++ b/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx @@ -1,12 +1,54 @@ -import { Sparkles, ArrowRight, Loader2 } from "lucide-react"; +import { ArrowRight, Loader2, SlidersHorizontal, Sparkles, X } from "lucide-react"; +import { useMemo } from "react"; + +import { Badge } from "../../../components/ui/badge"; import { Button } from "../../../components/ui/button"; -import { Card, CardContent, CardHeader, CardTitle, CardDescription } from "../../../components/ui/card"; -import { Textarea } from "../../../components/ui/textarea"; +import { Card, CardContent } from "../../../components/ui/card"; import { Input } from "../../../components/ui/input"; +import { NumberStepper } from "../../../components/ui/number-stepper"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "../../../components/ui/select"; +import { Switch } from "../../../components/ui/switch"; +import { Textarea } from "../../../components/ui/textarea"; import type { AgentPageProps } from "../../../pages/types"; +import { getValueByPath } from "../../simulation/utils"; +import { + buildAgentConfig, + checkDependency, + coerceFieldValue, + collectAgentSchemaFields, + formatFieldValue, + optionDisabled, + optionLabel, + optionMeta, + selectedConstraintChips, + type AgentSchemaField, +} from "../schema"; export function AgentExperimentStudio(props: AgentPageProps) { - const { goal, setGoal, jobName, setJobName, presets, busy, handleGeneratePlan } = props; + const { + clearConfigConstraint, + configConstraints, + configSchema, + goal, + setConfigConstraint, + setGoal, + jobName, + setJobName, + presets, + busy, + handleGeneratePlan, + } = props; + + const fields = useMemo(() => collectAgentSchemaFields(configSchema, { featuredOnly: true }), [configSchema]); + const effectiveConfig = useMemo( + () => buildAgentConfig(configSchema, configConstraints), + [configConstraints, configSchema], + ); + const chips = useMemo( + () => selectedConstraintChips(configSchema, configConstraints), + [configConstraints, configSchema], + ); + const visibleFields = fields.filter((field) => checkDependency(effectiveConfig, field.definition.depends_on)); return (
@@ -47,6 +89,49 @@ export function AgentExperimentStudio(props: AgentPageProps) { ))}
+ +
+
+
+ +

Default Config

+
+
+ + {visibleFields.length === 0 && ( +

Config schema is loading.

+ )} + {visibleFields.length > 0 && ( +
+ {visibleFields.map((field) => ( + setConfigConstraint(field.path, coerceFieldValue(field, value))} + /> + ))} +
+ )} + + {chips.length > 0 && ( +
+ {chips.map((chip) => ( + + {chip.label}: {chip.value} + + + ))} +
+ )} +
@@ -64,10 +149,87 @@ export function AgentExperimentStudio(props: AgentPageProps) { ) : ( <> - Generate Experiment Plan + Generate Plan )} ); -} \ No newline at end of file +} + +function parseOptionalNumber(value: unknown): number | undefined { + if (typeof value === "number" && Number.isFinite(value)) return value; + if (typeof value === "string" && value.trim().length > 0) { + const parsed = Number(value); + if (Number.isFinite(parsed)) return parsed; + } + return undefined; +} + +function ConstraintField({ + field, + value, + onChange, +}: { + field: AgentSchemaField; + value: unknown; + onChange: (value: unknown) => void; +}) { + const hint = field.definition.ui && typeof field.definition.ui.prompt_hint === "string" + ? field.definition.ui.prompt_hint + : null; + + return ( +
+
+

{field.label}

+ {field.section} +
+ {field.type === "select" && ( + + )} + {field.type === "number" && ( + + )} + {field.type === "bool" && ( +
+ {formatFieldValue(value)} + +
+ )} + {field.type !== "select" && field.type !== "number" && field.type !== "bool" && ( + onChange(event.target.value)} /> + )} + {hint &&

{hint}

} +
+ ); +} diff --git a/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx b/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx index 9ff9f2d..10dc599 100644 --- a/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx +++ b/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx @@ -1,23 +1,64 @@ -import { useState } from "react"; -import { Play, Clock, Cpu, AlertTriangle, Sparkles, MessageSquare, ClipboardList, Target } from "lucide-react"; -import { Button } from "../../../components/ui/button"; -import { Card, CardContent, CardHeader, CardTitle, CardDescription } from "../../../components/ui/card"; +import { useEffect, useMemo, useState } from "react"; +import { ClipboardList, MessageSquare, Play, Sparkles, Target } from "lucide-react"; +import { toast } from "sonner"; + +import { agentApi } from "../../../api/agent"; import { Badge } from "../../../components/ui/badge"; +import { Button } from "../../../components/ui/button"; +import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "../../../components/ui/card"; import { Input } from "../../../components/ui/input"; +import { NumberStepper } from "../../../components/ui/number-stepper"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "../../../components/ui/select"; +import { Switch } from "../../../components/ui/switch"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "../../../components/ui/tabs"; import { Textarea } from "../../../components/ui/textarea"; -import { toast } from "sonner"; -import { agentApi } from "../../../api/agent"; import type { AgentPageProps } from "../../../pages/types"; +import { getValueByPath, isRecord } from "../../simulation/utils"; +import { + buildAgentConfig, + checkDependency, + coerceFieldValue, + collectAgentSchemaFields, + formatFieldValue, + optionDisabled, + optionLabel, + optionMeta, + setConfigValue, + validateDisabledOptions, + type AgentSchemaField, +} from "../schema"; + +type LocalExperiment = { + name: string; + plan_summary: string; + config_patch: Record; + _rawJsonString?: string; + _jsonError?: boolean; + _activeTab?: "form" | "json"; +}; export function AgentPlanPreview(props: AgentPageProps) { - const { draftPlan, busy: globalBusy, handleExecutePlan, setWorkflowStep, setDraftPlan } = props; + const { + configSchema, + draftPlan, + busy: globalBusy, + handleExecutePlan, + setWorkflowStep, + setDraftPlan, + } = props; - const [localExperiments, setLocalExperiments] = useState(() => - draftPlan?.experiments ? JSON.parse(JSON.stringify(draftPlan.experiments)) : [] + const [localExperiments, setLocalExperiments] = useState(() => + draftPlan?.experiments ? JSON.parse(JSON.stringify(draftPlan.experiments)) : [], ); const [instruction, setInstruction] = useState(""); const [isRevising, setIsRevising] = useState(false); + const schemaFields = useMemo(() => collectAgentSchemaFields(configSchema), [configSchema]); + + useEffect(() => { + setLocalExperiments(draftPlan?.experiments ? JSON.parse(JSON.stringify(draftPlan.experiments)) : []); + }, [draftPlan?.job_id]); + if (!draftPlan) return null; const isBusy = globalBusy || isRevising; @@ -26,13 +67,15 @@ export function AgentPlanPreview(props: AgentPageProps) { setIsRevising(true); try { const response = await agentApi.revisePlan(draftPlan.job_id, { instruction }); - setLocalExperiments(response.experiments); + const nextExperiments = response.experiments as LocalExperiment[]; + setLocalExperiments(nextExperiments); setDraftPlan({ ...draftPlan, - experiments: response.experiments + experiments: response.experiments, + config_constraints: response.config_constraints ?? draftPlan.config_constraints, }); setInstruction(""); - toast.success("Plan updated successfully by Agent!"); + toast.success("Plan updated successfully by Agent."); } catch (error: any) { toast.error(error.message || "Failed to update plan via Agent."); } finally { @@ -40,34 +83,63 @@ export function AgentPlanPreview(props: AgentPageProps) { } }; - const handleConfigChange = (index: number, newJsonString: string) => { + function patchExperiment(index: number, patch: Partial): void { + setLocalExperiments((current) => current.map((exp, idx) => (idx === index ? { ...exp, ...patch } : exp))); + } + + function handleFormConfigChange(index: number, path: string, value: unknown): void { + setLocalExperiments((current) => + current.map((exp, idx) => { + if (idx !== index) return exp; + const nextConfig = setConfigValue(exp.config_patch ?? {}, path, value); + return { + ...exp, + config_patch: nextConfig, + _rawJsonString: JSON.stringify(nextConfig, null, 2), + _jsonError: false, + }; + }), + ); + } + + function handleJsonConfigChange(index: number, newJsonString: string): void { const updated = [...localExperiments]; try { - updated[index].config_patch = JSON.parse(newJsonString); + const parsed = JSON.parse(newJsonString); + if (!isRecord(parsed)) { + throw new Error("Config must be an object."); + } + updated[index].config_patch = parsed; updated[index]._jsonError = false; - } catch (e) { + } catch { updated[index]._jsonError = true; } updated[index]._rawJsonString = newJsonString; setLocalExperiments(updated); - }; + } const onRunClick = () => { - const hasErrors = localExperiments.some((exp: any) => exp._jsonError); - if (hasErrors) { + const hasJsonErrors = localExperiments.some((exp) => exp._jsonError); + if (hasJsonErrors) { toast.error("Please fix invalid JSON formatting before running."); return; } - const cleanExperiments = localExperiments.map((exp: any) => { - const { _rawJsonString, _jsonError, ...rest } = exp; + + const disabledErrors = localExperiments.flatMap((exp) => + validateDisabledOptions(buildAgentConfig(configSchema, exp.config_patch ?? {}), configSchema), + ); + if (disabledErrors.length > 0) { + toast.error(disabledErrors[0]); + return; + } + + const cleanExperiments = localExperiments.map((exp) => { + const { _rawJsonString, _jsonError, _activeTab, ...rest } = exp; return rest; }); void handleExecutePlan(cleanExperiments); }; - const experiments = draftPlan.experiments || []; - const totalMinutes = experiments.reduce((acc, curr) => acc + (curr.estimated_minutes || 0), 0); - return (
@@ -98,7 +170,7 @@ export function AgentPlanPreview(props: AgentPageProps) {
- setInstruction(e.target.value)} @@ -113,40 +185,6 @@ export function AgentPlanPreview(props: AgentPageProps) { - {/* Resource Estimation Summary */} - {/*
- - -
-
-

Est. Total Time

-

{totalMinutes > 0 ? `~${totalMinutes} mins` : "Unknown"}

-
-
-
- - -
-
-

Peak VRAM Required

-

- {experiments.length > 0 && experiments[0].estimated_gpu_vram_gb - ? `~${experiments[0].estimated_gpu_vram_gb} GB` : "Unknown"} -

-
-
-
- - -
-
-

System Mode

-

Simulation

-
-
-
-
*/} - @@ -154,43 +192,206 @@ export function AgentPlanPreview(props: AgentPageProps) { Proposed Execution Plan - The agent has defined the following {localExperiments.length} steps to execute. Edit the config patch manually or use the AI input above. + Edit each experiment with schema-aware controls or switch to JSON for advanced changes.
- {localExperiments.map((exp: any, idx: number) => { - const jsonString = exp._rawJsonString ?? JSON.stringify(exp.config_patch, null, 2); - return ( -
- -
+ {localExperiments.map((exp, idx) => ( + handleFormConfigChange(idx, path, value)} + onJsonChange={(text) => handleJsonConfigChange(idx, text)} + onTabChange={(tab) => patchExperiment(idx, { _activeTab: tab })} + /> + ))} +
+ + +
+ ); +} -
- Step {idx + 1} - {exp.name} -
- -
-

{exp.plan_summary}

- -
-

Configuration Parameters

-