study-buddy / app /agents /d3 /router.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
6.71 kB
from __future__ import annotations
import logging
import re
from typing import List, Optional
from pydantic import BaseModel, Field
from app.agents.cerebras_client import CerebrasClient
from app.agents.d3.registry import catalog_for_prompt
logger = logging.getLogger(__name__)
class D3Selection(BaseModel):
template_id: str = Field(description="Chosen template id from the catalog -- always pick one.")
reasoning: str = Field("", description="One sentence on why this template fits the request.")
selected_metric: str = Field("", description="The source metric to visualize, when the request or table identifies one.")
explicit_family: bool = Field(False, description="True when the user explicitly named the requested visual family.")
requires_metric_clarification: bool = Field(False, description="True only when multiple source metrics are equally plausible and none was requested.")
metric_candidates: list[str] = Field(default_factory=list, description="Two to five exact source column labels when clarification is required.")
class D3TemplateSelectionError(ValueError):
pass
_EXPLICIT_TEMPLATE_TERMS: tuple[tuple[tuple[str, ...], tuple[str, ...]], ...] = (
(("donut", "doughnut"), ("donut_chart",)),
(("pie chart", "pie graph", " pie "), ("pie_chart",)),
(("grouped bar", "clustered bar"), ("grouped_bar",)),
(("stacked bar",), ("stacked_bar",)),
(("bar chart", "bar graph"), ("bar_chart", "grouped_bar", "stacked_bar")),
(("multi-line", "multiple line", "multi line"), ("multi_line",)),
(("line chart", "line graph"), ("line", "multi_line")),
(("scatter",), ("scatter",)),
(("histogram",), ("histogram",)),
)
def _explicit_allowed_templates(concept: str) -> tuple[str, ...]:
padded = f" {concept.lower()} "
for terms, template_ids in _EXPLICIT_TEMPLATE_TERMS:
if any(term in padded for term in terms):
return template_ids
# A bare mention such as "simplify this table" describes the selected
# input, not the desired output. Lock only imperative/output phrasing.
if re.search(
r"\b(?:as|into)\s+(?:a\s+)?(?:compact\s+)?table\b|"
r"\b(?:create|make|render|show)\s+(?:me\s+)?(?:a\s+)?(?:compact\s+)?table\b",
concept.lower(),
):
return ("data_table",)
return ()
class D3TemplateRouter:
def __init__(self, client: Optional[CerebrasClient] = None) -> None:
self._client = client or CerebrasClient()
def select(self, concept: str, chunks: List[dict], familiarity: str) -> D3Selection:
chunk_text = "\n\n".join(f"[{c.get('source','?')}]: {c['text']}" for c in chunks)[:3000]
messages = [
{"role": "system", "content": (
"You choose the single best D3 chart template for a student's chart request. The "
"student has already asked for a chart -- always pick exactly one template id from "
"the CATALOG, never decline. Prefer what the student's own wording asks for (e.g. "
"'bar chart' -> a bar template, 'over time' -> a line template, a relationship/network "
"-> a network template). Use the source material as supplementary context for the data "
"itself, but it is never a reason to refuse a pick -- if the source doesn't contain the "
"data this chart needs, that is handled by a later synthesis step, not by you.\n\n"
f"CATALOG:\n{catalog_for_prompt()}\n\nSOURCE MATERIAL:\n{chunk_text}"
)},
{"role": "user", "content": f"Concept: '{concept}' (level: {familiarity}). Choose the best-fitting template."},
]
sel = self._client.structured_complete(messages, D3Selection, reasoning_effort="medium")
from app.agents.d3.registry import TEMPLATES
if sel.template_id not in TEMPLATES:
# Root logger defaults to WARNING (no logging.basicConfig() anywhere in the
# app) -- logger.info() here would be silently swallowed, so this must be a
# warning to actually surface a fallback that otherwise looks identical to a
# deliberate bar_chart pick.
logger.warning(
"D3TemplateRouter: model chose unknown template_id=%r for concept=%r -- falling back to bar_chart",
sel.template_id, concept,
)
sel = D3Selection(
template_id="bar_chart",
reasoning=f"Fallback: model chose unknown template id {sel.template_id!r}.",
)
return sel
def select_transform(self, concept: str, chunks: List[dict], familiarity: str) -> D3Selection:
"""Agent-select a trusted template, then enforce explicit user chart commands."""
chunk_text = "\n\n".join(f"[{c.get('source','selection')}]: {c.get('text','')}" for c in chunks)[:5000]
allowed = _explicit_allowed_templates(concept)
messages = [
{"role": "system", "content": (
"You route a selected paper table, plot, or image into one trusted visualization template. "
"Choose exactly one template_id from the live CATALOG. The user's explicit visual family is "
"a command: pie must remain pie, donut must remain donut, table must remain table, and so on. "
"Choose a source metric explicitly requested by the user. If none is named, prefer a clearly "
"labeled Avg/Average/overall column that applies to the requested rows. If there is no clear "
"aggregate and several columns are equally plausible, set requires_metric_clarification=true and "
"list their exact labels in metric_candidates. Do not choose values or "
"invent data; a separate extractor handles transcription.\n\n"
f"CATALOG:\n{catalog_for_prompt()}\n\nSELECTED CONTEXT:\n{chunk_text}"
)},
{"role": "user", "content": f"REQUEST: {concept}\nLEVEL: {familiarity}\nChoose the template and metric."},
]
selection = self._client.structured_complete(messages, D3Selection, reasoning_effort="low")
from app.agents.d3.registry import TEMPLATES
if selection.template_id not in TEMPLATES:
raise D3TemplateSelectionError(f"The visualization agent selected unknown template {selection.template_id!r}.")
if allowed and selection.template_id not in allowed:
raise D3TemplateSelectionError(
f"The visualization agent changed the explicitly requested chart family to {selection.template_id!r}."
)
return selection.model_copy(update={"explicit_family": bool(allowed)})