From e3141257d85f61d40c4d0e5295e841088892ad5f Mon Sep 17 00:00:00 2001 From: Ali Akbar Jamali Date: Fri, 3 Jul 2026 18:00:17 -0600 Subject: [PATCH 1/2] A few changes through validation --- app/ui_agent.py | 143 ++++++++++++++-- app/workflow_extras.py | 44 ++++- server/core/plan_rules.py | 290 +++++++++++++++++++++++++++++++- server/core/ui_config_fields.py | 2 + server/llm/plan_shared.py | 57 ++++++- 5 files changed, 506 insertions(+), 30 deletions(-) diff --git a/app/ui_agent.py b/app/ui_agent.py index 567c47b..cb826f6 100644 --- a/app/ui_agent.py +++ b/app/ui_agent.py @@ -282,6 +282,54 @@ def format_datetime_value(date_value, time_value) -> str: return f"{date_value:%Y-%m-%d} {time_value:%H:%M}" +def _is_ui_placeholder_experiment_date(value: str, *, which: str) -> bool: + from server.core.plan_rules import UI_DEFAULT_EXPERIMENT_END, UI_DEFAULT_EXPERIMENT_START + + value = s(value) + if which == "start": + return value in {UI_DEFAULT_EXPERIMENT_START, UI_DEFAULT_EXPERIMENT_START[:10]} + return value in {UI_DEFAULT_EXPERIMENT_END, UI_DEFAULT_EXPERIMENT_END[:10]} + + +def commit_experiment_datetime_from_widgets( + *, + prev_value: str, + widget_value: str, + which: str, +) -> str: + """Keep Input-tab placeholder dates out of session/plan until the user or plan sets them.""" + prev_value = s(prev_value) + widget_value = s(widget_value) + if prev_value: + return widget_value + if widget_value and not _is_ui_placeholder_experiment_date(widget_value, which=which): + return widget_value + return "" + + +def ui_experiment_dates_for_plan() -> tuple[str, str]: + """Return dates to write into the plan, omitting uncommitted UI placeholders.""" + from server.core.plan_rules import request_mentions_experiment_dates + + tstart = s(st.session_state.tstart) + tend = s(st.session_state.tend) + prompt = user_prompt_for_metadata() or s(st.session_state.get("nl_request", "")) + prompt_has_dates = request_mentions_experiment_dates(prompt) + if tstart and ( + prompt_has_dates or not _is_ui_placeholder_experiment_date(tstart, which="start") + ): + pass + else: + tstart = "" + if tend and ( + prompt_has_dates or not _is_ui_placeholder_experiment_date(tend, which="end") + ): + pass + else: + tend = "" + return tstart, tend + + def bump_input_panel_widget_versions() -> None: """Refresh Input-tab widgets that cache values by Streamlit widget key.""" bump_all_input_widget_versions() @@ -2764,9 +2812,10 @@ def sync_all_ui_fields_to_plan(*, refresh_editor: bool = False, force_editor: bo "domain_def": s(st.session_state.domain_def), "hydrological_model": current_hydrological_model(), "forcing_dataset": s(st.session_state.forcing_dataset), - "experiment_time_start": s(st.session_state.tstart), - "experiment_time_end": s(st.session_state.tend), } + plan_tstart, plan_tend = ui_experiment_dates_for_plan() + values["experiment_time_start"] = plan_tstart + values["experiment_time_end"] = plan_tend st.session_state.selected_pour_point = values["pour_point_coords"] st.session_state.selected_bounding_box = values["bounding_box_coords"] @@ -4384,6 +4433,8 @@ def run_cmd_stream(cmd: list[str], cwd: Path, output_box, log_path: Path | None def augment_request_with_ui(nl: str) -> str: + from server.core.plan_rules import UI_DEFAULT_EXPERIMENT_END, UI_DEFAULT_EXPERIMENT_START + lines = [s(nl), "", "Optional UI inputs (use ONLY if non-empty):"] if s(st.session_state.domain_name): @@ -4403,10 +4454,12 @@ def augment_request_with_ui(nl: str) -> str: if s(st.session_state.domain_def): lines.append(f"- domain_def: {s(st.session_state.domain_def)}") - if s(st.session_state.tstart): - lines.append(f"- experiment_time_start: {s(st.session_state.tstart)}") - if s(st.session_state.tend): - lines.append(f"- experiment_time_end: {s(st.session_state.tend)}") + tstart = s(st.session_state.tstart) + tend = s(st.session_state.tend) + if tstart and tstart != UI_DEFAULT_EXPERIMENT_START: + lines.append(f"- experiment_time_start: {tstart}") + if tend and tend != UI_DEFAULT_EXPERIMENT_END: + lines.append(f"- experiment_time_end: {tend}") lines = wx.augment_request_with_advanced(lines) lines = wx.augment_request_with_calibration(lines) @@ -4643,22 +4696,51 @@ def apply_plan_config_to_ui(plan: dict): if forcing_dataset: st.session_state.forcing_dataset = forcing_dataset + dates_changed = False if tstart: st.session_state.tstart = tstart + dates_changed = True + elif _is_ui_placeholder_experiment_date(s(st.session_state.tstart), which="start") or not s( + st.session_state.tstart + ): + # Planner omitted dates — do not keep Input-tab placeholder defaults. + if s(st.session_state.tstart): + dates_changed = True + st.session_state.tstart = "" if tend: st.session_state.tend = tend + dates_changed = True + elif _is_ui_placeholder_experiment_date(s(st.session_state.tend), which="end") or not s( + st.session_state.tend + ): + if s(st.session_state.tend): + dates_changed = True + st.session_state.tend = "" mpi_val = resolve_num_processes_from_plan_cfg(cfgp) if mpi_val is not None: st.session_state.mpi = mpi_val bump_input_panel_widget_versions() + if dates_changed: + from widget_keys import bump_experiment_datetime_widget_version + + bump_experiment_datetime_widget_version() mark_spatial_inputs_stale() sync_run_folder_from_session() wx.apply_advanced_config_from_plan(cfgp) + prompt = user_prompt_for_metadata() or s(st.session_state.get("nl_request", "")) + if user_forbids_mizuroute(prompt): + st.session_state.routing_model = "" + cfgp.pop("routing_model", None) + cfgp.pop("ROUTING_MODEL", None) + extra = cfgp.get("extra_config") + if isinstance(extra, dict): + extra.pop("routing_model", None) + extra.pop("ROUTING_MODEL", None) st.session_state.refresh_spatial_inputs = True @@ -6606,8 +6688,10 @@ def render_workflow_input_tab() -> None: render_start_load_run_section() - start_dt = parse_datetime_value(st.session_state.tstart, dt.datetime(2001, 1, 1, 1, 0)) - end_dt = parse_datetime_value(st.session_state.tend, dt.datetime(2001, 1, 10, 23, 0)) + prev_tstart = s(st.session_state.tstart) + prev_tend = s(st.session_state.tend) + start_dt = parse_datetime_value(prev_tstart, dt.datetime(2001, 1, 1, 1, 0)) + end_dt = parse_datetime_value(prev_tend, dt.datetime(2001, 1, 10, 23, 0)) ws1, ws2, ws3 = st.columns(3) with ws1: @@ -6697,7 +6781,11 @@ def render_workflow_input_tab() -> None: value=start_dt.time().replace(second=0, microsecond=0), key=experiment_datetime_widget_key("experiment_start_time"), ) - st.session_state.tstart = format_datetime_value(start_date, start_time) + st.session_state.tstart = commit_experiment_datetime_from_widgets( + prev_value=prev_tstart, + widget_value=format_datetime_value(start_date, start_time), + which="start", + ) with ws9: end_date = st.date_input( "End date", @@ -6709,9 +6797,19 @@ def render_workflow_input_tab() -> None: value=end_dt.time().replace(second=0, microsecond=0), key=experiment_datetime_widget_key("experiment_end_time"), ) - st.session_state.tend = format_datetime_value(end_date, end_time) - - st.caption(f"Experiment time window: `{st.session_state.tstart}` → `{st.session_state.tend}`") + st.session_state.tend = commit_experiment_datetime_from_widgets( + prev_value=prev_tend, + widget_value=format_datetime_value(end_date, end_time), + which="end", + ) + + if s(st.session_state.tstart) or s(st.session_state.tend): + st.caption( + f"Experiment time window: `{st.session_state.tstart or '—'}` → " + f"`{st.session_state.tend or '—'}`" + ) + else: + st.caption("Experiment time window: not set (placeholder dates are display-only).") st.markdown('', unsafe_allow_html=True) st.markdown( @@ -6868,10 +6966,23 @@ def sync_preview_artifacts() -> None: plan_to_save["config"]["domain_def"] = s(st.session_state.domain_def) if s(st.session_state.forcing_dataset): plan_to_save["config"]["forcing_dataset"] = s(st.session_state.forcing_dataset) - if s(st.session_state.tstart): - plan_to_save["config"]["experiment_time_start"] = s(st.session_state.tstart) - if s(st.session_state.tend): - plan_to_save["config"]["experiment_time_end"] = s(st.session_state.tend) + plan_tstart, plan_tend = ui_experiment_dates_for_plan() + if plan_tstart: + plan_to_save["config"]["experiment_time_start"] = plan_tstart + else: + plan_to_save["config"].pop("experiment_time_start", None) + if plan_tend: + plan_to_save["config"]["experiment_time_end"] = plan_tend + else: + plan_to_save["config"].pop("experiment_time_end", None) + prompt = user_prompt_for_metadata() or s(st.session_state.get("nl_request", "")) + if user_forbids_mizuroute(prompt): + plan_to_save["config"].pop("routing_model", None) + plan_to_save["config"].pop("ROUTING_MODEL", None) + extra = plan_to_save["config"].get("extra_config") + if isinstance(extra, dict): + extra.pop("routing_model", None) + extra.pop("ROUTING_MODEL", None) mpi_val = resolve_num_processes_from_plan_cfg(plan_to_save.get("config") or {}) or int( st.session_state.mpi diff --git a/app/workflow_extras.py b/app/workflow_extras.py index 8aad4a6..1dda4af 100644 --- a/app/workflow_extras.py +++ b/app/workflow_extras.py @@ -392,11 +392,15 @@ def get_advanced_config_values(plan_cfg: dict | None = None) -> dict[str, str]: def apply_advanced_config_from_plan(cfg: dict) -> None: + from server.core.ui_config_fields import user_forbids_mizuroute + extra = cfg.get("extra_config") if isinstance(cfg.get("extra_config"), dict) else {} for session_key, yaml_key, _ in ADVANCED_SESSION_FIELDS: val = s(cfg.get(session_key)) or s(cfg.get(yaml_key)) or s(extra.get(session_key)) or s(extra.get(yaml_key)) if val: st.session_state[session_key] = val + if user_forbids_mizuroute(_user_request_for_advanced_sync()): + st.session_state.routing_model = "" apply_calibration_config_from_plan(cfg) @@ -416,17 +420,42 @@ def merge_advanced_into_spec(spec: dict, plan_cfg: dict | None = None) -> dict: return merge_calibration_into_spec(spec, plan_cfg) +def _user_request_for_advanced_sync() -> str: + return s(st.session_state.get("user_prompt")) or s(st.session_state.get("nl_request")) + + def sync_advanced_config_to_plan() -> None: + from server.core.ui_config_fields import user_forbids_mizuroute + if not st.session_state.get("run_plan"): return plan = st.session_state.run_plan plan.setdefault("config", {}) values = get_advanced_config_values() + user_request = _user_request_for_advanced_sync() + if user_forbids_mizuroute(user_request): + values.pop("routing_model", None) + st.session_state.routing_model = "" + plan["config"].pop("routing_model", None) + plan["config"].pop("ROUTING_MODEL", None) extra = {k: v for k, v in values.items() if v} + existing_extra = plan["config"].get("extra_config") + if isinstance(existing_extra, dict): + merged_extra = dict(existing_extra) + merged_extra.update(extra) + if user_forbids_mizuroute(user_request): + merged_extra.pop("routing_model", None) + merged_extra.pop("ROUTING_MODEL", None) + extra = {k: v for k, v in merged_extra.items() if v} if extra: plan["config"]["extra_config"] = extra for k, v in extra.items(): plan["config"][k] = v + elif "extra_config" in plan["config"] and user_forbids_mizuroute(user_request): + extra_cfg = plan["config"].get("extra_config") + if isinstance(extra_cfg, dict): + extra_cfg.pop("routing_model", None) + extra_cfg.pop("ROUTING_MODEL", None) st.session_state.run_plan = plan @@ -457,9 +486,16 @@ def render_advanced_config_section() -> None: value=s(st.session_state.get("station_id")), key=input_panel_widget_key("adv_station_id"), ) + from server.core.ui_config_fields import user_forbids_mizuroute + + routing_default = ( + "" + if user_forbids_mizuroute(_user_request_for_advanced_sync()) + else "mizuRoute" + ) st.session_state.routing_model = st.text_input( "Routing model", - value=s(st.session_state.get("routing_model")) or "mizuRoute", + value=s(st.session_state.get("routing_model")) or routing_default, key=input_panel_widget_key("adv_routing_model"), ) with c2: @@ -494,7 +530,13 @@ def render_advanced_config_section() -> None: def augment_request_with_advanced(lines: list[str]) -> list[str]: + from server.core.ui_config_fields import user_forbids_mizuroute + + user_request = _user_request_for_advanced_sync() + forbid_mizu = user_forbids_mizuroute(user_request) for session_key, _, label in ADVANCED_SESSION_FIELDS: + if session_key == "routing_model" and forbid_mizu: + continue val = s(st.session_state.get(session_key)) if val: lines.append(f"- {session_key}: {val}") diff --git a/server/core/plan_rules.py b/server/core/plan_rules.py index ddbd22a..970281d 100644 --- a/server/core/plan_rules.py +++ b/server/core/plan_rules.py @@ -346,6 +346,12 @@ def normalize_committed_plan_config(cfg: dict | None) -> dict: return out +_DOMAIN_META_WORDS = frozenset({ + "name", "called", "named", "the", "a", "an", "is", "for", "my", "our", + "this", "that", "as", "to", "of", "in", "with", "from", +}) + + def extract_explicit_domain_name_from_request(user_request: str) -> str: """Return domain_name only when the user explicitly names it in the prompt.""" text = _s(user_request) @@ -353,8 +359,14 @@ def extract_explicit_domain_name_from_request(user_request: str) -> str: return "" patterns = ( + # High-specificity patterns first r"\buse\s+([A-Za-z0-9_\-]+)\s+as\s+(?:the\s+)?domain(?:\s+name|_name)\b", r"\b(?:set|change|update)\s+(?:the\s+)?domain(?:\s+name|_name)\s+(?:to\s+)?([A-Za-z0-9_\-]+)\b", + r"\bdomain\s+called\s+([A-Za-z0-9_\-]+)", + r"\b(?:called|named)\s+([A-Za-z0-9_\-]+)\s+for\s+experiment\b", + r"\bfor\s+domain\s+([A-Za-z0-9_\-]+)", + r"\bdomain\s+([A-Za-z0-9_\-]+)\s+with\s+experiment\b", + r"\bdomain\s+([A-Za-z0-9_\-]+)\s*,\s*experiment\b", r"\bdomain_name\s+to\s+([A-Za-z0-9_\-]+)", r"\b(?:set|change|update)\s+domain_name\s+to\s+([A-Za-z0-9_\-]+)", r"\bdomain_name\s*[=:]\s*[\"']?([A-Za-z0-9_\-]+)", @@ -371,10 +383,132 @@ def extract_explicit_domain_name_from_request(user_request: str) -> str: for pattern in patterns: match = re.search(pattern, text, flags=re.IGNORECASE) if match: - return normalize_domain_name_token(match.group(1)) + token = match.group(1) + if token.lower() in _DOMAIN_META_WORDS: + continue + return normalize_domain_name_token(token) return "" +UI_DEFAULT_EXPERIMENT_START = "2001-01-01 01:00" +UI_DEFAULT_EXPERIMENT_END = "2001-01-10 23:00" + + +def extract_experiment_dates_from_request(user_request: str) -> tuple[str, str]: + """Parse common natural-language experiment windows from the user prompt.""" + text = _s(user_request) + if not text: + return "", "" + + year_range = re.search(r"\bfrom\s+(\d{4})\s+to\s+(\d{4})\b", text, flags=re.IGNORECASE) + if year_range: + start_year, end_year = year_range.group(1), year_range.group(2) + return f"{start_year}-01-01 01:00", f"{end_year}-12-31 23:00" + + patterns = ( + r"\bfrom\s+(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})\b", + r"\bdates?\s+(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})\b", + r"\b(?:experiment\s+window|run)\s+(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})\b", + r"\b(?:calibration|evaluation)\s+period\s+(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})\b", + ) + for pattern in patterns: + match = re.search(pattern, text, flags=re.IGNORECASE) + if match: + return f"{match.group(1)} 01:00", f"{match.group(2)} 23:00" + return "", "" + + +def request_mentions_experiment_dates(user_request: str) -> bool: + text = _s(user_request) + if not text: + return False + if extract_experiment_dates_from_request(text) != ("", ""): + return True + return bool(re.search(r"\b\d{4}-\d{2}-\d{2}\b", text)) + + +def sanitize_planner_experiment_dates(cfg: dict, user_request: str) -> dict: + """Apply prompt dates and drop UI placeholder windows not mentioned by the user.""" + out = dict(cfg or {}) + start, end = extract_experiment_dates_from_request(user_request) + if start: + out["experiment_time_start"] = start + if end: + out["experiment_time_end"] = end + + prompt_has_dates = request_mentions_experiment_dates(user_request) + prompt_mentions_2001 = bool(re.search(r"\b2001\b", _s(user_request))) + + for key in ("experiment_time_start", "experiment_time_end"): + value = _s(out.get(key)) + if not value: + continue + is_ui_default = value in { + UI_DEFAULT_EXPERIMENT_START, + UI_DEFAULT_EXPERIMENT_END, + UI_DEFAULT_EXPERIMENT_START[:10], + UI_DEFAULT_EXPERIMENT_END[:10], + } + is_2001 = value.startswith("2001-") + if is_ui_default and not prompt_has_dates: + out.pop(key, None) + elif is_2001 and not prompt_mentions_2001: + out.pop(key, None) + return out + + +def apply_workflow_config_policies(cfg: dict, user_request: str = "") -> dict: + """Normalize discretization, lumped routing, and elevation extras from the prompt.""" + from server.core.ui_config_fields import ( + is_lumped_workflow, + normalize_discretization_value, + symfluence_discretization_from_plan, + ) + + out = dict(cfg or {}) + disc = normalize_discretization_value(_s(out.get("discretization"))) + if disc: + out["discretization"] = disc + + if re.search(r"\b(?:do\s+not\s+use|without|no)\s+mizu\s*route\b", user_request, flags=re.IGNORECASE): + out.pop("routing_model", None) + + if is_lumped_workflow(out, user_request): + extra = dict(out.get("extra_config") or {}) if isinstance(out.get("extra_config"), dict) else {} + extra.setdefault("ROUTING_DELINEATION", "lumped") + extra.setdefault("PARAMETER_REGIONALIZATION", "lumped") + out["extra_config"] = extra + out["discretization"] = symfluence_discretization_from_plan(out, user_request) + + band_match = re.search(r"\belevation\s+band\s+size\s+(\d+)", user_request, flags=re.IGNORECASE) + if band_match: + extra = dict(out.get("extra_config") or {}) if isinstance(out.get("extra_config"), dict) else {} + extra["ELEVATION_BAND_SIZE"] = int(band_match.group(1)) + out["extra_config"] = extra + + if re.search(r"\bGRUs?\b", user_request, flags=re.IGNORECASE): + out["discretization"] = "GRUs" + + return out + + +def domain_name_literal_in_request(domain_name: str, user_request: str) -> bool: + """True when the prompt names this basin in a domain context (not weak inference).""" + token = normalize_domain_name_token(domain_name) + text = _s(user_request) + if not token or not text: + return False + if extract_explicit_domain_name_from_request(text).lower() == token.lower(): + return True + patterns = ( + rf"\bfor\s+domain\s+{re.escape(token)}\b", + rf"\bdomain\s+{re.escape(token)}\b", + rf"\bcalled\s+{re.escape(token)}\b", + rf"\bdomain\s+called\s+{re.escape(token)}\b", + ) + return any(re.search(pattern, text, flags=re.IGNORECASE) for pattern in patterns) + + def domain_name_needs_user_input( cfg: dict | None, user_request: str = "", @@ -399,6 +533,9 @@ def domain_name_needs_user_input( if explicit and explicit.lower() == name.lower(): return False + if domain_name_literal_in_request(name, user_request): + return False + if is_weak_domain_name(name): return True @@ -441,13 +578,15 @@ def ensure_domain_name_user_input( cfg = mark_domain_name_confirmed(cfg) elif _s(cfg.get("domain_name")) and domain_name_needs_user_input(cfg, user_request, data_dir=data_dir): weak = is_weak_domain_name(_s(cfg.get("domain_name"))) - cfg.pop("domain_name", None) - if weak: + if weak and not domain_name_literal_in_request(_s(cfg.get("domain_name")), user_request): + cfg.pop("domain_name", None) _append_plan_note( out, "Removed weak inferred domain_name; provide a filesystem-safe basin name " "(e.g. Bow_at_Banff_semi_distributed).", ) + elif weak and domain_name_literal_in_request(_s(cfg.get("domain_name")), user_request): + cfg = apply_user_provided_domain_name(cfg, _s(cfg.get("domain_name"))) out["config"] = cfg if not domain_name_needs_user_input(cfg, user_request, data_dir=data_dir): @@ -1141,7 +1280,7 @@ def _plan_steps_are_gated(steps: list[str] | None) -> bool: def _plan_ready_for_step_inference(plan: dict) -> bool: needs = set(plan.get("needs_user_input") or []) - return not (needs - {"bounding_box_coords"}) + return not (needs - {"bounding_box_coords", "pour_point_coords"}) def infer_gated_plan_steps( @@ -1693,6 +1832,13 @@ def normalize_local_workflow_plan( out["config"] = cfg steps = list(out.get("steps") or []) + # Honor exclusive validate/dry-run requests before any step expansion. + if request_requests_validation_dry_run_only(user_request): + out["steps"] = ["validate_config", "dry_run"] + out = ensure_domain_name_user_input(out, user_request, data_dir=data_dir) + out = apply_explicit_step_constraints(out, user_request) + return _drop_satisfied_needs_user_input(out) + out = try_restore_local_recovery_plan(out, user_request, data_dir=data_dir) cfg = dict(out.get("config") or {}) out["config"] = cfg @@ -1710,6 +1856,7 @@ def normalize_local_workflow_plan( out = ensure_local_data_access_in_plan(out) out = ensure_online_data_when_missing(out, user_request, data_dir=data_dir) cfg = dict(out.get("config") or {}) + cfg = apply_workflow_config_policies(cfg, user_request) out["config"] = cfg steps = list(out.get("steps") or []) @@ -1724,6 +1871,10 @@ def normalize_local_workflow_plan( out = infer_gated_plan_steps(out, user_request, data_dir=data_dir) steps = list(out.get("steps") or []) + out = apply_explicit_step_constraints(out, user_request) + steps = list(out.get("steps") or []) + cfg = dict(out.get("config") or {}) + if not plan_uses_local_data(cfg, steps, user_request, data_dir=data_dir): return _drop_satisfied_needs_user_input(out) @@ -1754,12 +1905,97 @@ def extract_steps_from_request(user_request: str, allowed_steps: list[str]) -> l return found +def _user_forbids_step_inference(user_request: str, step_keyword: str) -> bool: + """True when the prompt explicitly forbids a step (e.g. 'do not run the model').""" + text = _s(user_request).lower() + kw = step_keyword.lower() + neg = ( + rf"\b(?:do\s+not|don'?t|without|skip|omit|exclude|no)\s+(?:\w+\s+){{0,3}}{re.escape(kw)}\b", + rf"\b(?:do\s+not|don'?t)\s+{re.escape(kw)}\b", + ) + return any(re.search(p, text) for p in neg) + + +def request_forbids_run_model(user_request: str) -> bool: + return ( + _user_forbids_step_inference(user_request, "run the model") + or _user_forbids_step_inference(user_request, "run model") + or _user_forbids_step_inference(user_request, "run_model") + or _user_forbids_step_inference(user_request, "model run") + ) + + +def request_forbids_calibrate(user_request: str) -> bool: + return ( + _user_forbids_step_inference(user_request, "calibrat") + or _user_forbids_step_inference(user_request, "calibration") + ) + + +def request_forbids_postprocess(user_request: str) -> bool: + return ( + _user_forbids_step_inference(user_request, "postprocess") + or _user_forbids_step_inference(user_request, "postprocessing") + or _user_forbids_step_inference(user_request, "post-process") + ) + + +def request_requests_validation_dry_run_only(user_request: str) -> bool: + """True when the user asks only for validate_config / dry_run.""" + text = _s(user_request).lower() + if not text: + return False + patterns = ( + r"\bonly\s+validate_config\s+and\s+dry_run\b", + r"\bonly\s+validate(?:_config)?\s+and\s+dry[_\s-]?run\b", + r"\bvalidate_config\s+and\s+dry_run\s+only\b", + r"\bonly\s+(?:do\s+)?validate(?:_config)?\s+and\s+(?:a\s+)?dry[_\s-]?run\b", + r"\bvalidate(?:_config)?\s+and\s+(?:run\s+a\s+)?dry[_\s-]?run\b.*\b(?:skip|do\s+not|don'?t|without)\b", + ) + return any(re.search(p, text) for p in patterns) + + +def apply_explicit_step_constraints(plan: dict, user_request: str = "") -> dict: + """Honor exclusive validate/dry-run requests and negative step constraints.""" + if not isinstance(plan, dict): + return plan + out = dict(plan) + text = _s(user_request) + if request_requests_validation_dry_run_only(text): + out["steps"] = ["validate_config", "dry_run"] + _append_plan_note(out, "Restricted steps to validate_config and dry_run per user request.") + return out + + steps = list(out.get("steps") or []) + if not steps: + return out + + remove: set[str] = set() + if request_forbids_run_model(text): + remove.update({"run_model", "calibrate_model"}) + if request_forbids_calibrate(text) or not re.search(r"\bcalibrat", text.lower()): + remove.add("calibrate_model") + if request_forbids_postprocess(text): + remove.add("postprocess_results") + + if remove: + out["steps"] = [step for step in steps if step not in remove] + return out + + def infer_goal_steps_from_request(user_request: str) -> list[str]: """Infer high-level workflow goals from natural-language hydrology prompts.""" text = _s(user_request).lower() if not text: return [] + if request_requests_validation_dry_run_only(user_request): + return ["validate_config", "dry_run"] + + forbids_model = request_forbids_run_model(user_request) + forbids_calibrate = request_forbids_calibrate(user_request) + forbids_postprocess = request_forbids_postprocess(user_request) + goals: list[str] = [] def add(step: str) -> None: @@ -1769,15 +2005,35 @@ def add(step: str) -> None: if re.search(r"\bprocess_observed\b", text) or re.search(r"\bobserved\s+streamflow\b", text): add("process_observed_data") - if ( + if not forbids_model and ( re.search(r"\bfrom\s+scratch\b", text) or re.search(r"\brun\s+(?:the\s+)?model\b", text) + or re.search(r"\b(?:and|then)\s+(?:the\s+)?model\b", text) or re.search(r"\brun\s+summa\b", text) or re.search(r"\bsumma\s+workflow\b", text) or (re.search(r"\bworkflow\b", text) and re.search(r"\bsumma\b", text)) ): add("run_model") + if not forbids_calibrate and re.search(r"\bcalibrat(?:e|ion)\s+(?:the\s+)?model\b", text): + add("calibrate_model") + + # Require an affirmative postprocess mention, not "skip postprocessing". + if not forbids_postprocess and re.search( + r"\b(?:include\s+)?postprocess(?:ing|_results)?\b", text + ) and not re.search( + r"\b(?:do\s+not|don'?t|without|skip|omit|exclude|no)\b.{0,20}\bpostprocess", + text, + ): + add("postprocess_results") + + if re.search(r"\bpreprocess", text): + add("model_agnostic_preprocessing") + add("model_specific_preprocessing") + + if re.search(r"\bsetup\b", text): + add("setup_project") + if re.search(r"\bsemi[- ]?distributed\b", text): add("define_domain") add("discretize_domain") @@ -1785,10 +2041,14 @@ def add(step: str) -> None: if re.search(r"\bprepare\s+.*\binput\b", text) or re.search(r"\bsumma\s+input\b", text): add("model_specific_preprocessing") - if re.search(r"\bforcing\b", text): + if re.search(r"\bforcings?\b", text) or re.search(r"\bdownload\b.*\bforcings?\b", text): add("acquire_forcings") - if re.search(r"\bacquire\s+attribute", text) or re.search(r"\battributes\s+and\b", text): + if ( + re.search(r"\bacquire\s+attribute", text) + or re.search(r"\bdownload\s+attribute", text) + or re.search(r"\battributes\s+and\b", text) + ): add("acquire_attributes") if ( @@ -1799,6 +2059,22 @@ def add(step: str) -> None: add("define_domain") for step in extract_steps_from_request(user_request, WORKFLOW_STEP_NAMES): + if step == "run_model" and forbids_model: + continue + if step == "calibrate_model" and forbids_calibrate: + continue + if step == "postprocess_results" and forbids_postprocess: + continue + if step == "dry_run": + add(step) + continue add(step) + if forbids_model: + goals = [g for g in goals if g not in ("run_model", "calibrate_model")] + if forbids_calibrate: + goals = [g for g in goals if g != "calibrate_model"] + if forbids_postprocess: + goals = [g for g in goals if g != "postprocess_results"] + return goals diff --git a/server/core/ui_config_fields.py b/server/core/ui_config_fields.py index 74e428a..6671359 100644 --- a/server/core/ui_config_fields.py +++ b/server/core/ui_config_fields.py @@ -334,6 +334,8 @@ def normalize_discretization_value(raw: str) -> str: return "lumped" if lower == "elevation": return "elevation" + if lower in {"e", "elev"} or lower.startswith("elev"): + return "elevation" return raw diff --git a/server/llm/plan_shared.py b/server/llm/plan_shared.py index 5c415df..e2daa5d 100644 --- a/server/llm/plan_shared.py +++ b/server/llm/plan_shared.py @@ -7,6 +7,8 @@ from typing import Any, Dict, List from server.core.plan_rules import ( + apply_explicit_step_constraints, + apply_workflow_config_policies, default_symfluence_data_dir, domain_name_needs_user_input, ensure_domain_name_user_input, @@ -17,6 +19,11 @@ plan_requires_bounding_box, pour_point_workflow_skips_bbox, request_indicates_local_data_reuse, + request_forbids_calibrate, + request_forbids_postprocess, + request_forbids_run_model, + request_requests_validation_dry_run_only, + sanitize_planner_experiment_dates, sort_plan_steps_by_workflow_order, try_restore_local_recovery_plan, ) @@ -39,6 +46,27 @@ def _extract_domain_name(text: str): return extract_explicit_domain_name_from_request(text) or None +def _apply_negative_step_constraints(steps: list[str], user_request: str) -> list[str]: + """Remove steps the user explicitly forbade and unrequested calibrate_model.""" + if not steps or not user_request: + return steps + if request_requests_validation_dry_run_only(user_request): + return ["validate_config", "dry_run"] + + remove: set[str] = set() + if request_forbids_run_model(user_request): + remove.update({"run_model", "calibrate_model"}) + if request_forbids_calibrate(user_request) or not re.search( + r"\bcalibrat", user_request.lower() + ): + remove.add("calibrate_model") + if request_forbids_postprocess(user_request): + remove.add("postprocess_results") + if not remove: + return steps + return [step for step in steps if step not in remove] + + def _compact_plan_config(cfg: dict) -> dict: core_keys = { "domain_name", @@ -253,6 +281,9 @@ def finalize_run_plan( if isinstance(end, str) and len(end.strip()) == 10: cfg["experiment_time_end"] = end.strip() + " 23:00" + cfg = sanitize_planner_experiment_dates(cfg, user_request) + cfg = apply_workflow_config_policies(cfg, user_request) + if not cfg.get("experiment_id"): cfg["experiment_id"] = "exp_001" @@ -323,12 +354,24 @@ def finalize_run_plan( bbox_missing = "bounding_box_coords" in missing if missing_core: - plan["needs_user_input"] = missing - plan["steps"] = ["validate_config", "dry_run"] - plan["notes"] = ( - f"Missing required inputs: {', '.join(missing)}. " - "Returning a safe validation/dry-run plan until those values are provided." - ) + if ( + missing_core == ["pour_point_coords"] + and request_indicates_local_data_reuse(user_request, cfg, data_dir=data_dir) + ): + plan["needs_user_input"] = ["pour_point_coords"] + from server.core.plan_rules import _append_plan_note + + _append_plan_note( + plan, + "Missing pour_point_coords. Workflow steps preserved for local-data reuse.", + ) + else: + plan["needs_user_input"] = missing + plan["steps"] = ["validate_config", "dry_run"] + plan["notes"] = ( + f"Missing required inputs: {', '.join(missing)}. " + "Returning a safe validation/dry-run plan until those values are provided." + ) elif bbox_missing: plan["needs_user_input"] = ["bounding_box_coords"] from server.core.plan_rules import _append_plan_note @@ -386,6 +429,8 @@ def finalize_run_plan( if ordered_steps: plan["steps"] = ordered_steps + plan["steps"] = _apply_negative_step_constraints(plan.get("steps", []), user_request) + extracted_fields = [] for k in [ "domain_name", From 310b1395fac53d3b51245e5515c25fb13fea75a1 Mon Sep 17 00:00:00 2001 From: Ali Akbar Jamali Date: Sat, 15 Aug 2026 16:12:25 -0600 Subject: [PATCH 2/2] Fix end-to-end planning for FUSE runs and Gemini/Claude providers. Insert discretize_domain before preprocessing, use space-free duplicate DOMAIN_NAME values so TauDEM can run, and raise structured-output robustness for newer Gemini/Claude models. --- .gitignore | 2 + app/ui_agent.py | 282 +++++++++++++----------- prompts/planner_prompt.txt | 1 + server/core/local_domain.py | 2 +- server/core/plan_rules.py | 119 +++++++++- server/core/run_naming.py | 23 +- server/llm/claude_provider.py | 205 +++++++++++++---- server/llm/gemini_provider.py | 175 +++++++++++---- tests/test_claude_provider.py | 1 + tests/test_gemini_provider.py | 1 + tests/test_plan_resolve_dependencies.py | 2 + tests/test_run_naming.py | 108 +++------ 12 files changed, 630 insertions(+), 291 deletions(-) create mode 100644 tests/test_claude_provider.py create mode 100644 tests/test_gemini_provider.py diff --git a/.gitignore b/.gitignore index 15e3a19..0c170ef 100644 --- a/.gitignore +++ b/.gitignore @@ -4,6 +4,8 @@ venv/ __pycache__/ *.pyc .DS_Store +.pytest_cache/ runs/* !runs/.gitkeep .streamlit/secrets.toml +figures/ diff --git a/app/ui_agent.py b/app/ui_agent.py index cb826f6..94a3c4a 100644 --- a/app/ui_agent.py +++ b/app/ui_agent.py @@ -119,6 +119,7 @@ ensure_skip_acquire_forcings_when_local_forcing, ensure_skip_model_agnostic_when_local_preprocessing, ensure_skip_process_observed_when_local_streamflow, + ensure_create_pour_point_before_define_domain, extract_explicit_domain_name_from_request, extract_station_id_from_request, resolve_station_id_from_plan, @@ -1424,24 +1425,33 @@ def build_symfluence_step_cmd(step: str, config_path: Path) -> list[str]: ALL_GEMINI_MODELS = [m for group in GEMINI_MODELS.values() for m in group] CLAUDE_MODELS = { - "Claude 4.x (Recommended)": [ - "claude-sonnet-4-20250514", - "claude-opus-4-20250514", + "Claude 5 (Recommended)": ["claude-sonnet-5"], + "Claude 4.x": [ + "claude-opus-4-8", + "claude-sonnet-4-6", + "claude-sonnet-4-5", ], - "Claude 3.7": ["claude-3-7-sonnet-latest"], - "Claude 3.5 (Legacy)": [ + "Claude 3.x (Legacy)": [ + "claude-3-7-sonnet-latest", "claude-3-5-sonnet-latest", - "claude-3-5-haiku-latest", + "claude-haiku-4-5", ], } ALL_CLAUDE_MODELS = [m for group in CLAUDE_MODELS.values() for m in group] CLAUDE_MODEL_LABELS = { - "claude-sonnet-4-20250514": "Sonnet 4", - "claude-opus-4-20250514": "Opus 4", + "claude-sonnet-5": "Sonnet 5", + "claude-opus-4-8": "Opus 4.8", + "claude-sonnet-4-6": "Sonnet 4.6", + "claude-sonnet-4-5": "Sonnet 4.5", "claude-3-7-sonnet-latest": "Sonnet 3.7", "claude-3-5-sonnet-latest": "Sonnet 3.5", - "claude-3-5-haiku-latest": "Haiku 3.5", + "claude-haiku-4-5": "Haiku 4.5", +} + +RETIRED_CLAUDE_MODEL_ALIASES = { + "claude-sonnet-4-20250514": "claude-sonnet-5", + "claude-opus-4-20250514": "claude-opus-4-8", } @@ -1470,10 +1480,20 @@ def llm_model_label(model_id: str, provider: str) -> str: DEFAULT_LLM_MODEL = { "openai": "gpt-5-mini", "gemini": "gemini-2.5-flash", - "claude": "claude-sonnet-4-20250514", + "claude": "claude-sonnet-5", } +def resolve_llm_model(model_id: str, provider: str) -> str: + model_id = s(model_id) or DEFAULT_LLM_MODEL.get(provider, "gpt-5-mini") + if provider == "claude": + model_id = RETIRED_CLAUDE_MODEL_ALIASES.get(model_id, model_id) + available = llm_models_for_provider("claude") + if model_id not in available: + model_id = DEFAULT_LLM_MODEL.get("claude", available[0]) + return model_id + + def llm_models_for_provider(provider: str) -> list[str]: if provider == "gemini": return ALL_GEMINI_MODELS @@ -3555,7 +3575,50 @@ def render_map_layer_checkboxes(key_prefix: str) -> int: # Scrollable swatch list up to this many classes; above that use compact gradient summary. MAP_LEGEND_MAX_SWATCHES = 120 MAP_LEGEND_TWO_COLUMN_MIN = 10 -MAP_LEGEND_SCROLL_MAX_HEIGHT_PX = 240 +MAP_LEGEND_SCROLL_MAX_HEIGHT_PX = 140 +MAP_LEGEND_MIN_WIDTH_PX = 150 +MAP_LEGEND_MAX_WIDTH_PX = 190 + + +def _add_workflow_map_css(m: folium.Map) -> None: + """Keep map controls legible and compact inside narrow Streamlit columns.""" + from branca.element import Element + + m.get_root().header.add_child( + Element( + f""" + + """ + ) + ) def _pour_point_legend_icon_html() -> str: @@ -3888,10 +3951,9 @@ def _add_choropleth_legend_to_map( template = Template( """ {% macro html(this, kwargs) %} -
{{ this.title }}
""" @@ -3922,28 +3984,27 @@ def _add_choropleth_legend_to_map( template = Template( """ {% macro html(this, kwargs) %} -
-
- {{ this.title }} +
+ {{ this.title }} {{ this.count }}
-
{% if this.value_header %} -
{{ this.value_header }}{{ this.value_header }}
{% endif %}
+ {% if this.two_column %}display:grid;grid-template-columns:1fr 1fr;column-gap:7px;{% endif %}"> {% for entry in this.entries %} -
- + {{ entry.label }} @@ -3964,10 +4025,9 @@ def _add_choropleth_legend_to_map( template = Template( """ {% macro html(this, kwargs) %} -
{{ this.title }} @@ -4214,7 +4274,12 @@ def build_pour_point_map( center, zoom = _map_view_center_zoom(pour_coords, bbox_bounds) center_lat, center_lon = center[0], center[1] - m = folium.Map(location=[center_lat, center_lon], zoom_start=zoom) + m = folium.Map( + location=[center_lat, center_lon], + zoom_start=zoom, + zoom_control=True, + ) + _add_workflow_map_css(m) if ( _bbox_visible_on_map() @@ -4483,7 +4548,7 @@ def run_generate_plan_from_nl_request() -> None: try: capture_user_prompt_from_session() planner_request = augment_request_with_ui(st.session_state.nl_request) - model = s(st.session_state.get("llm_model")) or DEFAULT_LLM_MODEL.get(provider, "gpt-5-mini") + model = resolve_llm_model(st.session_state.get("llm_model"), provider) if provider == "gemini": plan = GeminiProvider(api_key=key).generate_run_plan( model=model, @@ -4506,6 +4571,7 @@ def run_generate_plan_from_nl_request() -> None: convo, data_dir=SYMFLUENCE_DATA_DIR, ) + plan = resolve_requested_plan_dependencies(plan, convo) if not isinstance(plan, dict): raise RuntimeError(f"Planner returned non-dict: {type(plan)}") @@ -4783,7 +4849,8 @@ def force_steps(plan: dict, want_create_pour_point: bool) -> dict: ordered.insert(i, "setup_project") seen.add("setup_project") - if want_create_pour_point and "create_pour_point" not in seen: + need_pour_point = want_create_pour_point or "define_domain" in seen + if need_pour_point and "create_pour_point" not in seen: i = ordered.index("setup_project") + 1 ordered.insert(i, "create_pour_point") seen.add("create_pour_point") @@ -4796,7 +4863,7 @@ def force_steps(plan: dict, want_create_pour_point: bool) -> dict: return plan -st.set_page_config(page_title="HydroAgent: SYMFLUENCE Workflow Assistant Agent", layout="wide") +st.set_page_config(page_title="HydroAgent: Hydrological Workflow Assistant", layout="wide") st.markdown( """ @@ -4955,7 +5022,7 @@ def force_steps(plan: dict, want_create_pour_point: bool) -> dict:
HydroAgent
-
SYMFLUENCE Workflow Assistant Agent
+
Hydrological Workflow Assistant
✓ System status
@@ -5705,12 +5772,18 @@ def refine_plan_from_chat_message(user_text: str) -> None: conversation_text=conversation_text, emit_chat_messages=False, ) - stored = st.session_state.run_plan or pre_patched append_chat_message( "assistant", - reconcile_chat_reply("", current_plan, stored, user_message=text), + "Updated the plan steps from your message.", kind="text", ) + missing = (st.session_state.run_plan or {}).get("needs_user_input") or [] + if missing: + append_chat_message( + "assistant", + "Still missing before execution: " + ", ".join(missing), + kind="warning", + ) save_chat_messages_to_run_folder() return @@ -5763,11 +5836,6 @@ def refine_plan_from_chat_message(user_text: str) -> None: working_plan = pre_patched if pre_changed else current_plan context_text = build_chat_refinement_context() model = s(st.session_state.get("llm_model")) or DEFAULT_LLM_MODEL.get(provider, "gpt-5-mini") - preserve_steps = _chat_message_preserves_workflow_steps( - text, - before=current_plan, - after=pre_patched if pre_changed else current_plan, - ) try: reply, new_plan, updated = call_llm_refine_run_plan( @@ -5778,7 +5846,6 @@ def refine_plan_from_chat_message(user_text: str) -> None: current_plan=working_plan, conversation_text=conversation_text, context_text=context_text, - preserve_workflow_steps=preserve_steps, ) candidate = new_plan if updated else working_plan final_plan, _ = apply_chat_message_to_plan(candidate, text) @@ -5788,14 +5855,15 @@ def refine_plan_from_chat_message(user_text: str) -> None: final_plan, user_message=text, conversation_text=conversation_text, - emit_chat_messages=False, ) - stored = st.session_state.run_plan or final_plan - append_chat_message( - "assistant", - reconcile_chat_reply(reply, current_plan, stored, user_message=text), - kind="text", - ) + append_chat_message("assistant", reply) + missing = (st.session_state.run_plan or {}).get("needs_user_input") or [] + if missing: + append_chat_message( + "assistant", + "Still missing before execution: " + ", ".join(missing), + kind="warning", + ) save_chat_messages_to_run_folder() except Exception as e: if pre_changed: @@ -6235,6 +6303,10 @@ def render_workflow_assistant_panel() -> None: st.session_state.want_create_pour_point = st.checkbox( "Also run create_pour_point", value=st.session_state.want_create_pour_point, + help=( + "Creates the pour-point shapefile from coordinates. " + "Always inserted before define_domain when that shapefile is missing." + ), ) st.session_state.allow_run = st.checkbox( "Allow dangerous run steps", @@ -6246,15 +6318,13 @@ def render_workflow_assistant_panel() -> None: help="Required for run_model or calibrate_model.", ) - render_fix_missing_inputs_section( - (st.session_state.run_plan or {}).get("needs_user_input", []) or [] - ) - needs = (st.session_state.run_plan or {}).get("needs_user_input", []) or [] + needs = st.session_state.run_plan.get("needs_user_input", []) or [] if needs: st.warning( f"{len(needs)} required field(s) missing before execution. " - "Use the section above or the Input tab." + "Use the section below or the Input tab." ) + render_fix_missing_inputs_section(needs) assistant_shortcut_out = st.empty() wx.render_run_shortcuts_section( @@ -6688,31 +6758,21 @@ def render_workflow_input_tab() -> None: render_start_load_run_section() - prev_tstart = s(st.session_state.tstart) - prev_tend = s(st.session_state.tend) - start_dt = parse_datetime_value(prev_tstart, dt.datetime(2001, 1, 1, 1, 0)) - end_dt = parse_datetime_value(prev_tend, dt.datetime(2001, 1, 10, 23, 0)) + start_dt = parse_datetime_value(st.session_state.tstart, dt.datetime(2001, 1, 1, 1, 0)) + end_dt = parse_datetime_value(st.session_state.tend, dt.datetime(2001, 1, 10, 23, 0)) ws1, ws2, ws3 = st.columns(3) with ws1: - st.text_input( + st.session_state.domain_name = st.text_input( "Domain name", st.session_state.domain_name, key=input_panel_widget_key("input_domain_name"), - on_change=on_input_domain_name_change, - ) - st.session_state.domain_name = s( - st.session_state.get(input_panel_widget_key("input_domain_name")) ) with ws2: - st.text_input( + st.session_state.experiment_id = st.text_input( "Experiment ID", st.session_state.experiment_id, key=input_panel_widget_key("input_experiment_id"), - on_change=on_input_experiment_id_change, - ) - st.session_state.experiment_id = s( - st.session_state.get(input_panel_widget_key("input_experiment_id")) ) with ws3: st.session_state.run_folder = st.text_input( @@ -6781,11 +6841,7 @@ def render_workflow_input_tab() -> None: value=start_dt.time().replace(second=0, microsecond=0), key=experiment_datetime_widget_key("experiment_start_time"), ) - st.session_state.tstart = commit_experiment_datetime_from_widgets( - prev_value=prev_tstart, - widget_value=format_datetime_value(start_date, start_time), - which="start", - ) + st.session_state.tstart = format_datetime_value(start_date, start_time) with ws9: end_date = st.date_input( "End date", @@ -6797,19 +6853,9 @@ def render_workflow_input_tab() -> None: value=end_dt.time().replace(second=0, microsecond=0), key=experiment_datetime_widget_key("experiment_end_time"), ) - st.session_state.tend = commit_experiment_datetime_from_widgets( - prev_value=prev_tend, - widget_value=format_datetime_value(end_date, end_time), - which="end", - ) - - if s(st.session_state.tstart) or s(st.session_state.tend): - st.caption( - f"Experiment time window: `{st.session_state.tstart or '—'}` → " - f"`{st.session_state.tend or '—'}`" - ) - else: - st.caption("Experiment time window: not set (placeholder dates are display-only).") + st.session_state.tend = format_datetime_value(end_date, end_time) + + st.caption(f"Experiment time window: `{st.session_state.tstart}` → `{st.session_state.tend}`") st.markdown('
', unsafe_allow_html=True) st.markdown( @@ -6876,7 +6922,6 @@ def render_workflow_input_tab() -> None: st.caption("Bounding box mode: click first corner, then opposite corner.") sync_input_panel_to_run_plan() - render_run_single_steps_section() shortcut_output = st.empty() wx.render_run_shortcuts_section( execute_bundle_fn=lambda key: execute_step_bundle(key, shortcut_output), @@ -6966,23 +7011,10 @@ def sync_preview_artifacts() -> None: plan_to_save["config"]["domain_def"] = s(st.session_state.domain_def) if s(st.session_state.forcing_dataset): plan_to_save["config"]["forcing_dataset"] = s(st.session_state.forcing_dataset) - plan_tstart, plan_tend = ui_experiment_dates_for_plan() - if plan_tstart: - plan_to_save["config"]["experiment_time_start"] = plan_tstart - else: - plan_to_save["config"].pop("experiment_time_start", None) - if plan_tend: - plan_to_save["config"]["experiment_time_end"] = plan_tend - else: - plan_to_save["config"].pop("experiment_time_end", None) - prompt = user_prompt_for_metadata() or s(st.session_state.get("nl_request", "")) - if user_forbids_mizuroute(prompt): - plan_to_save["config"].pop("routing_model", None) - plan_to_save["config"].pop("ROUTING_MODEL", None) - extra = plan_to_save["config"].get("extra_config") - if isinstance(extra, dict): - extra.pop("routing_model", None) - extra.pop("ROUTING_MODEL", None) + if s(st.session_state.tstart): + plan_to_save["config"]["experiment_time_start"] = s(st.session_state.tstart) + if s(st.session_state.tend): + plan_to_save["config"]["experiment_time_end"] = s(st.session_state.tend) mpi_val = resolve_num_processes_from_plan_cfg(plan_to_save.get("config") or {}) or int( st.session_state.mpi @@ -7217,6 +7249,21 @@ def show_text(text: str): ) if isinstance(committed_steps, list) and committed_steps: plan["steps"] = list(committed_steps) + plan = force_steps( + plan, + want_create_pour_point=st.session_state.want_create_pour_point, + ) + plan = ensure_create_pour_point_before_define_domain( + plan, + conversation_text_for_plan_rules() + or user_prompt_for_metadata() + or s(st.session_state.get("nl_request", "")), + data_dir=SYMFLUENCE_DATA_DIR, + ) + sorted_exec_steps = sort_plan_steps_by_workflow_order(list(plan.get("steps") or [])) + if sorted_exec_steps: + plan["steps"] = sorted_exec_steps + st.session_state["_committed_plan_steps"] = list(plan.get("steps") or []) store_run_plan(plan) plan_cfg = (plan or {}).get("config", {}) or {} spec_dict = build_spec_dict(plan_cfg) @@ -7283,22 +7330,16 @@ def show_text(text: str): else: cfg.pop("BOUNDING_BOX_COORDS", None) + dump_yaml(cfg, out_yaml) + + write_run_metadata_files(outdir, plan, spec_dict) + user_request = user_prompt_for_metadata() or s(st.session_state.get("nl_request", "")) basin_domain = symfluence_domain_name( s(plan_cfg.get("domain_name")) or s(st.session_state.domain_name), s(plan_cfg.get("experiment_id")) or s(st.session_state.experiment_id), ) - symfluence_domain = basin_domain - plan = ensure_skip_acquire_forcings_when_local_forcing( - plan, - user_request, - data_dir=SYMFLUENCE_DATA_DIR, - ) - plan = ensure_skip_model_agnostic_when_local_preprocessing( - plan, - user_request, - data_dir=SYMFLUENCE_DATA_DIR, - ) + symfluence_domain = s(cfg.get("DOMAIN_NAME")) or basin_domain plan = ensure_skip_process_observed_when_local_streamflow( plan, user_request, @@ -7306,22 +7347,7 @@ def show_text(text: str): symfluence_domain=symfluence_domain, ) store_run_plan(plan) - plan_cfg = (plan or {}).get("config", {}) or {} - - station_id = resolve_station_id_from_plan(plan_cfg, user_request, fallback=s(st.session_state.station_id)) - if station_id: - cfg["STATION_ID"] = station_id - if plan_uses_local_data(plan_cfg, plan.get("steps") or [], user_request, data_dir=SYMFLUENCE_DATA_DIR) or ( - symfluence_domain and domain_has_local_streamflow(SYMFLUENCE_DATA_DIR, symfluence_domain) - ): - cfg["DOWNLOAD_WSC_DATA"] = False - dump_yaml(cfg, out_yaml) - - # plan.json must match skip-adjusted steps (ensure_skip_* helpers run above). - write_run_metadata_files(outdir, plan, spec_dict) - - symfluence_domain = s(cfg.get("DOMAIN_NAME")) or basin_domain steps = plan.get("steps", []) or [] pour_coords = ( diff --git a/prompts/planner_prompt.txt b/prompts/planner_prompt.txt index fd69069..9d26205 100644 --- a/prompts/planner_prompt.txt +++ b/prompts/planner_prompt.txt @@ -232,6 +232,7 @@ Rules: - calibrate_model - postprocess_results when all required inputs are provided and the user requested those steps. + If the user asks to delineate (domain_def delineate) and provides pour_point_coords, always include create_pour_point before define_domain. 53. steps must be an ordered list of workflow step names only. diff --git a/server/core/local_domain.py b/server/core/local_domain.py index 0485c9c..ba649e4 100644 --- a/server/core/local_domain.py +++ b/server/core/local_domain.py @@ -122,7 +122,7 @@ def seed_mac_duplicate_domain_from_basin( basin_domain: str, symfluence_domain: str, ) -> list[str]: - """When DOMAIN_NAME has a Mac-style (n) suffix, seed artifacts from the base basin domain.""" + """When DOMAIN_NAME is a duplicate (``Basin (2)`` or ``Basin_2``), seed from the base basin.""" if not basin_domain or not symfluence_domain or basin_domain == symfluence_domain: return [] from server.core.plan_rules import domain_has_complete_local_workflow, domain_has_local_summa_forcing diff --git a/server/core/plan_rules.py b/server/core/plan_rules.py index 970281d..f2dbca3 100644 --- a/server/core/plan_rules.py +++ b/server/core/plan_rules.py @@ -263,6 +263,23 @@ def domain_has_local_dem(data_dir: str | Path | None, domain_name: str) -> bool: return domain_dem_path(data_dir, domain_name).is_file() +def domain_pour_point_shapefile_path(data_dir: str | Path, domain_name: str) -> Path: + name = _s(domain_name) + return ( + Path(data_dir) + / f"domain_{name}" + / "shapefiles" + / "pour_point" + / f"{name}_pourPoint.shp" + ) + + +def domain_has_local_pour_point(data_dir: str | Path | None, domain_name: str) -> bool: + if not data_dir or not _s(domain_name): + return False + return domain_pour_point_shapefile_path(data_dir, domain_name).is_file() + + def domain_attributes_root(data_dir: str | Path, domain_name: str) -> Path: return Path(data_dir) / f"domain_{_s(domain_name)}" / "data" / "attributes" @@ -1457,6 +1474,37 @@ def merge_step_dependencies_preserving_order(steps: list[str], catalog: dict) -> ) +def ensure_create_pour_point_before_define_domain( + plan: dict, + user_request: str = "", + *, + data_dir: str | Path | None = None, +) -> dict: + """define_domain needs a pour-point shapefile; insert create_pour_point when omitted.""" + if not isinstance(plan, dict): + return plan + out = dict(plan) + steps = list(out.get("steps") or []) + if "define_domain" not in steps or "create_pour_point" in steps: + return out + if _user_forbids_step_inference(user_request, "create_pour_point"): + return out + + cfg = dict(out.get("config") or {}) + domain_name = _s(cfg.get("domain_name")) + if domain_name and data_dir and domain_has_local_pour_point(data_dir, domain_name): + return out + + idx = steps.index("define_domain") + steps.insert(idx, "create_pour_point") + out["steps"] = steps + _append_plan_note( + out, + "Inserted create_pour_point before define_domain (pour-point shapefile required for delineation).", + ) + return out + + def ensure_acquire_attributes_before_define_domain( plan: dict, user_request: str = "", @@ -1494,6 +1542,56 @@ def ensure_acquire_attributes_before_define_domain( return out +def ensure_discretize_domain_before_preprocessing( + plan: dict, + user_request: str = "", + *, + data_dir: str | Path | None = None, +) -> dict: + """model_agnostic_preprocessing needs HRU/GRU shapefiles from discretize_domain.""" + if not isinstance(plan, dict): + return plan + out = dict(plan) + steps = list(out.get("steps") or []) + if "discretize_domain" in steps: + return out + if _user_forbids_step_inference(user_request, "discretize_domain"): + return out + + needs_hrus = bool( + { + "model_agnostic_preprocessing", + "model_specific_preprocessing", + "run_model", + } + & set(steps) + ) + cfg = dict(out.get("config") or {}) + delineated = "define_domain" in steps or _s(cfg.get("domain_def")).lower() == "delineate" + if not needs_hrus or not delineated: + return out + + domain_name = _s(cfg.get("domain_name")) + experiment_id = _s(cfg.get("experiment_id")) or "run_1" + if domain_name and data_dir and domain_has_local_discretization( + data_dir, domain_name, experiment_id + ): + return out + + if "define_domain" in steps: + steps.insert(steps.index("define_domain") + 1, "discretize_domain") + elif "model_agnostic_preprocessing" in steps: + steps.insert(steps.index("model_agnostic_preprocessing"), "discretize_domain") + else: + steps.append("discretize_domain") + out["steps"] = sort_plan_steps_by_workflow_order(steps) + _append_plan_note( + out, + "Inserted discretize_domain before preprocessing (HRU/GRU shapefiles required).", + ) + return out + + def ensure_cloud_data_access_for_acquire_steps( plan: dict, user_request: str = "", @@ -1556,6 +1654,7 @@ def ensure_online_data_when_missing( if user_forbids_download_step(user_request, "acquire_attributes"): return out + out = ensure_create_pour_point_before_define_domain(out, user_request, data_dir=data_dir) out = ensure_acquire_attributes_before_define_domain(out, user_request, data_dir=data_dir) steps = list(out.get("steps") or []) @@ -1872,6 +1971,9 @@ def normalize_local_workflow_plan( steps = list(out.get("steps") or []) out = apply_explicit_step_constraints(out, user_request) + out = ensure_create_pour_point_before_define_domain(out, user_request, data_dir=data_dir) + out = ensure_acquire_attributes_before_define_domain(out, user_request, data_dir=data_dir) + out = ensure_discretize_domain_before_preprocessing(out, user_request, data_dir=data_dir) steps = list(out.get("steps") or []) cfg = dict(out.get("config") or {}) @@ -2008,10 +2110,14 @@ def add(step: str) -> None: if not forbids_model and ( re.search(r"\bfrom\s+scratch\b", text) or re.search(r"\brun\s+(?:the\s+)?model\b", text) + or re.search(r"\bmodel\s+run\b", text) or re.search(r"\b(?:and|then)\s+(?:the\s+)?model\b", text) or re.search(r"\brun\s+summa\b", text) or re.search(r"\bsumma\s+workflow\b", text) - or (re.search(r"\bworkflow\b", text) and re.search(r"\bsumma\b", text)) + or ( + re.search(r"\bworkflow\b", text) + and re.search(r"\b(?:fuse|summa|hbv|gr|mesh|hype|ngen|topmodel)\b", text) + ) ): add("run_model") @@ -2030,6 +2136,7 @@ def add(step: str) -> None: if re.search(r"\bpreprocess", text): add("model_agnostic_preprocessing") add("model_specific_preprocessing") + add("discretize_domain") if re.search(r"\bsetup\b", text): add("setup_project") @@ -2058,6 +2165,16 @@ def add(step: str) -> None: ): add("define_domain") + if "define_domain" in goals and not _user_forbids_step_inference( + user_request, "create_pour_point" + ): + add("create_pour_point") + + if "define_domain" in goals and ( + "model_agnostic_preprocessing" in goals or "run_model" in goals + ): + add("discretize_domain") + for step in extract_steps_from_request(user_request, WORKFLOW_STEP_NAMES): if step == "run_model" and forbids_model: continue diff --git a/server/core/run_naming.py b/server/core/run_naming.py index 24ee8e7..b044e98 100644 --- a/server/core/run_naming.py +++ b/server/core/run_naming.py @@ -11,6 +11,7 @@ PLACEHOLDER_EXPERIMENT = "exp" _MAC_DUP_SUFFIX_RE = re.compile(r" \((\d+)\)$") +_UNDERSCORE_DUP_SUFFIX_RE = re.compile(r"_(\d+)$") def sanitize_config_token(name: str) -> str: @@ -35,8 +36,22 @@ def parse_mac_duplicate_suffix(name: str) -> tuple[str, int | None]: return base, int(match.group(1)) +def duplicate_symfluence_domain(basin_domain: str, suffix_n: int) -> str: + """Space-free SYMFLUENCE DOMAIN_NAME for duplicate run ``n`` (TauDEM cannot use spaces).""" + basin = sanitize_config_token(basin_domain) or str(basin_domain or "").strip() + return f"{basin}_{int(suffix_n)}" + + def symfluence_domain_mac_suffix(symfluence_domain: str) -> tuple[str, int | None]: - return parse_mac_duplicate_suffix(str(symfluence_domain or "").strip()) + """Split ``BowRiver4 (2)`` or ``BowRiver4_2`` into ``(BowRiver4, 2)``.""" + text = str(symfluence_domain or "").strip() + base, suffix_n = parse_mac_duplicate_suffix(text) + if suffix_n is not None: + return base, suffix_n + match = _UNDERSCORE_DUP_SUFFIX_RE.search(text) + if not match: + return text, None + return text[: match.start()], int(match.group(1)) def symfluence_domain_for_run_folder( @@ -52,7 +67,7 @@ def symfluence_domain_for_run_folder( return basin base_part, suffix_n = parse_mac_duplicate_suffix(run_folder) if base_part == base_run and suffix_n is not None: - return f"{basin} ({suffix_n})" + return duplicate_symfluence_domain(basin, suffix_n) return basin @@ -190,7 +205,7 @@ def _allocate_unique_run_folder_excluding( for n in range(1, 1000): candidate = f"{base_run} ({n})" - sym_domain = f"{basin} ({n})" + sym_domain = duplicate_symfluence_domain(basin, n) if not is_run_workspace_taken( candidate, sym_domain, @@ -283,7 +298,7 @@ def allocate_unique_run_folder( for n in range(1, 1000): candidate = f"{base_run} ({n})" - sym_domain = f"{basin} ({n})" + sym_domain = duplicate_symfluence_domain(basin, n) if not is_run_workspace_taken( candidate, sym_domain, diff --git a/server/llm/claude_provider.py b/server/llm/claude_provider.py index e728a84..d9ec450 100644 --- a/server/llm/claude_provider.py +++ b/server/llm/claude_provider.py @@ -3,6 +3,7 @@ import json import os import re +from copy import deepcopy from pathlib import Path from typing import Any, Dict, Tuple @@ -21,6 +22,58 @@ ) +def _model_supports_temperature(model: str) -> bool: + """Newer Claude models reject non-default sampling parameters.""" + model_id = (model or "").lower() + deprecated_sampling_markers = ( + "claude-sonnet-5", + "claude-fable-5", + "claude-opus-4-7", + "claude-opus-4-8", + ) + return not any(marker in model_id for marker in deprecated_sampling_markers) + + +def _model_prefers_thinking_disabled(model: str) -> bool: + """Adaptive thinking is on by default on newer models and can exhaust output budget.""" + model_id = (model or "").lower() + adaptive_thinking_markers = ( + "claude-sonnet-5", + "claude-fable-5", + "claude-opus-4-7", + "claude-opus-4-8", + ) + return any(marker in model_id for marker in adaptive_thinking_markers) + + +def _sanitize_json_schema(node: Any) -> Any: + """Adapt JSON Schema for Anthropic structured output (union types → anyOf).""" + if isinstance(node, list): + return [_sanitize_json_schema(item) for item in node] + if not isinstance(node, dict): + return node + + out: Dict[str, Any] = {} + for key, value in node.items(): + if key == "type" and isinstance(value, list): + non_null = [item for item in value if item != "null"] + if "null" in value: + if len(non_null) == 1: + out["anyOf"] = [{"type": non_null[0]}, {"type": "null"}] + elif non_null: + out["anyOf"] = [{"type": item} for item in non_null] + [{"type": "null"}] + else: + out["type"] = "null" + else: + out["type"] = value + continue + if key == "enum" and isinstance(value, list): + out[key] = [item if item is not None else None for item in value] + continue + out[key] = _sanitize_json_schema(value) + return out + + class ClaudeProvider: def __init__(self, api_key: str | None = None): api_key = api_key or os.getenv("ANTHROPIC_API_KEY") or os.getenv("CLAUDE_API_KEY") @@ -46,65 +99,141 @@ def _loads_json_from_text(self, text: str) -> Dict[str, Any] | None: except Exception: return None - def _response_text(self, response) -> str: + def _response_text_chunks(self, response) -> list[str]: chunks: list[str] = [] for block in getattr(response, "content", []) or []: if getattr(block, "type", None) == "text": - chunks.append(getattr(block, "text", "") or "") - return "".join(chunks).strip() + text = getattr(block, "text", "") or "" + if isinstance(text, str) and text.strip(): + chunks.append(text) + return chunks - def _call_json_schema( + def _extract_json_from_response(self, response) -> Dict[str, Any] | None: + parsed_output = getattr(response, "parsed_output", None) + if isinstance(parsed_output, dict): + return parsed_output + + for block in getattr(response, "content", []) or []: + block_parsed = getattr(block, "parsed_output", None) + if isinstance(block_parsed, dict): + return block_parsed + + for chunk in self._response_text_chunks(response): + loaded = self._loads_json_from_text(chunk) + if loaded is not None: + return loaded + return None + + def _describe_response_failure(self, response) -> str: + details: list[str] = [] + stop_reason = getattr(response, "stop_reason", None) + if stop_reason: + details.append(f"stop_reason={stop_reason}") + + blocks = getattr(response, "content", []) or [] + if not blocks: + details.append("no content blocks") + else: + block_types = [getattr(block, "type", "?") for block in blocks] + details.append(f"content_types={','.join(block_types)}") + + text_preview = " ".join(self._response_text_chunks(response)) + if text_preview: + preview = text_preview[:240].replace("\n", " ") + details.append(f"text preview: {preview!r}") + else: + details.append("empty response text") + + return "; ".join(details) + + def _base_message_kwargs( self, *, model: str, - schema: Dict[str, Any], + max_output_tokens: int, system_prompt: str, user_prompt: str, - max_output_tokens: int, ) -> Dict[str, Any]: - create_kwargs: Dict[str, Any] = { + kwargs: Dict[str, Any] = { "model": model, "max_tokens": max_output_tokens, - "temperature": 0.2, "system": system_prompt, "messages": [{"role": "user", "content": user_prompt}], } + if _model_supports_temperature(model): + kwargs["temperature"] = 0.2 + if _model_prefers_thinking_disabled(model): + kwargs["thinking"] = {"type": "disabled"} + return kwargs - try: - response = self.client.messages.create( - **create_kwargs, - output_config={ - "format": { - "type": "json_schema", - "schema": schema, - } - }, - ) - parsed = self._loads_json_from_text(self._response_text(response)) - if parsed is not None: - return parsed - except TypeError: - pass - except Exception: - pass + def _call_json_schema( + self, + *, + model: str, + schema: Dict[str, Any], + system_prompt: str, + user_prompt: str, + max_output_tokens: int, + ) -> Dict[str, Any]: + claude_schema = _sanitize_json_schema(deepcopy(schema)) + token_budgets = [max_output_tokens] + if max_output_tokens < 8192: + token_budgets.append(8192) + last_failure = "unknown error" fallback_system = ( system_prompt + "\n\nIMPORTANT: Return ONLY a single valid JSON object. " "No numbering, no comments, no markdown, no trailing commas." ) - response2 = self.client.messages.create( - model=model, - max_tokens=max_output_tokens, - temperature=0.2, - system=fallback_system, - messages=[{"role": "user", "content": user_prompt}], - ) - parsed2 = self._loads_json_from_text(self._response_text(response2)) - if parsed2 is not None: - return parsed2 - raise RuntimeError("No JSON returned from Claude after schema + fallback attempt.") + for attempt_idx, token_budget in enumerate(token_budgets): + create_kwargs = self._base_message_kwargs( + model=model, + max_output_tokens=token_budget, + system_prompt=system_prompt, + user_prompt=user_prompt, + ) + + try: + response = self.client.messages.create( + **create_kwargs, + output_config={ + "format": { + "type": "json_schema", + "schema": claude_schema, + } + }, + ) + parsed = self._extract_json_from_response(response) + if parsed is not None: + return parsed + last_failure = self._describe_response_failure(response) + except TypeError as exc: + last_failure = f"structured output unsupported: {exc}" + except Exception as exc: + last_failure = f"structured output error: {type(exc).__name__}: {exc}" + + response2 = self.client.messages.create( + **self._base_message_kwargs( + model=model, + max_output_tokens=token_budget, + system_prompt=fallback_system, + user_prompt=user_prompt, + ), + ) + parsed2 = self._extract_json_from_response(response2) + if parsed2 is not None: + return parsed2 + last_failure = self._describe_response_failure(response2) + + if attempt_idx + 1 < len(token_budgets): + continue + + raise RuntimeError( + "No JSON returned from Claude after schema + fallback attempt " + f"({last_failure})." + ) def generate_run_plan(self, *, model: str, user_request: str) -> Dict[str, Any]: planner_prompt = PLANNER_PROMPT_PATH.read_text(encoding="utf-8") @@ -115,7 +244,7 @@ def generate_run_plan(self, *, model: str, user_request: str) -> Dict[str, Any]: schema=schema, system_prompt=planner_prompt, user_prompt=user_request, - max_output_tokens=4096, + max_output_tokens=8192, ) return finalize_run_plan(plan, user_request) @@ -143,7 +272,7 @@ def refine_run_plan( schema=schema, system_prompt=refinement_prompt, user_prompt=user_prompt, - max_output_tokens=4096, + max_output_tokens=8192, ) return finalize_plan_refinement( result, diff --git a/server/llm/gemini_provider.py b/server/llm/gemini_provider.py index 3d80c0b..3e028be 100644 --- a/server/llm/gemini_provider.py +++ b/server/llm/gemini_provider.py @@ -4,6 +4,7 @@ import os import re import tempfile +from copy import deepcopy from pathlib import Path from typing import Any, Dict @@ -23,6 +24,34 @@ ) +def _sanitize_schema_for_gemini(node: Any) -> Any: + """Adapt JSON Schema for Gemini structured output (union types → anyOf).""" + if isinstance(node, list): + return [_sanitize_schema_for_gemini(item) for item in node] + if not isinstance(node, dict): + return node + + out: Dict[str, Any] = {} + for key, value in node.items(): + if key == "type" and isinstance(value, list): + non_null = [item for item in value if item != "null"] + if "null" in value: + if len(non_null) == 1: + out["anyOf"] = [{"type": non_null[0]}, {"type": "null"}] + elif non_null: + out["anyOf"] = [{"type": item} for item in non_null] + [{"type": "null"}] + else: + out["type"] = "null" + else: + out["type"] = value + continue + if key == "enum" and isinstance(value, list): + out[key] = [item if item is not None else None for item in value] + continue + out[key] = _sanitize_schema_for_gemini(value) + return out + + class GeminiProvider: def __init__(self, api_key: str | None = None): api_key = api_key or os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY") @@ -48,6 +77,65 @@ def _loads_json_from_text(self, text: str) -> Dict[str, Any] | None: except Exception: return None + def _response_text_chunks(self, resp) -> list[str]: + chunks: list[str] = [] + text = getattr(resp, "text", None) + if isinstance(text, str) and text.strip(): + chunks.append(text) + + for candidate in getattr(resp, "candidates", None) or []: + content = getattr(candidate, "content", None) + for part in getattr(content, "parts", None) or []: + part_text = getattr(part, "text", None) + if isinstance(part_text, str) and part_text.strip(): + chunks.append(part_text) + + return chunks + + def _extract_json_from_response(self, resp) -> Dict[str, Any] | None: + parsed = getattr(resp, "parsed", None) + if isinstance(parsed, dict): + return parsed + + for chunk in self._response_text_chunks(resp): + loaded = self._loads_json_from_text(chunk) + if loaded is not None: + return loaded + return None + + def _describe_response_failure(self, resp) -> str: + details: list[str] = [] + prompt_feedback = getattr(resp, "prompt_feedback", None) + if prompt_feedback is not None: + block_reason = getattr(prompt_feedback, "block_reason", None) + if block_reason: + details.append(f"prompt blocked ({block_reason})") + + candidates = getattr(resp, "candidates", None) or [] + if not candidates: + details.append("no candidates returned") + for idx, candidate in enumerate(candidates): + finish_reason = getattr(candidate, "finish_reason", None) + if finish_reason: + details.append(f"candidate[{idx}] finish_reason={finish_reason}") + safety = getattr(candidate, "safety_ratings", None) or [] + blocked = [ + f"{getattr(r, 'category', '?')}={getattr(r, 'probability', '?')}" + for r in safety + if str(getattr(r, "probability", "")).upper() in {"HIGH", "MEDIUM"} + ] + if blocked: + details.append(f"candidate[{idx}] safety={', '.join(blocked)}") + + text_preview = " ".join(self._response_text_chunks(resp)) + if text_preview: + preview = text_preview[:240].replace("\n", " ") + details.append(f"text preview: {preview!r}") + else: + details.append("empty response text") + + return "; ".join(details) if details else "unknown Gemini response failure" + def _call_json_schema( self, *, @@ -57,48 +145,57 @@ def _call_json_schema( user_prompt: str, max_output_tokens: int, ) -> Dict[str, Any]: - contents = [ - types.Content( - role="user", - parts=[types.Part(text=f"{system_prompt}\n\nUser request:\n{user_prompt}")], - ) - ] - config = types.GenerateContentConfig( - temperature=0.2, - max_output_tokens=max_output_tokens, - response_mime_type="application/json", - response_json_schema=schema, - ) + gemini_schema = _sanitize_schema_for_gemini(deepcopy(schema)) + token_budgets = [max_output_tokens] + if max_output_tokens < 8192: + token_budgets.append(8192) - resp = self.client.models.generate_content( - model=model, - contents=contents, - config=config, - ) - parsed = self._loads_json_from_text(getattr(resp, "text", "") or "") - if parsed is not None: - return parsed + last_failure = "unknown error" + for attempt_idx, token_budget in enumerate(token_budgets): + config = types.GenerateContentConfig( + temperature=0.2, + max_output_tokens=token_budget, + system_instruction=system_prompt, + response_mime_type="application/json", + response_json_schema=gemini_schema, + ) + resp = self.client.models.generate_content( + model=model, + contents=user_prompt, + config=config, + ) + parsed = self._extract_json_from_response(resp) + if parsed is not None: + return parsed + last_failure = self._describe_response_failure(resp) - fallback_prompt = ( - f"{system_prompt}\n\n" - "IMPORTANT: Return ONLY a single valid JSON object. " - "No numbering, no comments, no markdown, no trailing commas.\n\n" - f"User request:\n{user_prompt}" - ) - resp2 = self.client.models.generate_content( - model=model, - contents=fallback_prompt, - config=types.GenerateContentConfig( + fallback_config = types.GenerateContentConfig( temperature=0.2, - max_output_tokens=max_output_tokens, + max_output_tokens=token_budget, + system_instruction=( + system_prompt + + "\n\nIMPORTANT: Return ONLY a single valid JSON object. " + "No numbering, no comments, no markdown, no trailing commas." + ), response_mime_type="application/json", - ), - ) - parsed2 = self._loads_json_from_text(getattr(resp2, "text", "") or "") - if parsed2 is not None: - return parsed2 + ) + resp2 = self.client.models.generate_content( + model=model, + contents=user_prompt, + config=fallback_config, + ) + parsed2 = self._extract_json_from_response(resp2) + if parsed2 is not None: + return parsed2 + last_failure = self._describe_response_failure(resp2) + + if attempt_idx + 1 < len(token_budgets): + continue - raise RuntimeError("No JSON returned from Gemini after schema + fallback attempt.") + raise RuntimeError( + "No JSON returned from Gemini after schema + fallback attempt " + f"({last_failure})." + ) def generate_run_plan(self, *, model: str, user_request: str) -> Dict[str, Any]: planner_prompt = PLANNER_PROMPT_PATH.read_text(encoding="utf-8") @@ -109,7 +206,7 @@ def generate_run_plan(self, *, model: str, user_request: str) -> Dict[str, Any]: schema=schema, system_prompt=planner_prompt, user_prompt=user_request, - max_output_tokens=2500, + max_output_tokens=8192, ) return finalize_run_plan(plan, user_request) @@ -137,7 +234,7 @@ def refine_run_plan( schema=schema, system_prompt=refinement_prompt, user_prompt=user_prompt, - max_output_tokens=3000, + max_output_tokens=8192, ) return finalize_plan_refinement( result, diff --git a/tests/test_claude_provider.py b/tests/test_claude_provider.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/test_claude_provider.py @@ -0,0 +1 @@ + diff --git a/tests/test_gemini_provider.py b/tests/test_gemini_provider.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/test_gemini_provider.py @@ -0,0 +1 @@ + diff --git a/tests/test_plan_resolve_dependencies.py b/tests/test_plan_resolve_dependencies.py index 0683df5..b6a43fa 100644 --- a/tests/test_plan_resolve_dependencies.py +++ b/tests/test_plan_resolve_dependencies.py @@ -90,4 +90,6 @@ def test_resolve_plan_step_dependencies_expands_minimal_cloud_plan( assert "validate_config" in steps assert "setup_project" in steps assert "run_model" in steps + assert "create_pour_point" in steps + assert steps.index("create_pour_point") < steps.index("define_domain") assert len(steps) > 3 diff --git a/tests/test_run_naming.py b/tests/test_run_naming.py index 91b46c3..dbd7a3d 100644 --- a/tests/test_run_naming.py +++ b/tests/test_run_naming.py @@ -1,93 +1,41 @@ -from __future__ import annotations - -import json -from pathlib import Path - from server.core.run_naming import ( - is_placeholder_assistant_run_folder, - placeholder_run_needs_rename, - rename_placeholder_assistant_run_folder, - resolve_run_workspace, + allocate_unique_run_folder, + duplicate_symfluence_domain, + run_folder_for_symfluence_domain, + symfluence_domain_for_run_folder, + symfluence_domain_mac_suffix, ) -def test_is_placeholder_assistant_run_folder_detects_domain_experiment_variants() -> None: - assert is_placeholder_assistant_run_folder("domain_exp001", "exp001") - assert is_placeholder_assistant_run_folder("domain_exp001 (11)", "exp001") - assert is_placeholder_assistant_run_folder("domain_exp", "exp001") - assert not is_placeholder_assistant_run_folder("BowRiver3_exp001", "exp001") - +def test_duplicate_symfluence_domain_has_no_spaces(): + assert duplicate_symfluence_domain("BowRiver4", 2) == "BowRiver4_2" + assert " " not in duplicate_symfluence_domain("BowRiver4", 1) -def test_placeholder_run_needs_rename_when_real_domain_arrives() -> None: - assert placeholder_run_needs_rename("domain_exp001 (11)", "BowRiver3", "exp001") - assert not placeholder_run_needs_rename("BowRiver3_exp001", "BowRiver3", "exp001") - assert not placeholder_run_needs_rename("domain_exp001", "domain", "exp001") - -def test_rename_placeholder_assistant_run_folder_moves_workspace(tmp_path: Path) -> None: - runs_dir = tmp_path / "runs" - data_dir = tmp_path / "data" - old_name = "domain_exp001 (11)" - old_dir = runs_dir / old_name - old_dir.mkdir(parents=True) - (old_dir / "plan.json").write_text(json.dumps({"config": {}}), encoding="utf-8") - - new_name = rename_placeholder_assistant_run_folder( - old_name, - "BowRiver3", - "exp001", - runs_dir=runs_dir, - data_dir=data_dir, +def test_run_folder_maps_to_space_free_domain(): + assert ( + symfluence_domain_for_run_folder("BowRiver4_exp061 (2)", "BowRiver4", "exp061") + == "BowRiver4_2" ) - - assert new_name == "BowRiver3_exp001" - assert not old_dir.exists() - assert (runs_dir / "BowRiver3_exp001" / "plan.json").is_file() - - -def test_rename_skips_when_target_workspace_already_exists(tmp_path: Path) -> None: - runs_dir = tmp_path / "runs" - data_dir = tmp_path / "data" - old_name = "domain_exp001" - old_dir = runs_dir / old_name - old_dir.mkdir(parents=True) - (old_dir / "plan.json").write_text("{}", encoding="utf-8") - - existing = runs_dir / "BowRiver3_exp001" - existing.mkdir() - (existing / "plan.json").write_text("{}", encoding="utf-8") - - new_name = rename_placeholder_assistant_run_folder( - old_name, - "BowRiver3", - "exp001", - runs_dir=runs_dir, - data_dir=data_dir, + assert ( + run_folder_for_symfluence_domain("BowRiver4_2", "exp061") + == "BowRiver4_exp061 (2)" ) - assert new_name == "BowRiver3_exp001 (1)" - assert not old_dir.exists() - assert (runs_dir / "BowRiver3_exp001 (1)" / "plan.json").is_file() - assert (existing / "plan.json").is_file() +def test_legacy_spaced_domain_suffix_still_parses(): + assert symfluence_domain_mac_suffix("BowRiver4 (2)") == ("BowRiver4", 2) + assert symfluence_domain_mac_suffix("BowRiver4_2") == ("BowRiver4", 2) -def test_resolve_run_workspace_renames_established_placeholder(tmp_path: Path) -> None: + +def test_allocate_unique_run_folder_uses_space_free_domain(tmp_path): runs_dir = tmp_path / "runs" data_dir = tmp_path / "data" - old_name = "domain_exp001 (11)" - old_dir = runs_dir / old_name - old_dir.mkdir(parents=True) - (old_dir / "plan.json").write_text("{}", encoding="utf-8") - - run_folder, sym_domain = resolve_run_workspace( - "BowRiver3", - "exp001", - old_name, - runs_dir=runs_dir, - data_dir=data_dir, - workspace_locked=True, - ) - - assert run_folder == "BowRiver3_exp001" - assert sym_domain == "BowRiver3" - assert (runs_dir / "BowRiver3_exp001" / "plan.json").is_file() + runs_dir.mkdir() + data_dir.mkdir() + (runs_dir / "BowRiver4_exp061").mkdir() + folder = allocate_unique_run_folder("BowRiver4", "exp061", runs_dir, data_dir) + assert folder == "BowRiver4_exp061 (1)" + domain = symfluence_domain_for_run_folder(folder, "BowRiver4", "exp061") + assert domain == "BowRiver4_1" + assert " " not in domain