File size: 5,087 Bytes
d64fd55
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
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 []