avfranco's picture
HF Space deploy snapshot (minimal allow-list)
d64fd55
Raw History Blame Contribute Delete
5.09 kB
from typing import List, Dict, Any, Optional
import logging
import json
import time
from observability import logger as obs_logger
from observability import components as obs_components
from agents.base import BaseAgent
from domain.training.charts import ChartSpec, ChartType
from llm.base import LLMClient
from router.models import RouteDecision
from router.compact import compact_raw_context
from pydantic import BaseModel, Field
logger = logging.getLogger(__name__)
class VisualizationAgentOutput(BaseModel):
"""The structured output from the VisualizationAgent."""
charts: List[ChartSpec] = Field(description="The list of charts to generate.")
class VisualizationAgent(BaseAgent):
"""
Agent responsible for specifying visualizations from run data using LLM.
Returns ChartSpec objects instead of executing code.
"""
def __init__(self, llm_client: LLMClient):
self.llm_client = llm_client
self.instruction = """You are a data visualization expert. Your goal is to specify charts to visualize run data.
Return a JSON object containing a 'charts' list. Each chart in the list must have:
- 'chart_type': One of ['pace', 'heart_rate', 'volume', 'frequency', 'zones', 'dynamic']
- 'title': A descriptive title for the chart.
- 'params': (Optional) A dictionary of parameters.
Chart Types:
- 'pace': Pace vs Date
- 'heart_rate': Heart Rate vs Date
- 'volume': Weekly mileage (distance)
- 'frequency': Number of runs per week
- 'zones': Distribution of Effort/HR zones
- 'dynamic': Custom chart with provided 'labels', 'values', and 'type' (bar/pie/line) in params.
When asked to generate charts:
1. If you need a custom chart (not pace/hr/volume/freq/zones), use 'dynamic' and specify labels and values in 'params'.
Example params for dynamic: {"labels": ["Mon", "Tue"], "values": [5, 10], "type": "bar"}
2. Consider the 'periods' (n_weeks, etc.) provided in the extracted intent.
3. Default to showing the last 4 weeks of data unless 'all_time' or a specific range is requested.
4. Always include 'pace' and 'heart_rate' charts by default if no specific query is provided.
"""
async def run(
self,
features: List[Dict[str, Any]],
query: str = "",
intent: Optional[RouteDecision] = None
) -> List[ChartSpec]:
with obs_logger.start_span("visualization_agent.run", obs_components.AGENT):
"""
Generates chart specifications using the LLM.
"""
intent_info = ""
if intent:
intent_info = f"\nExtracted Intent: Metric={intent.metric}, Period={intent.period}, Target Date={intent.target_date}"
if query:
# Use compact_raw_context for prompt injection as per user feedback
compacted_features = compact_raw_context({"features": features}, max_elements=25).get("features", [])
prompt = f"""
The user has requested a specific visualization: "{query}"{intent_info}
There are {len(features)} runs available in history.
**Relevant Data Context (Compacted):**
{json.dumps(compacted_features, default=str, indent=2)}
ACTION: Please specify the most appropriate chart(s) according to the schema.
If they ask for frequency or "by day", use 'frequency'.
"""
else:
prompt = f"""
Generate default visualizations for the provided {len(features)} runs.
Please request:
1. A pace chart ('pace')
2. A heart rate chart ('heart_rate')
"""
try:
res = await self.llm_client.generate(
prompt,
instruction=self.instruction,
schema=VisualizationAgentOutput,
name="visualization_agent",
)
logger.info(f"[VIZ DEBUG] LLM raw response type: {type(res)}")
if isinstance(res, VisualizationAgentOutput):
logger.info(f"[VIZ DEBUG] Specs generated: {len(res.charts)}")
return res.charts
logger.warning(f"[VIZ DEBUG] LLM returned non-schema response: {res}")
return []
except Exception as e:
logger.error(f"LLM specification failed: {e}")
# Fallback if no specs were generated
if not query:
return [
ChartSpec(chart_type=ChartType.PACE, title="Pace Chart"),
ChartSpec(chart_type=ChartType.HEART_RATE, title="Heart Rate Chart"),
]
return []