yyyc's picture
fist push the app
5194558
Raw
History Blame Contribute Delete
8.2 kB
"""LLM-backed agent and deterministic weather briefing helpers."""
from __future__ import annotations
import json
import os
import re
from dataclasses import dataclass
from typing import Any
import pandas as pd
from weatherpred.prompts import AGENT_RESPONSE_PROMPT, AGENT_SYSTEM_PROMPT
from weatherpred.tools import clamp_horizon, normalize_variable, run_tool
@dataclass
class AgentRun:
response: str
action: dict[str, Any]
tool_result: dict[str, Any]
comparison: pd.DataFrame | None = None
def has_gemini_key() -> bool:
return bool(os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY"))
def make_rule_based_brief(comparison: pd.DataFrame, summary: dict[str, Any]) -> str:
warmest = comparison.sort_values("model_high_c", ascending=False).iloc[0]
rainiest = comparison.sort_values("api_rain_mm", ascending=False).iloc[0]
confidence = round(summary["avg_model_confidence"] * 100)
return (
f"The ONNX model expects the warmest day around {warmest['date']} "
f"at {warmest['model_high_c']:.1f}C, while Open-Meteo peaks at "
f"{summary['max_api_high_c']:.1f}C. API rainfall is highest on "
f"{rainiest['date']} at {rainiest['api_rain_mm']:.1f} mm. "
f"The model confidence proxy averages {confidence}%, based on historical "
"seasonal variability rather than a calibrated meteorological probability."
)
def make_gemini_brief(comparison: pd.DataFrame, summary: dict[str, Any]) -> str:
if not has_gemini_key():
return make_rule_based_brief(comparison, summary)
try:
from google import genai
except ImportError:
return make_rule_based_brief(comparison, summary)
prompt = (
"Write a concise weather analyst briefing for Tokyo. Compare an ONNX "
"model forecast against Open-Meteo API forecast. Mention uncertainty and "
"avoid claiming this model is production-grade.\n\n"
f"Summary: {json.dumps(summary)}\n"
f"Forecast rows: {comparison.to_json(orient='records')}"
)
try:
client = genai.Client()
response = client.models.generate_content(
model=os.getenv("GEMINI_MODEL", "gemini-2.5-flash"),
contents=prompt,
)
return response.text or make_rule_based_brief(comparison, summary)
except Exception:
return make_rule_based_brief(comparison, summary)
def extract_json_object(text: str) -> dict[str, Any] | None:
try:
return json.loads(text)
except json.JSONDecodeError:
pass
match = re.search(r"\{.*\}", text, flags=re.DOTALL)
if not match:
return None
try:
return json.loads(match.group(0))
except json.JSONDecodeError:
return None
def fallback_action(user_message: str) -> dict[str, Any]:
text = user_message.lower()
horizon_match = re.search(r"(?:next|for)\s+(\d+)\s*(?:day|days)?", text)
horizon = clamp_horizon(horizon_match.group(1) if horizon_match else None)
if any(word in text for word in ["rain", "precip", "precipitation", "shower"]):
variable = "rain"
elif "wind" in text:
variable = "wind"
elif any(word in text for word in ["temperature", "temp", "hot", "cold", "high", "low"]):
variable = "temperature"
else:
variable = "all"
if any(word in text for word in ["season", "seasonal", "cycle", "cyclic", "monthly", "summer", "winter"]):
tool = "analyze_seasonality"
elif any(word in text for word in ["trend", "past", "historical", "history", "30 years", "long-term"]):
tool = "analyze_historical_trend"
elif any(word in text for word in ["model", "onnx", "confidence", "method", "how", "source", "data"]):
tool = "explain_model"
elif any(word in text for word in ["compare", "difference", "versus", "vs"]):
tool = "compare_model_vs_open_meteo"
elif horizon_match or any(word in text for word in ["predict", "forecast", "next"]):
tool = "forecast_next_days"
else:
tool = "show_dashboard_view"
return {"tool": tool, "variable": variable, "horizon_days": horizon, "reason": "fallback parser"}
def infer_action(user_message: str) -> dict[str, Any]:
if not has_gemini_key():
return fallback_action(user_message)
try:
from google import genai
client = genai.Client()
prompt = f"{AGENT_SYSTEM_PROMPT}\n\nUser request: {user_message}"
response = client.models.generate_content(
model=os.getenv("GEMINI_MODEL", "gemini-2.5-flash"),
contents=prompt,
)
action = extract_json_object(response.text or "")
if action and "tool" in action:
action["variable"] = normalize_variable(action.get("variable"))
action["horizon_days"] = clamp_horizon(action.get("horizon_days"))
return action
except Exception:
pass
return fallback_action(user_message)
def deterministic_agent_response(action: dict[str, Any], tool_result: dict[str, Any]) -> str:
tool = tool_result.get("tool", action.get("tool"))
summary = tool_result.get("summary", {})
variable = tool_result.get("variable", action.get("variable", "all"))
if tool == "forecast_next_days":
return (
f"Showing the next {tool_result['horizon_days']} days for {variable}. "
"Open-Meteo is shown beside the local model. Note: ONNX directly predicts five-day max temperature; "
"longer horizons, rain, wind, and low temperature use historical seasonal statistics."
)
if tool == "compare_model_vs_open_meteo":
return f"Model vs Open-Meteo comparison for {variable}: {json.dumps(summary.get('average_absolute_differences', {}))}."
if tool == "analyze_historical_trend":
return (
f"Historical {variable} trend from {summary['history_start']} to {summary['history_end']}: "
f"the recent five-year mean is {summary['last_5_year_mean']} {summary['unit']}, "
f"versus {summary['first_5_year_mean']} {summary['unit']} in the first five years. "
f"Change: {summary['change_last_vs_first']} {summary['unit']}."
)
if tool == "analyze_seasonality":
return (
f"Seasonality for {variable}: peak month is {summary['peak_month']} "
f"at {summary['peak_value']} {summary['unit']}; lowest month is {summary['low_month']} "
f"at {summary['low_value']} {summary['unit']}. Seasonal amplitude is "
f"{summary['seasonal_amplitude']} {summary['unit']}."
)
if tool == "explain_model":
return " ".join(str(value) for value in summary.values())
if tool == "show_dashboard_view":
return f"{variable.title()} is available in {summary['view']}."
return "I handled the request with the weather dashboard tools."
def summarize_agent_response(user_message: str, action: dict[str, Any], tool_result: dict[str, Any]) -> str:
if not has_gemini_key():
return deterministic_agent_response(action, tool_result)
try:
from google import genai
safe_result = {key: value for key, value in tool_result.items() if key != "comparison"}
prompt = AGENT_RESPONSE_PROMPT.format(
user_message=user_message,
tool_result=json.dumps(safe_result, default=str),
)
client = genai.Client()
response = client.models.generate_content(
model=os.getenv("GEMINI_MODEL", "gemini-2.5-flash"),
contents=prompt,
)
return response.text or deterministic_agent_response(action, tool_result)
except Exception:
return deterministic_agent_response(action, tool_result)
def run_weather_agent(user_message: str) -> AgentRun:
action = infer_action(user_message)
tool_result = run_tool(action)
response = summarize_agent_response(user_message, action, tool_result)
comparison = tool_result.get("comparison")
if not isinstance(comparison, pd.DataFrame):
comparison = None
return AgentRun(
response=response,
action=action,
tool_result=tool_result,
comparison=comparison,
)