| """ |
| Visualization Configuration Module |
| |
| Centralizes all rules, strategies, spec builders, and style tokens for the |
| Data360 visualization system. |
| |
| Design principles: |
| - World Bank Data Visualization Style Guide (colors, typography, grid) |
| - FT Visual Vocabulary (chart-type selection by data relationship) |
| - All functions here are pure (no async, no I/O) → fully unit-testable |
| """ |
|
|
| from __future__ import annotations |
|
|
| import re |
| from collections import defaultdict |
| from dataclasses import dataclass, field |
| from enum import Enum |
| from numbers import Integral |
| from typing import Any, Literal |
|
|
| import pandas as pd |
|
|
| |
| |
| |
| TemporalFreq = Literal["annual", "quarterly", "monthly", "daily"] |
|
|
| |
| |
| |
| |
|
|
| WB_CAT_COLORS: list[str] = [ |
| "#34A7F2", |
| "#FF9800", |
| "#664AB6", |
| "#4EC2C0", |
| "#F3578E", |
| "#081079", |
| "#0C7C68", |
| "#AA0000", |
| "#DDDA21", |
| ] |
|
|
| WB_REGION_COLORS: dict[str, str] = { |
| "NAC": "#34A7F2", |
| "SSF": "#FF9800", |
| "MEA": "#664AB6", |
| "SAS": "#4EC2C0", |
| "EAS": "#F3578E", |
| "LCN": "#0C7C68", |
| "ECS": "#AA0000", |
| "AFW": "#DDDA21", |
| "AFE": "#FF9800", |
| "WLD": "#081079", |
| } |
|
|
| WB_GENDER_COLORS: dict[str, str] = { |
| "F": "#FF9800", |
| "M": "#664AB6", |
| "_T": "#4EC2C0", |
| "female": "#FF9800", |
| "male": "#664AB6", |
| } |
|
|
| WB_INCOME_COLORS: dict[str, str] = { |
| "HIC": "#016B6C", |
| "UMC": "#73AF48", |
| "LMC": "#DB95D7", |
| "LIC": "#3B4DA6", |
| } |
|
|
| WB_SEQ_GOOD: list[str] = ["#FDF6DB", "#A1CBCF", "#5D99C2", "#2868A0", "#023B6F"] |
| WB_SEQ_BAD: list[str] = ["#E3F6FD", "#91C5F0", "#8B8AC0", "#88506E", "#691B15"] |
| WB_SEQ_BLUE: list[str] = ["#E3F6FD", "#75CCEC", "#089BD4", "#0169A1", "#023B6F"] |
| WB_DIV_DEFAULT: list[str] = [ |
| "#920000", |
| "#BD6126", |
| "#E3A763", |
| "#EFEFEF", |
| "#80BDE7", |
| "#3587C3", |
| "#025288", |
| ] |
|
|
| WB_TEXT = "#111111" |
| WB_TEXT_SUBTLE = "#666666" |
| WB_GRID_COLOR = "#CED4DE" |
| WB_ZERO_COLOR = "#8A969F" |
| WB_REFERENCE = "#8A969F" |
| WB_NO_DATA = "#CED4DE" |
| WB_WHITE = "#FFFFFF" |
| WB_BACKGROUND = "#FFFFFF" |
| WB_FONT_FAMILY = "Noto Sans, Arial, sans-serif" |
|
|
|
|
| |
| |
| |
|
|
| KNOWN_ORDINAL_SORT_ORDERS: dict[str, list[str]] = { |
| "sex": ["female", "male", "_T"], |
| "gender": ["female", "male", "total"], |
| "ipc_phase": ["Minimal", "Stressed", "Crisis", "Emergency", "Famine"], |
| "income_group": ["Low income", "Lower middle income", "Upper middle income", "High income"], |
| "development_stage": ["Pre-transition", "Transition", "Post-transition"], |
| "age": [ |
| "80 years old and over", "80+", "80-84", "85-89", "90-94", "95-99", "100+", |
| "75 to 79 years old", "75-79", |
| "70 to 74 years old", "70-74", |
| "65 to 69 years old", "65-69", |
| "60 to 64 years old", "60-64", |
| "55 to 59 years old", "55-59", |
| "50 to 54 years old", "50-54", |
| "45 to 49 years old", "45-49", |
| "40 to 44 years old", "40-44", |
| "35 to 39 years old", "35-39", |
| "30 to 34 years old", "30-34", |
| "25 to 29 years old", "25-29", |
| "20 to 24 years old", "20-24", |
| "15 to 19 years old", "15-19", |
| "10 to 14 years old", "10-14", |
| "5 to 9 years old", "5-9", |
| "under 5 years old", "0-4", |
| "under 15 years old", "15 to 64 years old", "65 years old and over" |
| ] |
| } |
|
|
| def _get_dimension_sort_order(field_name: str, values: list[str]) -> list[str] | None: |
| """Return an explicit sorting list if the field or values match a known ordinal dimension.""" |
| normalized_field = (field_name or "").strip().lower() |
| if normalized_field in KNOWN_ORDINAL_SORT_ORDERS: |
| return KNOWN_ORDINAL_SORT_ORDERS[normalized_field] |
|
|
| |
| val_set = {str(v).strip().lower() for v in values} |
| for key, order in KNOWN_ORDINAL_SORT_ORDERS.items(): |
| order_set = {str(o).strip().lower() for o in order} |
| if val_set.intersection(order_set) and len(val_set.intersection(order_set)) >= 2: |
| sorted_avail = [] |
| for o in order: |
| o_norm = o.strip().lower() |
| matched_val = None |
| for val in values: |
| if str(val).strip().lower() == o_norm: |
| matched_val = val |
| break |
| if matched_val is not None: |
| sorted_avail.append(matched_val) |
| return sorted_avail |
|
|
| return None |
|
|
|
|
| |
| |
| |
|
|
|
|
| def wb_altair_config() -> dict: |
| """Return World Bank style config dict for injection into Vega-Lite specs.""" |
| return { |
| "background": WB_BACKGROUND, |
| "font": WB_FONT_FAMILY, |
| "title": { |
| "fontSize": 16, |
| "fontWeight": "bold", |
| "color": WB_TEXT, |
| |
| |
| "lineHeight": 22, |
| "anchor": "start", |
| "offset": 8, |
| "subtitleFontSize": 12, |
| "subtitleColor": WB_TEXT_SUBTLE, |
| "subtitleFontWeight": "normal", |
| "subtitlePadding": 4, |
| |
| "subtitleLineHeight": 18, |
| }, |
| "axis": { |
| "grid": False, |
| "labelColor": WB_TEXT_SUBTLE, |
| "labelFontSize": 12, |
| "labelFont": WB_FONT_FAMILY, |
| "titleColor": WB_TEXT, |
| "titleFontSize": 12, |
| "titleFont": WB_FONT_FAMILY, |
| "titleFontWeight": "bold", |
| "gridColor": WB_GRID_COLOR, |
| "gridDash": [4, 2], |
| "gridWidth": 1, |
| "domainColor": WB_GRID_COLOR, |
| "tickColor": WB_GRID_COLOR, |
| "tickCount": 5, |
| "labelOverlap": "greedy", |
| }, |
| "legend": { |
| "labelColor": WB_TEXT, |
| "labelFont": WB_FONT_FAMILY, |
| "labelFontSize": 12, |
| "labelFontWeight": "bold", |
| "labelLimit": 300, |
| "titleColor": WB_TEXT, |
| "titleFont": WB_FONT_FAMILY, |
| "titleFontSize": 12, |
| "orient": "top", |
| "direction": "horizontal", |
| }, |
| "range": {"category": WB_CAT_COLORS}, |
| "view": {"stroke": "transparent"}, |
| "line": {"strokeWidth": 3, "strokeCap": "round"}, |
| "point": {"size": 60, "stroke": WB_WHITE, "strokeWidth": 1}, |
| "bar": {"cornerRadiusTopLeft": 2, "cornerRadiusTopRight": 2}, |
| } |
|
|
|
|
| def _suppress_redundant_legends(node): |
| if isinstance(node, dict): |
| if "encoding" in node: |
| enc = node["encoding"] |
| if isinstance(enc, dict): |
| color_enc = enc.get("color") |
| if isinstance(color_enc, dict) and "field" in color_enc: |
| color_field = color_enc.get("field") |
| if color_field: |
| x_field = enc.get("x", {}).get("field") if isinstance(enc.get("x"), dict) else None |
| y_field = enc.get("y", {}).get("field") if isinstance(enc.get("y"), dict) else None |
| if color_field in (x_field, y_field): |
| color_enc["legend"] = None |
| for k, v in node.items(): |
| _suppress_redundant_legends(v) |
| elif isinstance(node, list): |
| for item in node: |
| _suppress_redundant_legends(item) |
|
|
|
|
| def inject_wb_config(vl_spec: dict) -> dict: |
| """Merge WB style config into a Vega-Lite spec without overwriting user settings.""" |
| _suppress_redundant_legends(vl_spec) |
| wb_cfg = wb_altair_config() |
| if "config" not in vl_spec: |
| vl_spec["config"] = wb_cfg |
| else: |
| for section, props in wb_cfg.items(): |
| if section not in vl_spec["config"]: |
| vl_spec["config"][section] = props |
| elif isinstance(props, dict) and isinstance( |
| vl_spec["config"].get(section), dict |
| ): |
| for k, v in props.items(): |
| vl_spec["config"][section].setdefault(k, v) |
| return vl_spec |
|
|
|
|
| |
| |
| |
|
|
| |
| |
| |
| _TOOLTIP_SPECS: dict[str, dict] = { |
| "value": {"title": "Value", "format": ",.2f", "type": "quantitative"}, |
| "country": {"title": "Economy", "type": "nominal"}, |
| "sex": {"title": "Sex", "type": "nominal"}, |
| "age": {"title": "Age Group", "type": "nominal"}, |
| "urbanisation": {"title": "Urbanisation", "type": "nominal"}, |
| "residence": {"title": "Residence", "type": "nominal"}, |
| "comp_breakdown_1": {"title": "Dimension 1", "type": "nominal"}, |
| "comp_breakdown_2": {"title": "Dimension 2", "type": "nominal"}, |
| "comp_breakdown_3": {"title": "Dimension 3", "type": "nominal"}, |
| "time_period": {"title": "Period", "type": "temporal"}, |
| "obs_value": {"title": "Value", "format": ",.2f", "type": "quantitative"}, |
| "ref_area": {"title": "Economy", "type": "nominal"}, |
| "region": {"title": "Region", "type": "nominal"}, |
| } |
|
|
| _TOOLTIP_PRIORITY = [ |
| "year", |
| "time_period", |
| "value", |
| "obs_value", |
| "country", |
| "ref_area", |
| "region", |
| "sex", |
| "age", |
| "urbanisation", |
| "residence", |
| "comp_breakdown_1", |
| "comp_breakdown_2", |
| ] |
|
|
| _VOWELS = {"a", "e", "i", "o", "u"} |
|
|
|
|
| def _year_range_label(year_series: pd.Series) -> str | None: |
| """Min–max year label, e.g. ``1990-2024`` or ``2020`` when only one year.""" |
| if year_series.empty: |
| return None |
| try: |
| if pd.api.types.is_datetime64_any_dtype(year_series): |
| ynum = year_series.dt.year |
| else: |
| ynum = pd.to_numeric(year_series, errors="coerce") |
| yvalid = ynum.dropna() |
| if yvalid.empty: |
| return None |
| y0, y1 = int(yvalid.min()), int(yvalid.max()) |
| return f"{y0}-{y1}" if y0 != y1 else str(y0) |
| except (TypeError, ValueError): |
| return None |
|
|
|
|
| def format_chart_context_subtitle(df: pd.DataFrame) -> str | None: |
| """Build geography list + year range for chart subtitle (product-style). |
| |
| Example: ``\"Philippines, Belgium, 1990-2024\"``. Long geography lists are truncated. |
| """ |
| parts: list[str] = [] |
| if "country" in df.columns: |
| vals = sorted( |
| {str(v).strip() for v in df["country"].dropna() if str(v).strip()}, |
| key=str.casefold, |
| ) |
| if vals: |
| if len(vals) <= 3: |
| parts.append(", ".join(vals)) |
| year_lbl = None |
| if "year" in df.columns: |
| year_lbl = _year_range_label(df["year"]) |
| if year_lbl: |
| parts.append(year_lbl) |
| if not parts: |
| return None |
| return ", ".join(parts) |
|
|
|
|
| def build_chart_title_with_context( |
| main_title: str | list[str] | dict, |
| unit_subtitle: str | None, |
| df: pd.DataFrame, |
| ) -> str | dict | list: |
| """Vega-Lite title: main text plus subtitle lines (geography + years, unit). |
| |
| Subtitle is returned as a **list of strings** so Vega-Lite v5 renders each |
| part on its own line. This prevents the single-line overflow that occurs |
| when country names, year ranges, units, and trim notes are concatenated. |
| """ |
| if isinstance(main_title, str) and main_title.strip().startswith("{"): |
| try: |
| import json |
| main_title = json.loads(main_title) |
| except Exception: |
| pass |
|
|
| if isinstance(main_title, dict): |
| return main_title |
| ctx = format_chart_context_subtitle(df) |
| subtitle_parts: list[str] = [] |
| if ctx: |
| subtitle_parts.append(ctx) |
| if unit_subtitle and str(unit_subtitle).strip(): |
| subtitle_parts.append(str(unit_subtitle).strip()) |
| if not subtitle_parts: |
| return main_title |
| return {"text": main_title, "subtitle": subtitle_parts} |
| def _clean_label_generic(label: str) -> str: |
| """Generically cleans long indicator/dimension labels by stripping trailing parentheticals.""" |
| if not label or not isinstance(label, str): |
| return label |
| s = label.strip() |
| |
| for _ in range(2): |
| if s.endswith(")"): |
| idx = s.rfind("(") |
| if idx != -1: |
| content = s[idx+1:-1].lower() |
| |
| strip_keywords = {"estimate", "modeled", "ilo", "gdp", "%", "percent", "constant", "current", "low to", "base", "index"} |
| if any(kw in content for kw in strip_keywords) or len(content) > 10: |
| s = s[:idx].strip() |
| else: |
| break |
| else: |
| break |
| else: |
| break |
| return s |
|
|
|
|
|
|
| def _clean_unit_label(label: str) -> str: |
| """Returns a short, clean unit label suitable for axis titles.""" |
| if not label: |
| return "Value" |
| l = label.lower() |
| if "(" in label or "per" in l or "/" in label: |
| return label.strip() |
| if "percent" in l or "percentage" in l or "%" in l: |
| if "gdp" in l: |
| return "% of GDP" |
| return "%" if label == "%" else "Percentage" |
| if "index" in l or "score" in l: |
| return "Index" |
| if "usd" in l or "us$" in l or "dollar" in l: |
| if "constant" in l or "current" in l: |
| return label.strip() |
| return "USD" |
| if "share" in l or "proportion" in l or "ratio" in l: |
| return "Share" |
| if "estimate" in l: |
| return "Estimate" |
| if "co2" in l or "greenhouse" in l: |
| return "Tonnes CO2-eq" |
| return label.strip() |
|
|
|
|
| |
| _CUSTOM_BREAKDOWN_DIMS = {"comp_breakdown_1", "comp_breakdown_2", "comp_breakdown_3"} |
|
|
|
|
| def _generate_color_shades(hex_color: str, n: int) -> list[str]: |
| """Return n shades of hex_color spread from dark to light (HSL lightness). |
| |
| n=1 → returns the base color unchanged. |
| n=2 → [dark, base] (darker shade + original). |
| n=3 → [dark, base, light]. |
| n>3 → evenly distributed from 0.28 L to 0.75 L. |
| """ |
| import colorsys |
| if n <= 0: |
| return [] |
| if n == 1: |
| return [hex_color] |
| h_str = hex_color.lstrip("#") |
| r, g, b = (int(h_str[i : i + 2], 16) / 255.0 for i in (0, 2, 4)) |
| h, _l, s = colorsys.rgb_to_hls(r, g, b) |
| shades: list[str] = [] |
| for i in range(n): |
| factor = i / (n - 1) |
| new_l = 0.28 + factor * 0.47 |
| nr, ng, nb = colorsys.hls_to_rgb(h, new_l, min(s, 0.90)) |
| shades.append(f"#{int(nr * 255):02x}{int(ng * 255):02x}{int(nb * 255):02x}") |
| return shades |
|
|
|
|
| def _compute_legend_layout(labels: list[str], chart_width: int = 680) -> dict: |
| """Compute Vega-Lite legend config (orient/direction/columns/labelLimit) from data. |
| |
| Uses the number of items, their rendered pixel width, and the chart width to |
| determine whether the legend fits in one horizontal row, needs multiple rows |
| (grid), or must fall back to a vertical list. |
| |
| Approximate rendered width per item at 11 px font: |
| 6.5 px/char × label_length + symbol (20 px) + padding (16 px) |
| """ |
| import math |
| n = len(labels) |
| if n == 0: |
| return {"orient": "bottom", "labelFontSize": 11, "symbolSize": 80} |
| max_lbl = max(len(l) for l in labels) |
| item_px = int(max_lbl * 6.5 + 36) |
| items_per_row = max(1, chart_width // item_px) |
|
|
| base = {"labelFontSize": 11, "symbolSize": 80} |
|
|
| if n <= items_per_row: |
| |
| return {**base, "orient": "bottom", "direction": "horizontal", |
| "labelLimit": max(150, item_px - 36)} |
|
|
| if items_per_row >= 2: |
| |
| cols = min(items_per_row, max(2, math.ceil(n / 3))) |
| return {**base, "orient": "bottom", "direction": "horizontal", |
| "columns": cols, "labelLimit": max(150, item_px - 36)} |
|
|
| |
| return {**base, "orient": "bottom", "direction": "vertical", "labelLimit": 0} |
|
|
|
|
| def _estimate_legend_height( |
| n_items: int, |
| layout: dict, |
| has_title: bool = True, |
| ) -> int: |
| """Estimate legend pixel height from the layout dict returned by _compute_legend_layout. |
| |
| Used to shrink panel heights so the total figure height stays within budget. |
| """ |
| import math |
| direction = layout.get("direction", "vertical") |
| columns = layout.get("columns", 0) |
| row_px = 22 |
| title_px = 20 if has_title else 0 |
| padding = 16 |
|
|
| if direction == "horizontal" and columns: |
| rows = math.ceil(n_items / columns) |
| elif direction == "horizontal": |
| rows = 1 |
| else: |
| rows = n_items |
|
|
| return title_px + rows * row_px + padding |
|
|
|
|
|
|
| def _detect_scale_incompatibility( |
| df: pd.DataFrame, |
| breakdown_dim: str, |
| magnitude_threshold: float = 1.5, |
| ) -> bool: |
| """Return True when breakdown series have incompatible Y-axis scales. |
| |
| Computes the order of magnitude (log10 of max |value|) for each |
| breakdown series. If the spread between the largest and smallest |
| magnitudes exceeds *magnitude_threshold*, the series cannot share a |
| Y-axis without visually compressing the small-magnitude series. |
| |
| A threshold of 1.5 corresponds to roughly a 30× difference between |
| the dominant series and the smallest (e.g., a WGI percentile rank |
| peaking at ~65 vs. a standard error peaking at ~0.2). |
| |
| Only meaningful for *_CUSTOM_BREAKDOWN_DIMS*; standard demographic |
| dims (sex, age) almost always share a unit and should never fire. |
| |
| Returns False when: |
| - breakdown_dim is not in df.columns or 'value' is missing |
| - fewer than 2 unique breakdown values are present |
| - all series are zero or NaN (no meaningful magnitudes to compare) |
| |
| Examples:: |
| |
| WGI: EST max|val|≈0.6 (mag≈-0.22), SC max|val|≈65 (mag≈1.81) |
| spread = 1.81 − (−0.22) = 2.03 → True (exceeds 1.5) |
| |
| IPC phases: all person counts, max|val|∈[1000, 5000] |
| mags ≈ [3.0, 3.5, 3.7], spread = 0.7 → False |
| """ |
| import math |
|
|
| if breakdown_dim not in df.columns or "value" not in df.columns: |
| return False |
| bd_vals = df[breakdown_dim].dropna().unique() |
| if len(bd_vals) < 2: |
| return False |
|
|
| mags: list[float] = [] |
| for v in bd_vals: |
| series_vals = df.loc[df[breakdown_dim] == v, "value"].dropna() |
| if series_vals.empty: |
| continue |
| max_abs = float(series_vals.abs().max()) |
| if max_abs == 0: |
| mags.append(0.0) |
| else: |
| mags.append(math.log10(max_abs)) |
|
|
| if len(mags) < 2: |
| return False |
| return (max(mags) - min(mags)) >= magnitude_threshold |
|
|
|
|
| def _format_breakdown_subtitle(df: pd.DataFrame, color_dim: str | None) -> str | None: |
| """Return a compact subtitle note when color_dim is a heterogeneous custom breakdown. |
| |
| Appended to chart subtitles so end users can see which series are present |
| and understand they may carry different units or scales. |
| |
| Returns None when: |
| - color_dim is a standard dimension (country, sex, age, …) |
| - there is only one unique breakdown value |
| - series share a compatible scale (log10 magnitude spread ≤ 1.5) — the |
| unit warning is suppressed because _detect_scale_incompatibility returns |
| False. This correctly handles summary-measure breakdowns such as |
| "Arithmetic mean" vs "Median" which share the same currency unit. |
| """ |
| if color_dim not in _CUSTOM_BREAKDOWN_DIMS: |
| return None |
| if color_dim not in df.columns: |
| return None |
| vals = sorted(str(v) for v in df[color_dim].dropna().unique()) |
| if len(vals) <= 1: |
| return None |
| series_list = ", ".join(vals) |
| |
| |
| |
| if _detect_scale_incompatibility(df, color_dim): |
| return f"Series: {series_list} — series may have different units/scales" |
| return f"Series: {series_list}" |
|
|
|
|
| def _append_breakdown_note( |
| title: str | dict, |
| df: pd.DataFrame, |
| color_dim: str | None, |
| ) -> str | dict: |
| """Inject breakdown note into a Vega-Lite title dict's subtitle. |
| |
| When subtitle is a list (Vega-Lite multi-line form), the note is appended |
| as a new line. When subtitle is a string, it is appended with ' · '. |
| """ |
| note = _format_breakdown_subtitle(df, color_dim) |
| if not note: |
| return title |
| if isinstance(title, dict): |
| existing = title.get("subtitle", "") |
| if isinstance(existing, list): |
| return {**title, "subtitle": existing + [note]} |
| new_sub = f"{existing} · {note}" if existing else note |
| return {**title, "subtitle": new_sub} |
| |
| return {"text": title, "subtitle": note} |
|
|
|
|
| def _cap_cardinality( |
| df: pd.DataFrame, |
| dim: str, |
| max_n: int, |
| ) -> tuple[pd.DataFrame, int | None]: |
| """Cap the number of unique values for *dim* to *max_n*. |
| |
| Shared utility called by every spec builder that renders one visual element |
| per dim value (facet panels, bar rows, color lines). This is standard |
| chart best practice: beyond ~8–12 elements embedded charts overflow the |
| chatbot UI and individual items become unreadable. |
| |
| Selection strategy: top-N by most-recent data point, ties broken by row |
| count (more data = more informative panel). Rows outside the top-N are |
| dropped from the returned DataFrame. |
| |
| Args: |
| df: Input DataFrame. Must have a ``year`` column for recency sort. |
| dim: Dimension column whose cardinality to cap (e.g. ``country``). |
| max_n: Maximum number of unique values to retain. |
| |
| Returns: |
| (trimmed_df, original_n) where *original_n* is the pre-trim count, or |
| *None* when no trimming was needed (df is returned unchanged). |
| """ |
| if dim not in df.columns: |
| return df, None |
| n_total = df[dim].nunique() |
| if n_total <= max_n: |
| return df, None |
|
|
| if "year" in df.columns: |
| latest = df.groupby(dim)["year"].max() |
| else: |
| latest = pd.Series(dtype="object", index=df[dim].unique()) |
| count = df.groupby(dim).size() |
| rank = pd.DataFrame( |
| {"latest": latest.reindex(count.index).fillna(pd.Timestamp.min), "count": count} |
| ) |
| top = ( |
| rank.sort_values(["latest", "count"], ascending=False) |
| .head(max_n) |
| .index.tolist() |
| ) |
| return df[df[dim].isin(top)].copy(), n_total |
|
|
|
|
| def _append_trim_note( |
| title: str | dict, |
| dim_label: str, |
| shown: int, |
| original: int | None, |
| ) -> str | dict: |
| """Inject a 'Showing N of M' note into the chart subtitle when cardinality |
| was capped by :func:`_cap_cardinality`. |
| |
| No-op when *original* is None (no trimming occurred). |
| When subtitle is a list (Vega-Lite multi-line form), the note is appended |
| as a new line. When subtitle is a string, it is appended with ' · '. |
| """ |
| if original is None: |
| return title |
| dim_title = ( |
| _TOOLTIP_SPECS.get(dim_label, {}).get("title") |
| or dim_label.replace("_", " ") |
| ).strip().lower() |
| |
| |
| if dim_label.startswith("comp_breakdown_"): |
| dim_title = "breakdown" |
| |
| if dim_title.endswith("y") and len(dim_title) > 2 and dim_title[-2] not in _VOWELS: |
| dim_plural = f"{dim_title[:-1]}ies" |
| elif dim_title.endswith(("s", "x", "z", "ch", "sh")): |
| dim_plural = f"{dim_title}es" |
| else: |
| dim_plural = f"{dim_title}s" |
| note = ( |
| f"Showing {shown} of {original} {dim_plural} by most recent data — " |
| "specify a subset for the full view" |
| ) |
| if isinstance(title, dict): |
| existing = title.get("subtitle", "") |
| if isinstance(existing, list): |
| return {**title, "subtitle": existing + [note]} |
| return {**title, "subtitle": f"{existing} · {note}" if existing else note} |
| return {"text": title, "subtitle": note} |
|
|
|
|
| |
| _MULTI_IND_TOOLTIP_DIMS: tuple[str, ...] = ( |
| "year", |
| "time_period", |
| "country", |
| "ref_area", |
| "region", |
| "sex", |
| "age", |
| "urbanisation", |
| "residence", |
| ) |
|
|
| |
| _LINE_HOVER_POINT: dict[str, object] = {"filled": True, "size": 56} |
|
|
|
|
| def _multi_indicator_tooltip_columns( |
| df_columns: list[str], value_col: str |
| ) -> list[str]: |
| colset = set(df_columns) |
| out: list[str] = [] |
| for c in _MULTI_IND_TOOLTIP_DIMS: |
| if c in colset: |
| out.append(c) |
| if value_col in colset and value_col not in out: |
| out.append(value_col) |
| return out |
|
|
|
|
| def _tooltip_spec_for_time_dim( |
| col: str, |
| viz_data: pd.DataFrame | None, |
| temporal_freq: TemporalFreq | None = None, |
| ) -> dict: |
| """Return a Vega-Lite tooltip spec for year/time_period columns. |
| |
| When a temporal frequency is known (or can be detected from the values), |
| the spec uses ``type: temporal`` with the correct timeUnit + format so that |
| Vega-Lite formats the internally-parsed epoch timestamp correctly. Without |
| this, charts with a temporal X-axis display the raw millisecond number |
| (e.g. 1596240000000) instead of a human-readable date string. |
| """ |
| title = "Year" if col == "year" else "Period" |
|
|
| |
| _FREQ_TOOLTIP: dict[str, dict] = { |
| "annual": {"timeUnit": "utcyear", "format": "%Y"}, |
| "monthly": {"timeUnit": "utcyearmonth", "format": "%b %Y"}, |
| "quarterly": {"timeUnit": "utcyearquarter", "format": "Q%q %Y"}, |
| "daily": {"timeUnit": "utcyearmonthdate", "format": "%Y-%m-%d"}, |
| } |
|
|
| |
| |
| if viz_data is not None and col in viz_data.columns and viz_data[col].nunique() <= 1: |
| return {"field": col, "title": title, "type": "nominal"} |
|
|
| freq: TemporalFreq | None = temporal_freq |
| if freq is None and viz_data is not None and col in viz_data.columns: |
| |
| freq = _detect_temporal_frequency(viz_data[col]) |
|
|
| if freq is not None: |
| cfg = _FREQ_TOOLTIP.get(freq, _FREQ_TOOLTIP["annual"]) |
| return { |
| "field": col, |
| "title": title, |
| "type": "temporal", |
| "timeUnit": cfg["timeUnit"], |
| "format": cfg["format"], |
| } |
|
|
| return {"field": col, "title": title, "type": "nominal"} |
|
|
|
|
| def build_structured_tooltips( |
| columns: list[str], |
| mark_type: str, |
| indicator_labels: dict[str, str] | None = None, |
| value_format: str = ",.2f", |
| viz_data: pd.DataFrame | None = None, |
| temporal_freq: TemporalFreq | None = None, |
| dim_name_labels: dict[str, str] | None = None, |
| indicator_name: str | None = None, |
| ) -> list[dict]: |
| """Build typed, labelled tooltip list for a Vega-Lite encoding. |
| |
| indicator_labels: optional {col_name: human_label} for indicator value columns |
| in multi-indicator charts (e.g. {"gdp_per_capita": "GDP per capita (USD)"}). |
| value_format: D3 format string for quantitative value fields. |
| viz_data: when set, ``year`` / ``time_period`` frequency is detected from the |
| column values to produce correctly formatted temporal tooltip labels. |
| temporal_freq: explicit temporal frequency; overrides auto-detection from |
| viz_data. Pass ``result.temporal_frequency`` from temporal chart builders |
| so the tooltip date format matches the X-axis format exactly. |
| dim_name_labels: optional {col_name: human_label} for comp_breakdown_* columns |
| sourced from the disaggregation API. Overrides the generic "Dimension N" |
| fallback in ``_TOOLTIP_SPECS`` for those fields. |
| indicator_name: human-readable indicator title. When provided, appended as a |
| constant tooltip entry (``{"value": indicator_name, "title": "Indicator"}``) |
| so hovering always shows which indicator is displayed. |
| """ |
| ordered = [c for c in _TOOLTIP_PRIORITY if c in columns] |
| ordered += [c for c in columns if c not in _TOOLTIP_PRIORITY and not c.startswith("_")] |
|
|
| tooltips = [] |
| for col in ordered: |
| if col in ("year", "time_period"): |
| tooltips.append(_tooltip_spec_for_time_dim(col, viz_data, temporal_freq)) |
| continue |
| if dim_name_labels and col in dim_name_labels: |
| title = dim_name_labels[col] |
| col_type = "quantitative" if "value" in col.lower() else "nominal" |
| tip = {"field": col, "title": title, "type": col_type} |
| if col_type == "quantitative": |
| tip["format"] = value_format |
| elif indicator_labels and col in indicator_labels: |
| tip = { |
| "field": col, |
| "title": indicator_labels[col], |
| "format": value_format, |
| "type": "quantitative", |
| } |
| elif col in _TOOLTIP_SPECS: |
| spec = _TOOLTIP_SPECS[col] |
| |
| |
| title = ( |
| dim_name_labels.get(col) |
| if (dim_name_labels and col in dim_name_labels) |
| else spec["title"] |
| ) |
| tip = {"field": col, "title": title, "type": spec["type"]} |
| if "format" in spec: |
| |
| if col in ("value", "obs_value"): |
| tip["format"] = value_format |
| else: |
| tip["format"] = spec["format"] |
| else: |
| tip = {"field": col, "title": col.replace("_", " ").title()} |
| tooltips.append(tip) |
|
|
| |
| |
| |
| if indicator_name: |
| tooltips.append({"value": indicator_name, "title": "Indicator"}) |
|
|
| return tooltips |
|
|
|
|
| def apply_structured_tooltips( |
| vl_spec: dict, |
| columns: list[str], |
| mark_type: str, |
| indicator_labels: dict[str, str] | None = None, |
| viz_data: pd.DataFrame | None = None, |
| ) -> dict: |
| tips = build_structured_tooltips( |
| columns, mark_type, indicator_labels, viz_data=viz_data |
| ) |
| vl_spec.setdefault("encoding", {})["tooltip"] = tips |
| return vl_spec |
|
|
|
|
| |
| |
| |
|
|
|
|
| class ChartStrategy(str, Enum): |
| """Named chart strategies mapped to FT Visual Vocabulary categories.""" |
|
|
| TEMPORAL_SINGLE = "temporal_single" |
| TEMPORAL_MULTI_IND = ( |
| "temporal_multi_indicator" |
| ) |
| CORRELATION = "correlation" |
| CORRELATION_TEMPORAL = "correlation_temporal" |
| CROSS_SECTIONAL = ( |
| "cross_sectional" |
| ) |
| DISTRIBUTION = "distribution" |
| BREAKDOWN_COMPARISON = ( |
| "breakdown_comparison" |
| ) |
| SMALL_MULTIPLES = ( |
| "small_multiples" |
| ) |
| HEATMAP = "heatmap" |
| STACKED_AREA = "stacked_area" |
| STACKED_BAR = "stacked_bar" |
| CHOROPLETH = "choropleth" |
| FALLBACK_LINE = "fallback_line" |
|
|
|
|
| @dataclass |
| class StrategyResult: |
| strategy: ChartStrategy |
| reason: str |
| |
| indicator_cols: list[str] = field( |
| default_factory=list |
| ) |
| color_dim: str | None = None |
| facet_dim: str | None = None |
| |
| |
| |
| |
| secondary_color_dim: str | None = None |
| x_dim: str | None = None |
| y_dim: str | None = None |
| scale_incompatible: bool = False |
| temporal_frequency: TemporalFreq = "annual" |
| |
| scale_compatibility: dict | None = None |
| |
| |
| dim_name_labels: dict[str, str] = field(default_factory=dict) |
| |
| |
| mark_hint: str | None = None |
| scale_type: str | None = None |
| unit_mult: int = 0 |
| raw_hint: str | None = None |
| |
| data_profile: dict | None = None |
| refusal_reason: str | None = None |
|
|
|
|
| from typing import Protocol |
|
|
|
|
| @dataclass |
| class RoutingContext: |
| df: pd.DataFrame |
| n_indicators: int |
| hint: str | None |
| raw_hint: str | None |
| ind_cols: list[str] |
|
|
| |
| year_count: int |
| country_count: int |
| breakdown_counts: dict[str, int] |
| n_breakdowns: int |
| max_years_per_country: int = 1 |
| avg_years_per_country: float = 0.0 |
| scale_type: str | None = None |
| unit_mult: int = 0 |
| skewness: float = 0.0 |
| |
| indicator_value_ranges: dict[str, dict[str, float]] = field(default_factory=dict) |
| |
| sparse_country_count: int = 0 |
| completeness_pct: float = 100.0 |
| same_unit: bool = True |
| can_share_axis: bool = True |
| refusal_reason: str | None = None |
|
|
| def register_refusal(self, hint: str, reason: str): |
| self.refusal_reason = f"Cannot honor requested '{hint}': {reason}" |
|
|
| @classmethod |
| def build( |
| cls, |
| df: pd.DataFrame, |
| n_indicators: int, |
| chart_type_hint: str | None, |
| indicator_cols: list[str] | None, |
| raw_unit: str | None = None, |
| raw_unit_mult: int = 0, |
| ) -> RoutingContext: |
| hint = parse_chart_type_hint(chart_type_hint) |
| cols = set(df.columns) |
|
|
| year_count = df["year"].nunique() if "year" in cols else 0 |
| country_count = df["country"].nunique() if "country" in cols else 0 |
|
|
| max_years_per_country = 1 |
| avg_years_per_country = 0.0 |
| if country_count > 0 and "country" in cols and "year" in cols and not df.empty: |
| try: |
| max_years_per_country = int(df.groupby("country")["year"].nunique().max()) |
| avg_years_per_country = float(df.groupby("country")["year"].nunique().mean()) |
| except Exception: |
| pass |
|
|
| sex_count = df["sex"].nunique() if "sex" in cols else 0 |
| age_count = df["age"].nunique() if "age" in cols else 0 |
| urban_count = df["urbanisation"].nunique() if "urbanisation" in cols else 0 |
| residence_count = df["residence"].nunique() if "residence" in cols else 0 |
| cb1_count = df["comp_breakdown_1"].nunique() if "comp_breakdown_1" in cols else 0 |
| cb2_count = df["comp_breakdown_2"].nunique() if "comp_breakdown_2" in cols else 0 |
| cb3_count = df["comp_breakdown_3"].nunique() if "comp_breakdown_3" in cols else 0 |
| unit_count = df["unit_measure"].nunique() if "unit_measure" in cols else 0 |
|
|
| |
| |
| |
| err_band_dims = set() |
| if country_count <= 1: |
| for col_name in ("comp_breakdown_1", "comp_breakdown_2", "comp_breakdown_3"): |
| if col_name in cols: |
| unique_vals = [str(x).lower().strip() for x in df[col_name].dropna().unique()] |
| if unique_vals: |
| is_err = True |
| for val in unique_vals: |
| is_val_err = False |
| if "estimate" in val: |
| is_val_err = True |
| elif val == "est" or val.startswith("est ") or val.endswith(" est") or " est " in val: |
| is_val_err = True |
| elif "standard error" in val or "std error" in val or "std. error" in val: |
| is_val_err = True |
| elif val in ("stderr", "std_err", "std.err", "std. err", "s.e."): |
| is_val_err = True |
| elif val == "se" or val.startswith("se ") or val.endswith(" se") or " se " in val: |
| is_val_err = True |
|
|
| if not is_val_err: |
| is_err = False |
| break |
| if is_err: |
| err_band_dims.add(col_name) |
|
|
| breakdown_counts = { |
| k: v |
| for k, v in [ |
| ("sex", sex_count), |
| ("age", age_count), |
| ("urbanisation", urban_count), |
| ("residence", residence_count), |
| ("comp_breakdown_1", cb1_count if "comp_breakdown_1" not in err_band_dims else 0), |
| ("comp_breakdown_2", cb2_count if "comp_breakdown_2" not in err_band_dims else 0), |
| ("comp_breakdown_3", cb3_count if "comp_breakdown_3" not in err_band_dims else 0), |
| ("unit_measure", unit_count), |
| ] |
| if v > 1 |
| } |
|
|
| |
| normalized = (raw_unit or "").upper().strip() |
| if "$" in normalized or "USD" in normalized or "CURRENCY" in normalized or "DOLLARS" in normalized: |
| scale_type = "currency" |
| elif normalized == "PS" or "PEOPLE" in normalized or "PERSONS" in normalized or "COUNT" in normalized or "HEADCOUNT" in normalized or "PERSON" in normalized: |
| scale_type = "persons" |
| elif "%" in normalized or "PERCENT" in normalized or "RATE" in normalized or "SHARE" in normalized or "PROPORTION" in normalized: |
| scale_type = "percentage" |
| else: |
| scale_type = "index" |
|
|
| |
| unit_mult = raw_unit_mult |
| if "unit_mult" in df.columns: |
| _mults = df["unit_mult"].dropna().unique() |
| if len(_mults) == 1: |
| try: |
| unit_mult = int(_mults[0]) |
| except (ValueError, TypeError): |
| pass |
|
|
| |
| skewness = 0.0 |
| if "value" in df.columns: |
| val_series = pd.to_numeric(df["value"], errors="coerce").dropna() |
| if not val_series.empty: |
| skew = val_series.skew() |
| skewness = float(skew) if not pd.isna(skew) else 0.0 |
|
|
| |
| indicator_value_ranges: dict[str, dict[str, float]] = {} |
| if indicator_cols: |
| for col in indicator_cols: |
| if col in df.columns: |
| series = pd.to_numeric(df[col], errors="coerce").dropna() |
| if not series.empty: |
| indicator_value_ranges[col] = { |
| "min": float(series.min()), |
| "max": float(series.max()), |
| } |
|
|
| return cls( |
| df=df, |
| n_indicators=n_indicators, |
| hint=hint, |
| raw_hint=chart_type_hint, |
| ind_cols=indicator_cols or [], |
| year_count=year_count, |
| country_count=country_count, |
| max_years_per_country=max_years_per_country, |
| avg_years_per_country=avg_years_per_country, |
| breakdown_counts=breakdown_counts, |
| n_breakdowns=len(breakdown_counts), |
| scale_type=scale_type, |
| unit_mult=unit_mult, |
| skewness=skewness, |
| indicator_value_ranges=indicator_value_ranges, |
| ) |
|
|
| class RoutingRule(Protocol): |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: ... |
|
|
| class ExplicitScatterRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("point", "scatter", "connected_scatter") and ctx.n_indicators == 2 and len(ctx.ind_cols) == 2: |
| if ctx.year_count > 1: |
| return StrategyResult( |
| ChartStrategy.CORRELATION_TEMPORAL, |
| "User requested scatter; 2 indicators, multi-year → connected scatter", |
| indicator_cols=ctx.ind_cols, |
| color_dim="country" if ctx.country_count > 0 else None, |
| ) |
| return StrategyResult( |
| ChartStrategy.CORRELATION, |
| "User requested scatter; 2 indicators, single year → scatterplot", |
| indicator_cols=ctx.ind_cols, |
| color_dim="country" if ctx.country_count > 0 else None, |
| ) |
| return None |
|
|
| class ExplicitStackedAreaMultiIndicatorRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("area", "stacked_area") and ctx.n_indicators >= 2 and len(ctx.ind_cols) >= 2: |
| if not (ctx.same_unit and ctx.can_share_axis): |
| ctx.register_refusal(ctx.hint, "Indicators represent different units or incompatible scales") |
| return None |
| if ctx.year_count <= 1: |
| ctx.register_refusal(ctx.hint, "Area charts require multiple years of data") |
| return None |
| return StrategyResult( |
| ChartStrategy.STACKED_AREA, |
| f"User requested stacked area; {ctx.n_indicators} indicators, {ctx.year_count} years → stacked area chart", |
| indicator_cols=ctx.ind_cols, |
| color_dim="indicator", |
| ) |
| return None |
|
|
| class ExplicitStackedBarRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("stacked_bar", "bar"): |
| if ctx.n_indicators >= 2 and len(ctx.ind_cols) >= 2: |
| if not (ctx.same_unit and ctx.can_share_axis): |
| ctx.register_refusal(ctx.hint, "Indicators represent different units or incompatible scales") |
| return None |
| return StrategyResult( |
| ChartStrategy.STACKED_BAR, |
| f"User requested stacked bar; {ctx.n_indicators} indicators → stacked bar chart", |
| indicator_cols=ctx.ind_cols, |
| color_dim="indicator", |
| ) |
| elif ctx.n_breakdowns == 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| return StrategyResult( |
| ChartStrategy.STACKED_BAR, |
| f"User requested stacked bar; 1 breakdown ({color_dim}) → stacked bar chart", |
| color_dim=color_dim, |
| ) |
| elif ctx.country_count > 1: |
| |
| |
| |
| if ctx.hint == "stacked_bar" or ctx.year_count > 1: |
| return StrategyResult( |
| ChartStrategy.STACKED_BAR, |
| f"User requested bar; {ctx.country_count} economies, {ctx.year_count} years → stacked bar chart (color=country)", |
| color_dim="country", |
| ) |
| return None |
|
|
| class ExplicitMapRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("map", "choropleth", "geoshape"): |
| if ctx.country_count > 0: |
| return StrategyResult( |
| ChartStrategy.CHOROPLETH, |
| f"User requested map; {ctx.country_count} economies → choropleth map", |
| indicator_cols=ctx.ind_cols, |
| ) |
| return None |
|
|
|
|
|
|
| class ExplicitDistributionRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("strip", "beeswarm", "tick", "distribution"): |
| |
| |
| |
| if ctx.country_count >= 10 and ctx.year_count <= 1: |
| return StrategyResult( |
| ChartStrategy.DISTRIBUTION, |
| f"User requested distribution; {ctx.country_count} economies, single year → strip/beeswarm", |
| color_dim="country", |
| ) |
| return None |
|
|
|
|
| class ExplicitSmallMultiplesRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("facet", "small_multiples"): |
| facet_dim = None |
| color_dim = None |
| if ctx.n_breakdowns >= 1: |
| facet_dim = list(ctx.breakdown_counts.keys())[0] |
| if ctx.n_breakdowns >= 2: |
| color_dim = list(ctx.breakdown_counts.keys())[1] |
| elif ctx.country_count > 1: |
| facet_dim = "country" |
| elif ctx.year_count > 1: |
| facet_dim = "year" |
|
|
| if facet_dim: |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"User requested small multiples; faceting by {facet_dim}", |
| facet_dim=facet_dim, |
| color_dim=color_dim, |
| ) |
| return None |
|
|
|
|
| def _axes_are_incompatible(df: "pd.DataFrame", ind_cols: list, scale_type: str | None = None) -> bool: |
| """Return True when indicators cannot meaningfully share a single Y-axis. |
| |
| Three complementary tests — all data-driven and domain-agnostic: |
| |
| 1. **Magnitude ratio ≥ 10×** (1.0 log₁₀ unit). |
| Standard dataviz threshold at which the smaller series becomes visually |
| compressed against the larger one. |
| |
| 2. **Sign-domain mismatch** — one indicator's *median* is negative while |
| another's is non-negative. |
| Growth rates, balances, and returns cross zero; levels and counts do not. |
| Placing them on a shared axis creates a misleading zero-reference. |
| |
| 3. **Range non-overlap** — the full [min, max] interval of one indicator |
| does not intersect the interval of another. |
| This catches cases where magnitude ratios are similar but the scales are |
| completely disjoint (e.g. GDP growth [-6, 14] vs life expectancy [51, 67]: |
| ratio is only ~5×, no negative medians, but the ranges are entirely |
| separate and combining them on one axis is meaningless). |
| """ |
| import math as _m |
| mags: list[float] = [] |
| medians: list[float] = [] |
| ranges: list[tuple] = [] |
|
|
| for col in ind_cols: |
| if col not in df.columns: |
| continue |
| s = pd.to_numeric(df[col], errors="coerce").dropna() |
| if s.empty: |
| continue |
| lo, hi = float(s.min()), float(s.max()) |
| ranges.append((lo, hi)) |
| medians.append(float(s.median())) |
| max_abs = float(s.abs().max()) |
| if max_abs > 0: |
| mags.append(_m.log10(max_abs)) |
|
|
| |
| valid_cols = [col for col in ind_cols if col in df.columns and not pd.to_numeric(df[col], errors="coerce").dropna().empty] |
| for i in range(len(mags)): |
| for j in range(i + 1, len(mags)): |
| col_i, col_j = valid_cols[i], valid_cols[j] |
| s_i = pd.to_numeric(df[col_i], errors="coerce").dropna() |
| s_j = pd.to_numeric(df[col_j], errors="coerce").dropna() |
| is_pct_i = any(x in col_i.lower() for x in ["%", "percent", "pct", "proportion", "share", "rate"]) and s_i.max() <= 100.0 and s_i.min() >= 0.0 |
| is_pct_j = any(x in col_j.lower() for x in ["%", "percent", "pct", "proportion", "share", "rate"]) and s_j.max() <= 100.0 and s_j.min() >= 0.0 |
| same_numeric_scale = not ((s_i.max() <= 1.0) ^ (s_j.max() <= 1.0)) |
|
|
| is_pct = (scale_type == "percentage") or (is_pct_i and is_pct_j and same_numeric_scale) |
| thresh = float("inf") if is_pct else 1.0 |
| if abs(mags[i] - mags[j]) >= thresh: |
| return True |
|
|
| |
| if len(medians) >= 2: |
| if any(m < 0 for m in medians) and any(m >= 0 for m in medians): |
| return True |
|
|
| |
| |
| |
| is_pct_scale = (scale_type == "percentage") |
| if not is_pct_scale: |
| for i in range(len(ranges)): |
| for j in range(i + 1, len(ranges)): |
| a, b = ranges[i] |
| c, d = ranges[j] |
| if b < c or d < a: |
| return True |
|
|
| return False |
|
|
|
|
|
|
| class TwoIndicatorRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.n_indicators == 2 and len(ctx.ind_cols) == 2: |
| if ctx.year_count <= 1 and ctx.country_count > 1: |
| if ctx.hint == "bar": |
| return StrategyResult( |
| ChartStrategy.BREAKDOWN_COMPARISON, |
| f"User requested bar; 2 indicators, {ctx.country_count} economies, single year → grouped bar", |
| indicator_cols=ctx.ind_cols, |
| color_dim="indicator", |
| ) |
| return StrategyResult( |
| ChartStrategy.CORRELATION, |
| f"2 indicators, {ctx.country_count} economies, single year → scatterplot", |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| x_dim=ctx.ind_cols[0], |
| y_dim=ctx.ind_cols[1], |
| ) |
| if ctx.year_count > 1 and ctx.country_count > 1: |
| |
| |
| |
| is_scatter_hint = ctx.hint == "point" |
| if ( |
| is_scatter_hint |
| and ctx.country_count <= CORRELATION_TEMPORAL_AUTO_MAX_COUNTRIES |
| and ctx.year_count <= CORRELATION_TEMPORAL_AUTO_MAX_YEARS |
| ): |
| return StrategyResult( |
| ChartStrategy.CORRELATION_TEMPORAL, |
| ( |
| f"2 indicators, {ctx.country_count} economies, " |
| f"{ctx.year_count} years → connected scatter " |
| f"(≤{CORRELATION_TEMPORAL_AUTO_MAX_COUNTRIES} countries " |
| f"× ≤{CORRELATION_TEMPORAL_AUTO_MAX_YEARS} years threshold)" |
| ), |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| ) |
| try: |
| is_incompatible = _axes_are_incompatible(ctx.df, ctx.ind_cols, ctx.scale_type) |
| except Exception: |
| is_incompatible = False |
|
|
| if is_incompatible: |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"2 indicators (scale-incompatible), {ctx.country_count} economies, {ctx.year_count} years → small multiples (facet=indicator, color=country)", |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| facet_dim="indicator", |
| scale_incompatible=True, |
| ) |
|
|
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| ( |
| f"2 indicators (scale-compatible), " |
| f"{ctx.country_count} economies, {ctx.year_count} years → small multiples (facet=indicator, color=country)" |
| ), |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| facet_dim="indicator", |
| ) |
| if ctx.year_count <= 1: |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| "2 indicators, 1 country, single year → small multiples (facet=indicator, color=indicator)", |
| indicator_cols=ctx.ind_cols, |
| color_dim="indicator", |
| facet_dim="indicator", |
| ) |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"2 indicators, 1 country, {ctx.year_count} years → small multiples (facet=indicator, color=indicator)", |
| indicator_cols=ctx.ind_cols, |
| color_dim="indicator", |
| facet_dim="indicator", |
| ) |
| return None |
|
|
| class ThreePlusIndicatorRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.n_indicators >= 2 and len(ctx.ind_cols) >= 2: |
| mark_hint = "bar" if ctx.year_count <= 1 else "line" |
|
|
| |
| |
| |
| if ctx.year_count > 1 and ctx.country_count >= 1: |
| try: |
| is_incompatible = _axes_are_incompatible(ctx.df, ctx.ind_cols, ctx.scale_type) |
| except Exception: |
| is_incompatible = False |
|
|
| if is_incompatible: |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| ( |
| f"{ctx.n_indicators} indicators (scale-incompatible), " |
| f"{ctx.country_count} economies, {ctx.year_count} years " |
| f"→ small multiples (facet=indicator, color=country)" |
| ), |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| facet_dim="indicator", |
| scale_incompatible=True, |
| ) |
| elif ctx.country_count > 1: |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| ( |
| f"{ctx.n_indicators} indicators, " |
| f"{ctx.country_count} economies, {ctx.year_count} years " |
| f"→ small multiples (facet=indicator, color=country)" |
| ), |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| facet_dim="indicator", |
| ) |
|
|
| if ctx.country_count > 1: |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| ( |
| f"{ctx.n_indicators} indicators, " |
| f"{ctx.country_count} economies → " |
| f"small multiples (facet=indicator, color=country)" |
| ), |
| indicator_cols=ctx.ind_cols, |
| color_dim="country", |
| facet_dim="indicator", |
| ) |
|
|
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"{ctx.n_indicators} indicators, 1 country → small multiples (facet=indicator, color=indicator)", |
| indicator_cols=ctx.ind_cols, |
| color_dim="indicator", |
| facet_dim="indicator", |
| ) |
| return None |
|
|
|
|
| class StackedAreaRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint in ("area", "stacked_area") and "unit_measure" not in ctx.breakdown_counts: |
| if ctx.n_indicators >= 2: |
| if not (ctx.same_unit and ctx.can_share_axis): |
| ctx.register_refusal(ctx.hint, "Indicators represent different units or incompatible scales") |
| return None |
| if ctx.year_count <= 1: |
| ctx.register_refusal(ctx.hint, "Area charts require multiple years of data") |
| return None |
| if ctx.year_count > 1: |
| color_dim = None |
| if ctx.breakdown_counts: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| elif ctx.country_count > 1: |
| color_dim = "country" |
| return StrategyResult( |
| ChartStrategy.STACKED_AREA, |
| f"User requested area; {ctx.year_count} years → stacked area chart", |
| color_dim=color_dim, |
| ) |
| return None |
|
|
| class ExplicitHeatmapRule: |
| """Honour an explicit ``heatmap`` hint from the caller. |
| |
| Requires multiple time periods so there is a meaningful country-x-year |
| matrix. Single-year requests fall through to the auto-routing rules |
| (usually CrossSectional / BreakdownComparison). |
| """ |
|
|
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint == "heatmap" and ctx.year_count > 1 and ctx.country_count > 1: |
| return StrategyResult( |
| ChartStrategy.HEATMAP, |
| f"User requested heatmap; {ctx.country_count} economies, {ctx.year_count} years → heatmap", |
| color_dim="value", |
| ) |
| return None |
|
|
|
|
| class HeatmapRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.country_count >= HIGH_CARDINALITY_THRESHOLDS["heatmap_threshold"] and ctx.year_count > 1: |
| if ctx.n_breakdowns == 0: |
| return StrategyResult( |
| ChartStrategy.HEATMAP, |
| f"{ctx.country_count} economies, {ctx.year_count} years → heatmap", |
| color_dim="value", |
| ) |
| return None |
|
|
| class ExplicitBarBreakdownRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint == "bar" and ctx.n_breakdowns == 1 and 0 < ctx.country_count <= 4: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| return StrategyResult( |
| ChartStrategy.BREAKDOWN_COMPARISON, |
| f"User requested bar; 1 breakdown ({color_dim}), {ctx.country_count} economies → grouped bar", |
| color_dim=color_dim, |
| ) |
| return None |
|
|
| class IncompatibleUnitsRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if "unit_measure" in ctx.breakdown_counts: |
| facet_dim = "unit_measure" |
| other_breakdowns = [k for k in ctx.breakdown_counts if k != "unit_measure"] |
| color_dim = other_breakdowns[0] if other_breakdowns else None |
| secondary_color_dim = "country" if ctx.country_count > 1 else None |
| n_other = len(other_breakdowns) |
|
|
| unit_count = ctx.breakdown_counts["unit_measure"] |
| reason_detail = ( |
| f"unit_measure ({unit_count} units)" |
| + (f" + {n_other} other breakdown(s)" if n_other else "") |
| + f", {ctx.country_count} econom{'y' if ctx.country_count == 1 else 'ies'}" |
| + (" + country combo" if secondary_color_dim else "") |
| + " → faceted by unit (independent Y-axes)" |
| ) |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| reason_detail, |
| color_dim=color_dim, |
| facet_dim=facet_dim, |
| secondary_color_dim=secondary_color_dim, |
| scale_incompatible=True, |
| ) |
| return None |
|
|
| class IncompatibleCustomBreakdownRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.n_breakdowns == 1 and ctx.country_count > 1: |
| bd_dim = list(ctx.breakdown_counts.keys())[0] |
| if bd_dim in _CUSTOM_BREAKDOWN_DIMS and _detect_scale_incompatibility(ctx.df, bd_dim): |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"1 breakdown ({bd_dim}), {ctx.country_count} economies, scale-incompatible " |
| f"→ scale-split panels (color=economy)", |
| color_dim="country", |
| facet_dim=bd_dim, |
| scale_incompatible=True, |
| ) |
| return None |
|
|
| class GenericSmallMultiplesRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.n_breakdowns >= 2 or (ctx.n_breakdowns >= 1 and ctx.country_count > 1): |
| if ctx.country_count > 1: |
| facet_dim = "country" |
| color_dim = list(ctx.breakdown_counts.keys())[0] if ctx.breakdown_counts else None |
| elif ctx.n_breakdowns >= 2: |
| bd_keys = list(ctx.breakdown_counts.keys()) |
| facet_dim = bd_keys[0] |
| color_dim = bd_keys[1] if len(bd_keys) >= 2 else None |
| else: |
| facet_dim = list(ctx.breakdown_counts.keys())[0] |
| color_dim = None |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"{ctx.n_breakdowns} breakdowns, {ctx.country_count} econom{'y' if ctx.country_count == 1 else 'ies'} " |
| f"→ small multiples (facet={facet_dim}, color={color_dim})", |
| color_dim=color_dim, |
| facet_dim=facet_dim, |
| ) |
| return None |
|
|
| class TemporalBreakdownRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.n_breakdowns == 1 and ctx.year_count > 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| if color_dim in _CUSTOM_BREAKDOWN_DIMS and _detect_scale_incompatibility(ctx.df, color_dim): |
| return StrategyResult( |
| ChartStrategy.SMALL_MULTIPLES, |
| f"1 breakdown ({color_dim}), {ctx.year_count} years, scale-incompatible " |
| f"→ faceted (independent Y-axes)", |
| color_dim=None, |
| facet_dim=color_dim, |
| scale_incompatible=True, |
| ) |
| return StrategyResult( |
| ChartStrategy.TEMPORAL_SINGLE, |
| f"1 breakdown ({color_dim}), {ctx.year_count} years → multi-series line chart", |
| color_dim=color_dim, |
| ) |
| return None |
|
|
| class BreakdownComparisonGroupedBarRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.n_breakdowns == 1 and ctx.country_count <= 4 and ctx.year_count <= 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| return StrategyResult( |
| ChartStrategy.BREAKDOWN_COMPARISON, |
| f"1 breakdown ({color_dim}), {ctx.breakdown_counts[color_dim]} values, single year → grouped bar", |
| color_dim=color_dim, |
| ) |
| return None |
|
|
| class ExplicitBarCrossSectionalRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.hint == "bar" and ctx.year_count <= 1 and ctx.country_count > 0: |
| return StrategyResult( |
| ChartStrategy.CROSS_SECTIONAL, |
| f"User requested bar; {ctx.country_count} economies, single year → horizontal bar", |
| color_dim="country", |
| ) |
| return None |
|
|
| class HighCardinalityCrossSectionalRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.country_count > HIGH_CARDINALITY_THRESHOLDS["beeswarm_threshold"] and ctx.year_count <= 1: |
| if ctx.hint in ("strip", "beeswarm", "tick", "distribution"): |
| return StrategyResult( |
| ChartStrategy.DISTRIBUTION, |
| f"User requested strip/beeswarm; {ctx.country_count} economies, single year → strip/beeswarm", |
| color_dim="country", |
| ) |
| return StrategyResult( |
| ChartStrategy.CHOROPLETH, |
| f"{ctx.country_count} economies, single year → default to choropleth map", |
| ) |
| return None |
|
|
| class CrossSectionalRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if (ctx.year_count <= 1 or ctx.avg_years_per_country < 1.7) and ctx.country_count > 0: |
| return StrategyResult( |
| ChartStrategy.CROSS_SECTIONAL, |
| f"{ctx.country_count} economies, sparse or single year (avg {ctx.avg_years_per_country:.2f} yrs) → horizontal bar", |
| color_dim="country", |
| ) |
| return None |
|
|
| class TemporalSingleRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult | None: |
| if ctx.year_count > 1 and ctx.avg_years_per_country >= 1.7: |
| phrase = chart_type_phrase_for_reason(ctx.hint) |
| return StrategyResult( |
| ChartStrategy.TEMPORAL_SINGLE, |
| f"Single indicator, {ctx.year_count} years, {ctx.country_count} economies → {phrase}", |
| color_dim="country" if ctx.country_count > 0 else None, |
| mark_hint=ctx.hint if ctx.hint in ("bar", "line") else None, |
| ) |
| return None |
|
|
| class FallbackRule: |
| def evaluate(self, ctx: RoutingContext) -> StrategyResult: |
| return StrategyResult( |
| ChartStrategy.FALLBACK_LINE, |
| f"Default fallback → {chart_type_phrase_for_reason(ctx.hint)}", |
| color_dim="country" if ctx.country_count > 0 else None, |
| ) |
|
|
| ROUTING_RULES: list[RoutingRule] = [ |
| ExplicitScatterRule(), |
| ExplicitStackedAreaMultiIndicatorRule(), |
| ExplicitStackedBarRule(), |
| ExplicitMapRule(), |
| ExplicitSmallMultiplesRule(), |
| ExplicitHeatmapRule(), |
| ExplicitDistributionRule(), |
| TwoIndicatorRule(), |
| ThreePlusIndicatorRule(), |
| StackedAreaRule(), |
| HeatmapRule(), |
| ExplicitBarBreakdownRule(), |
| IncompatibleUnitsRule(), |
| IncompatibleCustomBreakdownRule(), |
| GenericSmallMultiplesRule(), |
| TemporalBreakdownRule(), |
| BreakdownComparisonGroupedBarRule(), |
| ExplicitBarCrossSectionalRule(), |
| HighCardinalityCrossSectionalRule(), |
| CrossSectionalRule(), |
| TemporalSingleRule(), |
| ] |
|
|
| def select_strategy( |
| df: pd.DataFrame, |
| n_indicators: int = 1, |
| chart_type_hint: str | None = None, |
| indicator_cols: list[str] | None = None, |
| raw_unit: str | None = None, |
| raw_unit_mult: int = 0, |
| data_profile: dict | None = None, |
| strategy_override: str | None = None, |
| ) -> StrategyResult: |
| """ |
| Applies the Rule Engine to select the correct ChartStrategy. |
| |
| Parameters |
| ---------- |
| data_profile : dict | None |
| Pre-computed data profile from ``_build_data_profile``. When provided, |
| coverage quality signals (sparse country count, completeness %) are |
| injected into the ``RoutingContext`` so routing rules can use real |
| sparsity data rather than structural heuristics alone. |
| strategy_override : str | None |
| Explicit ChartStrategy value or name (e.g. "stacked_bar", "vconcat_panels"). |
| If provided and valid, bypasses the rule engine and directly forces |
| the requested strategy, using the routing context to configure sensible |
| color and facet defaults. |
| """ |
| ctx = RoutingContext.build( |
| df, |
| n_indicators, |
| chart_type_hint, |
| indicator_cols, |
| raw_unit=raw_unit, |
| raw_unit_mult=raw_unit_mult |
| ) |
|
|
| |
| if data_profile: |
| cq = data_profile.get("coverage_quality", {}) |
| ctx.sparse_country_count = len(cq.get("sparse_countries", [])) |
| ctx.completeness_pct = cq.get("completeness_pct", 100.0) |
|
|
| scale_comp = data_profile.get("scale_compatibility", {}) |
| ctx.same_unit = scale_comp.get("same_unit", True) |
| ctx.can_share_axis = scale_comp.get("can_share_axis", True) |
|
|
| inds = data_profile.get("indicators", []) |
| if inds and all(ind.get("scale_type") == "percentage" for ind in inds): |
| ctx.scale_type = "percentage" |
|
|
| res: StrategyResult | None = None |
|
|
| |
| if strategy_override: |
| matched_strategy = None |
| for s in ChartStrategy: |
| if s.value.lower() == strategy_override.lower() or s.name.lower() == strategy_override.lower(): |
| matched_strategy = s |
| break |
| if matched_strategy: |
| res = _build_overridden_strategy(matched_strategy, ctx) |
|
|
| |
| if res is None: |
| for rule in ROUTING_RULES: |
| res = rule.evaluate(ctx) |
| if res is not None: |
| break |
| if res is None: |
| res = FallbackRule().evaluate(ctx) |
|
|
| res.scale_type = ctx.scale_type |
| res.unit_mult = ctx.unit_mult |
| res.raw_hint = ctx.raw_hint |
| res.refusal_reason = ctx.refusal_reason |
|
|
| |
| if data_profile is not None: |
| res.data_profile = data_profile |
|
|
| |
| if ctx.indicator_value_ranges and len(ctx.indicator_value_ranges) >= 2: |
| maxes = [r["max"] for r in ctx.indicator_value_ranges.values() if r["max"] > 0] |
| mins_nz = [r["min"] for r in ctx.indicator_value_ranges.values() if r["min"] > 0] |
| if maxes and mins_nz: |
| ratio = max(maxes) / min(mins_nz) |
| can_share = ratio <= 10.0 |
| res.scale_compatibility = { |
| "max_min_ratio": round(ratio, 2), |
| "can_share_axis": can_share, |
| "reason": ( |
| f"Value ratio {ratio:.1f}x — " |
| + ("within" if can_share else "exceeds") |
| + " 10x threshold" |
| ), |
| } |
|
|
| return res |
|
|
|
|
| def _build_overridden_strategy(strategy: ChartStrategy, ctx: RoutingContext) -> StrategyResult: |
| """Build a StrategyResult for a user/LLM overridden strategy, populating sensible defaults.""" |
| color_dim = None |
| facet_dim = None |
| indicator_cols = ctx.ind_cols |
|
|
| if strategy in (ChartStrategy.STACKED_AREA, ChartStrategy.STACKED_BAR): |
| if len(indicator_cols) >= 2: |
| color_dim = "indicator" |
| elif ctx.n_breakdowns == 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| elif ctx.country_count > 1: |
| color_dim = "country" |
|
|
| elif strategy == ChartStrategy.TEMPORAL_SINGLE: |
| if ctx.n_breakdowns == 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
| elif ctx.country_count > 1: |
| color_dim = "country" |
|
|
| elif strategy == ChartStrategy.SMALL_MULTIPLES: |
| if len(indicator_cols) >= 2: |
| facet_dim = "indicator" |
| color_dim = "indicator" |
| elif ctx.country_count > 1: |
| facet_dim = "country" |
| if ctx.n_breakdowns == 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
|
|
| elif strategy == ChartStrategy.BREAKDOWN_COMPARISON: |
| if len(indicator_cols) >= 2: |
| color_dim = "indicator" |
| elif ctx.n_breakdowns == 1: |
| color_dim = list(ctx.breakdown_counts.keys())[0] |
|
|
| elif strategy in (ChartStrategy.CORRELATION, ChartStrategy.CORRELATION_TEMPORAL): |
| if ctx.country_count > 0: |
| color_dim = "country" |
|
|
| return StrategyResult( |
| strategy=strategy, |
| reason=f"Forced via strategy_override: {strategy.value}", |
| indicator_cols=indicator_cols, |
| color_dim=color_dim, |
| facet_dim=facet_dim, |
| ) |
|
|
|
|
| def explain_chart_routing( |
| n_indicators: int, |
| country_count: int, |
| year_count: int, |
| avg_years_per_country: float, |
| breakdown_dims: list[str] | None = None, |
| chart_type_hint: str | None = None, |
| scale_type: str | None = None, |
| indicator_scales: list[dict] | None = None, |
| ) -> dict: |
| """Explain which chart strategy the routing engine would select for the given data shape. |
| |
| Synthesises a minimal DataFrame from the structural descriptors, runs the full |
| rule waterfall, checks multi-indicator scale compatibility, and returns a |
| structured explanation the LLM can use to decide *which* viz tool to call and |
| *how* to configure it — before committing to a potentially expensive data fetch. |
| |
| Args: |
| n_indicators: Number of distinct indicators (1, 2, or 3+). |
| country_count: Number of unique countries in the data. |
| year_count: Total number of unique time periods in the data. |
| avg_years_per_country: Mean unique years per country |
| (use year_count for single-country data; lower if countries have sparse coverage). |
| breakdown_dims: Disaggregation dimensions present in the data |
| (e.g. ["sex"], ["age", "urbanisation"]). Omit or pass [] for none. |
| chart_type_hint: Optional explicit chart type requested by the user |
| (e.g. "heatmap", "bar", "scatter"). Pass None to use auto-routing. |
| scale_type: Shared unit type for single-indicator data |
| (e.g. "percentage", "currency", "persons", "index"). |
| indicator_scales: Per-indicator scale info for multi-indicator data. |
| Each entry: {"label": str, "scale_type": str, "approx_max": float}. |
| Used to determine whether indicators can share a Y-axis (ratio <= 10x). |
| |
| Returns: |
| dict with keys: |
| strategy : ChartStrategy value string (e.g. "temporal_single") |
| layout : "single_panel" | "vconcat_panels" (multi-indicator only) |
| reason : Human-readable explanation of why this strategy was chosen |
| recommended_viz_tool : "data360_get_viz_spec" | "data360_get_multi_indicator_viz_spec" |
| chart_type_hint : Suggested chart_type argument value, or null |
| scale_notes : Notes on scale compatibility for multi-indicator data |
| routing_inputs : Echo of the structural inputs used for reproducibility |
| """ |
| import numpy as np |
|
|
| |
| countries = [f"C{i}" for i in range(max(country_count, 1))] |
| years = list(range(2020, 2020 + max(year_count, 1))) |
|
|
| |
| rows = [] |
| for c in countries: |
| c_years = years[:max(1, round(avg_years_per_country))] |
| for y in c_years: |
| row: dict = {"country": c, "year": y, "value": float(np.random.uniform(10, 100))} |
| for dim in (breakdown_dims or []): |
| row[dim] = f"{dim}_val" |
| rows.append(row) |
| df = pd.DataFrame(rows) if rows else pd.DataFrame({"country": ["C0"], "year": [2020], "value": [50.0]}) |
|
|
| |
| ind_cols: list[str] | None = None |
| if n_indicators >= 2 and indicator_scales: |
| ind_cols = [] |
| for i, ind in enumerate(indicator_scales[:n_indicators]): |
| col_name = f"IND_{i}" |
| approx_max = float(ind.get("approx_max", 100.0)) |
| df[col_name] = [float(np.random.uniform(approx_max * 0.5, approx_max)) for _ in range(len(df))] |
| ind_cols.append(col_name) |
|
|
| |
| raw_unit = scale_type or "" |
| result = select_strategy( |
| df, |
| n_indicators=n_indicators, |
| chart_type_hint=chart_type_hint, |
| indicator_cols=ind_cols, |
| raw_unit=raw_unit, |
| ) |
| strategy_value = result.strategy.value |
|
|
| |
| layout = "single_panel" |
| scale_notes = "" |
| if n_indicators >= 2 and indicator_scales and ind_cols: |
| try: |
| max_vals = [float(ind.get("approx_max", 100.0)) for ind in indicator_scales[:n_indicators]] |
| if max_vals and min(max_vals) > 0: |
| ratio = max(max_vals) / min(max_vals) |
| if ratio <= 10.0: |
| layout = "single_panel" |
| scale_notes = ( |
| f"Indicators are scale-compatible (max ratio {ratio:.1f}x ≤ 10x). " |
| f"They will share a single Y-axis in one panel." |
| ) |
| else: |
| layout = "vconcat_panels" |
| scale_notes = ( |
| f"Indicators differ in scale by {ratio:.0f}x (> 10x threshold). " |
| f"They will be split into stacked panels with independent Y-axes." |
| ) |
| else: |
| layout = "vconcat_panels" |
| scale_notes = "Could not determine scale ratio — defaulting to stacked panels." |
| except Exception: |
| layout = "vconcat_panels" |
| scale_notes = "Scale check failed — defaulting to stacked panels." |
| elif strategy_value == "temporal_multi_indicator": |
| scale_notes = ( |
| "No per-indicator scale info provided. Pass indicator_scales to determine " |
| "whether indicators can share a Y-axis." |
| ) |
|
|
| |
| recommended_tool = ( |
| "data360_get_multi_indicator_viz_spec" |
| if n_indicators >= 2 |
| else "data360_get_viz_spec" |
| ) |
|
|
| |
| suggested_hint: str | None = chart_type_hint |
| if not chart_type_hint: |
| |
| hint_map = { |
| "cross_sectional": "bar", |
| "heatmap": "heatmap", |
| "distribution": "strip", |
| "correlation": "scatter", |
| "correlation_temporal": "connected_scatter", |
| "stacked_area": "stacked_area", |
| } |
| suggested_hint = hint_map.get(strategy_value) |
|
|
| return { |
| "strategy": strategy_value, |
| "layout": layout, |
| "reason": result.reason, |
| "recommended_viz_tool": recommended_tool, |
| "chart_type_hint": suggested_hint, |
| "scale_notes": scale_notes, |
| "routing_inputs": { |
| "n_indicators": n_indicators, |
| "country_count": country_count, |
| "year_count": year_count, |
| "avg_years_per_country": avg_years_per_country, |
| "breakdown_dims": breakdown_dims or [], |
| "chart_type_hint": chart_type_hint, |
| "scale_type": scale_type, |
| "indicator_scales": indicator_scales or [], |
| }, |
| } |
|
|
|
|
| |
| |
| |
|
|
|
|
|
|
| def _vl_schema() -> str: |
| return "https://vega.github.io/schema/vega-lite/v5.json" |
|
|
|
|
| def _axis_style(title: str | None = None, temporal: bool = False) -> dict: |
| ax: dict = { |
| "grid": False, |
| "gridColor": WB_GRID_COLOR, |
| "gridDash": [4, 2], |
| "labelColor": WB_TEXT_SUBTLE, |
| "titleColor": WB_TEXT, |
| "titleFontWeight": "bold", |
| "tickCount": 5, |
| } |
| if temporal: |
| ax["title"] = None |
| ax["format"] = "%Y" |
| ax["tickCount"] = 5 |
| ax["labelAngle"] = 0 |
| elif title is not None: |
| ax["title"] = title |
| return ax |
|
|
|
|
| def _resolve_axis_title( |
| y_label: str | None, |
| indicator_name: str | None, |
| ) -> str | None: |
| """Return the best available axis title for a value axis, using only the short unit of measure.""" |
| _GENERIC = {"value", ""} |
| unit_ok = y_label is not None and y_label.lower() not in _GENERIC |
| if unit_ok: |
| return _clean_unit_label(y_label) |
| if indicator_name: |
| return indicator_name |
| return None |
|
|
|
|
| def _detect_temporal_frequency(series: pd.Series) -> TemporalFreq: |
| """Infer temporal frequency from raw TIME_PERIOD values. |
| |
| Handles all formats the Data360 API produces: |
| - Annual: "2019", "2020" |
| - Monthly: "2019-09", "2019-09-01", "2019M09" |
| - Quarterly: "2019-Q1", "2019Q1", "2019-q1" |
| - Daily: "2019-09-15" |
| |
| Logic: |
| - Parse unique values as datetime. If all land on Jan-1 (or are bare |
| 4-digit integers), treat as annual. |
| - If distinct parsed periods show > 1 period per year → sub-annual. |
| Distinguish quarterly (avg ~4/year) vs monthly (avg ~12/year). |
| - Fall back to annual on any parse error or empty series. |
| """ |
| values = series.dropna().astype(str).unique() |
| if len(values) == 0: |
| return "annual" |
|
|
| |
| if all(v.strip().isdigit() and len(v.strip()) == 4 for v in values): |
| return "annual" |
|
|
| |
| q_pattern = re.compile(r"\d{4}[-\s]?[Qq]\d", re.IGNORECASE) |
| if any(q_pattern.search(v) for v in values): |
| return "quarterly" |
|
|
| |
| try: |
| parsed = pd.to_datetime(pd.Series(values), errors="coerce").dropna() |
| except Exception: |
| return "annual" |
|
|
| if parsed.empty: |
| return "annual" |
|
|
| |
| if (parsed.dt.month == 1).all() and (parsed.dt.day == 1).all(): |
| return "annual" |
|
|
| n_years = max(parsed.dt.year.nunique(), 1) |
| |
| |
| n_year_months = parsed.dt.to_period("M").nunique() |
| avg_months_per_year = n_year_months / n_years |
|
|
| |
| |
| if avg_months_per_year >= 5: |
| return "monthly" |
| if avg_months_per_year >= 3: |
| return "quarterly" |
| |
| return "annual" |
|
|
|
|
| def _format_time_period_series( |
| series: pd.Series, freq: TemporalFreq |
| ) -> pd.Series: |
| """Convert raw TIME_PERIOD strings to the correct format for a given frequency. |
| |
| Annual → "2019" (4-digit year string; Vega-Lite ordinal) |
| Monthly → "2019-09" (ISO yearmonth; Vega-Lite temporal + timeUnit yearmonth) |
| Quarterly→ "2019-Q1" (ISO yearquarter; Vega-Lite temporal + timeUnit yearquarter) |
| Daily → "2019-09-15" (ISO date; Vega-Lite temporal) |
| |
| Values that fail parsing are left as-is (graceful fallback). |
| """ |
| try: |
| parsed = pd.to_datetime(series.astype(str), errors="coerce") |
| except Exception: |
| return series |
|
|
| if freq == "annual": |
| return parsed.dt.year.astype("Int64").astype(str).where(parsed.notna(), series) |
| elif freq == "monthly": |
| return parsed.dt.to_period("M").astype(str).where(parsed.notna(), series) |
| elif freq == "quarterly": |
| return parsed.dt.to_period("Q").astype(str).where(parsed.notna(), series) |
| else: |
| return parsed.dt.strftime("%Y-%m-%d").where(parsed.notna(), series) |
|
|
|
|
| |
| |
| |
| _TEMPORAL_X_ENCODING: dict[TemporalFreq, dict] = { |
| "annual": { |
| "field": "year", |
| "type": "temporal", |
| "timeUnit": "utcyear", |
| "axis": { |
| "title": None, |
| "format": "%Y", |
| "tickCount": 5, |
| "labelAngle": 0, |
| "grid": False, |
| "gridColor": WB_GRID_COLOR, |
| "gridDash": [4, 2], |
| "labelColor": WB_TEXT_SUBTLE, |
| "titleColor": WB_TEXT, |
| "titleFontWeight": "bold", |
| }, |
| }, |
| "monthly": { |
| "field": "year", |
| "type": "temporal", |
| "timeUnit": "utcyearmonth", |
| "axis": { |
| "title": None, |
| "format": "%b %Y", |
| "tickCount": 8, |
| "labelAngle": -45, |
| "grid": False, |
| "gridColor": WB_GRID_COLOR, |
| "gridDash": [4, 2], |
| "labelColor": WB_TEXT_SUBTLE, |
| "titleColor": WB_TEXT, |
| "titleFontWeight": "bold", |
| }, |
| }, |
| "quarterly": { |
| "field": "year", |
| "type": "temporal", |
| "timeUnit": "utcyearquarter", |
| "axis": { |
| "title": None, |
| "format": "Q%q %Y", |
| "tickCount": 6, |
| "labelAngle": -45, |
| "grid": False, |
| "gridColor": WB_GRID_COLOR, |
| "gridDash": [4, 2], |
| "labelColor": WB_TEXT_SUBTLE, |
| "titleColor": WB_TEXT, |
| "titleFontWeight": "bold", |
| }, |
| }, |
| "daily": { |
| "field": "year", |
| "type": "temporal", |
| "timeUnit": "utcyearmonthdate", |
| "axis": { |
| "title": None, |
| "format": "%d %b %Y", |
| "tickCount": 6, |
| "labelAngle": -45, |
| "grid": False, |
| "gridColor": WB_GRID_COLOR, |
| "gridDash": [4, 2], |
| "labelColor": WB_TEXT_SUBTLE, |
| "titleColor": WB_TEXT, |
| "titleFontWeight": "bold", |
| }, |
| }, |
| } |
|
|
|
|
| def _x_temporal_encoding(freq: TemporalFreq = "annual") -> dict: |
| """Return the correct Vega-Lite x encoding for the given temporal frequency. |
| |
| Returns a deep copy so callers can mutate axis overrides (e.g. labels=False |
| for non-bottom panels) without affecting subsequent calls. |
| """ |
| import copy |
| return copy.deepcopy(_TEMPORAL_X_ENCODING.get(freq, _TEMPORAL_X_ENCODING["annual"])) |
|
|
|
|
| def _is_proportion_indicator(df, unit_measure=None, scale_type=None): |
| """Check if the indicator unit, scale_type, or name suggests it's a proportion/rate/share.""" |
| normalized_unit = (unit_measure or "").upper() |
| if scale_type == "percentage" or (("%" in normalized_unit or "PERCENT" in normalized_unit) and "PERSON" not in normalized_unit): |
| return True |
|
|
| keywords = {"PROPORTION", "SHARE", "RATE", "RATIO", "FRACTION"} |
| if any(k in normalized_unit for k in keywords): |
| return True |
|
|
| if df is not None and "indicator" in df.columns: |
| ind_names = df["indicator"].dropna().unique() |
| for name in ind_names: |
| name_upper = str(name).upper() |
| if any(k in name_upper for k in keywords) or "PERCENT" in name_upper or "%" in name_upper: |
| return True |
|
|
| return False |
|
|
|
|
| def _value_label_expr(unit_measure: str | None = None, scale_type: str | None = None) -> str: |
| """Vega expression for custom k/m/b/t axis label formatting.""" |
| normalized = (unit_measure or "").upper().strip() |
| if scale_type == "proportion": |
| return "format(datum.value, '.0%')" |
| is_currency = scale_type == "currency" or "$" in normalized or "USD" in normalized |
| prefix = "$" if is_currency else "" |
| is_percentage = scale_type == "percentage" or (("%" in normalized or "PERCENT" in normalized) and "PERSON" not in normalized) |
| if is_percentage: |
| return "format(datum.value, '.1~f') + '%'" |
| if unit_measure == "T": |
| tiers = [("1e12", "Gt"), ("1e9", "Mt"), ("1e6", "Kt")] |
| elif unit_measure == "W_POP": |
| tiers = [("1e12", "Gw"), ("1e9", "Mw"), ("1e6", "Kw")] |
| elif unit_measure in ("BITS", "BIT_S_IU"): |
| tiers = [("1e12", "Gb"), ("1e9", "Mb"), ("1e6", "Kb")] |
| else: |
| tiers = [("1e12", "t"), ("1e9", "b"), ("1e6", "m"), ("1e3", "k")] |
| parts = [ |
| f"abs(datum.value)>={t} ? '{prefix}'+format(datum.value/{t},'.1~f')+'{s}'" |
| for t, s in tiers |
| ] |
| tail = ( |
| f" : abs(datum.value)>=10 ? '{prefix}'+format(datum.value,',.1~f')" |
| f" : abs(datum.value)>=1 ? '{prefix}'+format(datum.value,'.1~f')" |
| f" : '{prefix}'+format(datum.value,'.2~f')" |
| ) |
| return " : ".join(parts) + tail |
|
|
|
|
| def _compute_tooltip_format( |
| max_abs: float | None = None, unit_measure: str | None = None |
| ) -> str: |
| """Returns D3 format string for tooltip quantitative fields. |
| |
| Avoids D3's SI-prefix format (~s) which uses G/M/k (giga/mega/kilo) — |
| these conflict with our axis labelExpr which uses b/m/k (billion/million/thousand). |
| Large values are formatted as plain integers with comma separators instead. |
| """ |
| normalized = (unit_measure or "").upper().strip() |
| if ("PROPORTION" in normalized or "SHARE" in normalized) and max_abs is not None and 0.0 < max_abs <= 1.0: |
| return ".1%" |
| if unit_measure == "%" or "%" in normalized or "PERCENT" in normalized: |
| return ".1f" |
| if "$" in normalized or "USD" in normalized: |
| return "$,.2f" |
| if max_abs is None or max_abs < 1: |
| return ".2f" |
| if max_abs < 10: |
| return ".1f" |
| if max_abs < 1000: |
| return ",.1f" |
| |
| |
| return ",.0f" |
|
|
|
|
| def _color_encoding( |
| field: str, |
| domain: list | None = None, |
| mark_type: str = "point", |
| n_items: int = 0, |
| legend_title: str | None = None, |
| domain_labels: list[str] | None = None, |
| ) -> dict: |
| """Build a Vega-Lite color encoding channel. |
| |
| Legend title resolves in this priority order: |
| 1. Caller-supplied ``legend_title`` |
| 2. Human-readable label from ``_TOOLTIP_SPECS`` (e.g. "Dimension 1") |
| 3. Title-cased field name (e.g. "Comp Breakdown 2") |
| |
| Legend orientation is chosen dynamically: |
| - When the longest label in ``domain_labels`` exceeds 40 chars the legend |
| switches to ``orient: bottom`` / ``direction: vertical`` so labels are not |
| truncated and are not clipped in narrow containers (e.g. chatbot panels). |
| - Otherwise ``orient: top`` / ``direction: horizontal`` is used. |
| """ |
| resolved_title = ( |
| legend_title |
| or _TOOLTIP_SPECS.get(field, {}).get("title") |
| or field.replace("_", " ").title() |
| ) |
| scale = {"range": WB_CAT_COLORS} |
| if domain: |
| sort_order = _get_dimension_sort_order(field, [str(d) for d in domain]) |
| if sort_order: |
| domain = sorted(domain, key=lambda x: sort_order.index(str(x)) if str(x) in sort_order else len(sort_order)) |
| scale["domain"] = domain |
| |
| if n_items == 1 and field != "country": |
| legend = None |
| else: |
| _LONG_LABEL_THRESHOLD = 40 |
| max_label_len = ( |
| max((len(lbl) for lbl in domain_labels), default=0) |
| if domain_labels |
| else 0 |
| ) |
| if max_label_len > _LONG_LABEL_THRESHOLD: |
| legend: dict | None = { |
| "orient": "right", |
| "direction": "vertical", |
| "title": resolved_title, |
| "labelLimit": 1000, |
| } |
| else: |
| legend = { |
| "orient": "right", |
| "direction": "vertical", |
| "title": resolved_title, |
| "labelLimit": 250, |
| } |
| if mark_type == "line": |
| legend["symbolType"] = "stroke" |
| return { |
| "field": field, |
| "type": "nominal", |
| "scale": scale, |
| "legend": legend, |
| } |
|
|
|
|
| def _adjust_end_label_y(df: pd.DataFrame, color_dim: str) -> pd.DataFrame: |
| """Compute a '_label_y' column in df to prevent direct end labels from overlapping. |
| |
| Uses a 1D relaxation (spring/force) algorithm on the final year's data values. |
| """ |
| if df.empty or "value" not in df.columns or "year" not in df.columns or not color_dim or color_dim not in df.columns: |
| return df |
|
|
| |
| df = df.copy() |
| df["_label_y"] = df["value"] |
|
|
| try: |
| |
| df = df.reset_index(drop=True) |
| |
| last_indices = df.groupby(color_dim)["year"].idxmax() |
| last_rows = df.loc[last_indices] |
|
|
| if len(last_rows) < 2: |
| return df |
|
|
| y_min = df["value"].min() |
| y_max = df["value"].max() |
| y_range = y_max - y_min if y_max != y_min else 1.0 |
| if y_range <= 0: |
| return df |
|
|
| threshold = 0.04 * y_range |
|
|
| last_points = [] |
| for idx, row in last_rows.iterrows(): |
| last_points.append({ |
| "idx": idx, |
| "val": float(row["value"]) |
| }) |
|
|
| last_points.sort(key=lambda x: x["val"]) |
|
|
| |
| for _ in range(10): |
| for i in range(len(last_points) - 1): |
| p1 = last_points[i] |
| p2 = last_points[i+1] |
| diff = p2["val"] - p1["val"] |
| if diff < threshold: |
| overlap = threshold - diff |
| p1["val"] -= overlap / 2.0 |
| p2["val"] += overlap / 2.0 |
|
|
| for p in last_points: |
| df.at[p["idx"], "_label_y"] = p["val"] |
|
|
| except Exception as e: |
| import logging |
| logging.getLogger(__name__).warning(f"Failed to adjust end labels: {e}") |
|
|
| return df |
|
|
|
|
| def build_temporal_single_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Line or grouped-bar chart: 1 indicator, multi-year, ≤8 countries. |
| |
| When result.mark_hint == "bar", renders a grouped bar chart with: |
| - x = year as temporal (same encoding as line chart, timeUnit+format ensures correct display) |
| - xOffset = country (side-by-side bars within each year band) |
| - mark = bar with rounded top corners |
| """ |
| color_dim = result.color_dim |
| original_n = None |
| if color_dim and color_dim in df.columns: |
| df, original_n = _cap_cardinality(df, color_dim, HIGH_CARDINALITY_THRESHOLDS["line_max_series"]) |
|
|
| |
| is_bar = result.mark_hint == "bar" |
| n_series = df[color_dim].nunique() if (color_dim and color_dim in df.columns) else 0 |
| needs_end_labels = ( |
| not is_bar |
| and color_dim |
| and 2 <= n_series <= MAX_END_LABEL_SERIES |
| and "year" in df.columns |
| ) |
| if needs_end_labels: |
| df = _adjust_end_label_y(df, color_dim) |
|
|
| rows = df.to_dict(orient="records") |
| max_abs = float(df["value"].abs().max()) if "value" in df.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
| is_bar = result.mark_hint == "bar" |
| y_title = _resolve_axis_title(y_label, indicator_name) |
| y_ax = { |
| **_axis_style(), |
| "title": y_title, |
| "labelExpr": _value_label_expr(unit_measure), |
| } |
|
|
| |
| |
| |
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
|
|
| encoding: dict = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_ax, |
| "scale": {"zero": is_bar}, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "bar" if is_bar else "line", |
| indicator_labels, |
| value_format=tt_fmt, |
| viz_data=df, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| indicator_name=indicator_name, |
| ), |
| } |
| if result.color_dim: |
| n_items = ( |
| df[result.color_dim].nunique() if result.color_dim in df.columns else 0 |
| ) |
| domain_labels = ( |
| list(df[result.color_dim].unique()) if result.color_dim in df.columns else None |
| ) |
| legend_title = result.dim_name_labels.get(result.color_dim) |
| encoding["color"] = _color_encoding( |
| result.color_dim, |
| mark_type="bar" if is_bar else "line", |
| n_items=n_items, |
| legend_title=legend_title, |
| domain_labels=domain_labels, |
| ) |
| if is_bar and result.color_dim in df.columns: |
| encoding["xOffset"] = {"field": result.color_dim, "type": "nominal"} |
|
|
| |
| annotated_title = _append_breakdown_note(title, df, result.color_dim) |
| annotated_title = _append_trim_note(annotated_title, color_dim, df[color_dim].nunique() if color_dim in df.columns else 0, original_n) |
|
|
| if is_bar: |
| mark_spec: dict = { |
| "type": "bar", |
| "opacity": 0.85, |
| "cornerRadiusTopLeft": 2, |
| "cornerRadiusTopRight": 2, |
| } |
| else: |
| mark_spec = { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| } |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": annotated_title, |
| "data": {"values": rows}, |
| "mark": mark_spec, |
| "encoding": encoding, |
| "width": 600, |
| "height": 350, |
| } |
|
|
| |
| |
| |
| |
| |
| |
| needs_zero_line = ( |
| not is_bar |
| and "value" in df.columns |
| and df["value"].min() < 0 < df["value"].max() |
| ) |
|
|
| |
| |
| |
| |
| |
| color_dim = result.color_dim |
| n_series = df[color_dim].nunique() if (color_dim and color_dim in df.columns) else 0 |
| needs_end_labels = ( |
| not is_bar |
| and color_dim |
| and 2 <= n_series <= MAX_END_LABEL_SERIES |
| and "year" in df.columns |
| ) |
|
|
| if needs_zero_line or needs_end_labels: |
| |
| |
| |
| main_layer: dict = {"mark": spec.pop("mark"), "encoding": spec.pop("encoding")} |
| if needs_end_labels: |
| if "color" in main_layer["encoding"] and isinstance(main_layer["encoding"]["color"], dict): |
| main_layer["encoding"]["color"]["legend"] = None |
| layers: list[dict] = [main_layer] |
|
|
| if needs_zero_line: |
| zero_layer: dict = { |
| "mark": { |
| "type": "rule", |
| "color": "#999999", |
| "strokeWidth": 1.0, |
| "strokeDash": [4, 3], |
| "opacity": 0.8, |
| "tooltip": False, |
| }, |
| "encoding": {"y": {"datum": 0}}, |
| } |
| layers.append(zero_layer) |
|
|
| if needs_end_labels: |
| |
| |
| |
| year_type = "temporal" if df["year"].dtype == "datetime64[ns]" else "ordinal" |
| end_label_layer: dict = { |
| "transform": [ |
| { |
| "aggregate": [ |
| {"op": "argmax", "field": "year", "as": "_last"} |
| ], |
| "groupby": [color_dim], |
| }, |
| { |
| "calculate": "datum._last._label_y", |
| "as": "_end_value", |
| }, |
| { |
| "calculate": f"datum._last.year", |
| "as": "_end_year", |
| }, |
| ], |
| "mark": { |
| "type": "text", |
| "align": "left", |
| "dx": 5, |
| "fontSize": 10, |
| "fontWeight": "normal", |
| "tooltip": False, |
| }, |
| "encoding": { |
| "x": {"field": "_end_year", "type": year_type}, |
| "y": {"field": "_end_value", "type": "quantitative"}, |
| "text": {"field": color_dim, "type": "nominal"}, |
| "color": encoding.get("color", {}), |
| }, |
| } |
| layers.append(end_label_layer) |
|
|
| spec["layer"] = layers |
|
|
| return inject_wb_config(spec) |
|
|
|
|
| def build_cross_sectional_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| x_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Horizontal bar: 1 indicator, single year. |
| |
| Rows are capped at HIGH_CARDINALITY_THRESHOLDS["cross_sectional_max_items"] |
| and sorted descending by value (highest performing country at top). |
| """ |
| bar_dim = result.color_dim or "country" |
| df, original_n = _cap_cardinality( |
| df.sort_values("value", ascending=False), |
| bar_dim, |
| HIGH_CARDINALITY_THRESHOLDS["cross_sectional_max_items"], |
| ) |
| title = _append_trim_note(title, bar_dim, df[bar_dim].nunique() if bar_dim in df.columns else 0, original_n) |
|
|
| rows = df.to_dict(orient="records") |
| max_abs = float(df["value"].abs().max()) if "value" in df.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| color_enc = ( |
| _color_encoding(result.color_dim) |
| if result.color_dim |
| else {"value": WB_CAT_COLORS[0]} |
| ) |
| |
| if isinstance(color_enc, dict) and "field" in color_enc: |
| color_enc["legend"] = None |
| x_title = _resolve_axis_title(x_label, indicator_name) |
| x_ax = { |
| **_axis_style(), |
| "title": x_title, |
| "labelExpr": _value_label_expr(unit_measure), |
| } |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": rows}, |
| "mark": { |
| "type": "bar", |
| "cornerRadiusTopRight": 3, |
| "cornerRadiusBottomRight": 3, |
| }, |
| "encoding": { |
| "y": { |
| "field": "country", |
| "type": "nominal", |
| "sort": "-x", |
| "axis": { |
| "title": None, |
| "labelColor": WB_TEXT, |
| "labelFontWeight": "bold", |
| "labelLimit": 150, |
| }, |
| }, |
| "x": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": x_ax, |
| "scale": {"zero": True}, |
| }, |
| "color": color_enc, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "bar", |
| indicator_labels, |
| value_format=tt_fmt, |
| indicator_name=indicator_name, |
| ), |
| }, |
| "width": 500, |
| "height": max(180, len(df) * 28), |
| } |
| return inject_wb_config(spec) |
|
|
|
|
| def build_distribution_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| x_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Strip/beeswarm: 1 indicator, >8 countries, single year.""" |
| df, original_n = _cap_cardinality( |
| df.sort_values("value", ascending=False), |
| "country", |
| HIGH_CARDINALITY_THRESHOLDS["top_n_series"], |
| ) |
| title = _append_trim_note(title, "country", df["country"].nunique() if "country" in df.columns else 0, original_n) |
|
|
| rows = df.to_dict(orient="records") |
| max_abs = float(df["value"].abs().max()) if "value" in df.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": rows}, |
| "mark": {"type": "bar", "cornerRadiusEnd": 3}, |
| "encoding": { |
| "x": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": { |
| **_axis_style(), |
| "title": x_label, |
| "labelExpr": _value_label_expr(unit_measure), |
| }, |
| }, |
| "y": { |
| "field": "country", |
| "type": "nominal", |
| "sort": "-x", |
| "axis": { |
| "title": None, |
| "labelColor": WB_TEXT, |
| "labelFontWeight": "bold", |
| "labelLimit": 160, |
| }, |
| }, |
| |
| |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "bar", |
| indicator_labels, |
| value_format=tt_fmt, |
| indicator_name=indicator_name, |
| ), |
| }, |
| "width": 500, |
| "height": max(250, len(df) * 22), |
| } |
| return inject_wb_config(spec) |
|
|
|
|
| def build_breakdown_comparison_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Grouped bar: 1 indicator, 1 breakdown (sex/age/urban), 2-4 values, ≤4 countries.""" |
| rows = df.to_dict(orient="records") |
| color_dim = result.color_dim or "sex" |
| max_abs = float(df["value"].abs().max()) if "value" in df.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| if color_dim == "sex": |
| domain = [k for k in WB_GENDER_COLORS if k in df[color_dim].unique()] |
| color_range = [WB_GENDER_COLORS[k] for k in domain] |
| color_scale = {"domain": domain, "range": color_range} |
| else: |
| color_scale = {"range": WB_CAT_COLORS} |
|
|
| x_field = "country" if df.get("country", pd.Series()).nunique() > 1 else "year" |
| if x_field == "year" and df.get("year", pd.Series()).nunique() <= 1: |
| x_field = color_dim |
|
|
| x_enc: dict |
| if x_field == "year": |
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
| else: |
| x_enc = { |
| "field": x_field, |
| "type": "nominal", |
| "axis": {"title": None, "labelFontWeight": "bold"}, |
| } |
| y_ax = { |
| **_axis_style(), |
| "title": None, |
| "labelExpr": _value_label_expr(unit_measure), |
| } |
|
|
| |
| legend_title = _TOOLTIP_SPECS.get(color_dim, {}).get("title") or color_dim.replace("_", " ").title() |
|
|
| n_categories = df[x_field].nunique() if x_field in df.columns else 1 |
| mark_spec: dict = {"type": "bar"} |
| if n_categories == 1: |
| mark_spec["size"] = 30 |
| elif n_categories == 2: |
| mark_spec["size"] = 35 |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": rows}, |
| "mark": mark_spec, |
| "encoding": { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_ax, |
| "scale": {"zero": True}, |
| }, |
| "color": { |
| "field": color_dim, |
| "type": "nominal", |
| "scale": color_scale, |
| "legend": { |
| "orient": "top", |
| "title": legend_title, |
| "labelLimit": 100, |
| "columns": 3, |
| }, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "bar", |
| indicator_labels, |
| value_format=tt_fmt, |
| ), |
| }, |
| "width": max(300, df[x_field].nunique() * 80), |
| "height": 320, |
| } |
| if x_field != color_dim: |
| spec["encoding"]["xOffset"] = {"field": color_dim, "type": "nominal"} |
| return inject_wb_config(spec) |
|
|
|
|
| def _group_breakdowns_by_scale( |
| df: pd.DataFrame, |
| breakdown_dim: str, |
| grouping_threshold: float = 0.75, |
| ) -> list[list[str]]: |
| """Cluster breakdown values into scale-compatible groups. |
| |
| Groups breakdown series so that all members within a group can share a |
| Y-axis without any one series visually dominating the others. Series in |
| different groups will be rendered as separate panels. |
| |
| Algorithm |
| --------- |
| 1. Compute ``log10(max |value|)`` for each breakdown value. |
| 2. Sort breakdown values by this magnitude. |
| 3. Greedily build groups: add the next value to the current group if the |
| group's magnitude span (max_mag − min_mag) stays within |
| *grouping_threshold*. Otherwise start a new group. |
| |
| The default *grouping_threshold* of **0.75** (≈5.6× difference max within |
| a group) is intentionally tighter than the detection threshold of 1.5 |
| (≈30×) used by :func:`_detect_scale_incompatibility`. This ensures that |
| within-group series are visually comparable on a shared Y-axis. |
| |
| This function is **purely data-driven**: it does not use any hardcoded |
| dimension names, indicator codes, or external metadata. It works for any |
| indicator and any number of breakdown values. |
| |
| Args: |
| df: DataFrame containing *breakdown_dim* and ``value`` columns. |
| breakdown_dim: Column containing breakdown series identifiers. |
| grouping_threshold: Maximum log10 span within a group. Default |
| ``0.75`` ≈ 5.6× — half the detection threshold. |
| |
| Returns: |
| Ordered list of groups; each group is an ordered list of breakdown |
| value strings. Every unique non-null breakdown value appears in |
| exactly one group. Falls back to one singleton group per value if |
| computation fails. |
| """ |
| import math |
|
|
| bd_vals = sorted( |
| str(v) for v in df[breakdown_dim].dropna().unique() |
| ) if breakdown_dim in df.columns else [] |
|
|
| if len(bd_vals) <= 1: |
| return [bd_vals] if bd_vals else [] |
|
|
| if "value" not in df.columns: |
| return [[v] for v in bd_vals] |
|
|
| |
| mags: dict[str, float] = {} |
| for v in bd_vals: |
| series = df.loc[df[breakdown_dim] == v, "value"].dropna() |
| if series.empty: |
| mags[v] = 0.0 |
| continue |
| max_abs = float(series.abs().max()) |
| mags[v] = math.log10(max_abs) if max_abs > 0 else 0.0 |
|
|
| |
| sorted_vals = sorted(bd_vals, key=lambda v: mags[v]) |
|
|
| groups: list[list[str]] = [] |
| current_group: list[str] = [] |
| group_min_mag: float = 0.0 |
| group_max_mag: float = 0.0 |
|
|
| for v in sorted_vals: |
| mag = mags[v] |
| if not current_group: |
| current_group = [v] |
| group_min_mag = group_max_mag = mag |
| else: |
| new_min = min(group_min_mag, mag) |
| new_max = max(group_max_mag, mag) |
| if new_max - new_min <= grouping_threshold: |
| current_group.append(v) |
| group_min_mag = new_min |
| group_max_mag = new_max |
| else: |
| groups.append(current_group) |
| current_group = [v] |
| group_min_mag = group_max_mag = mag |
|
|
| if current_group: |
| groups.append(current_group) |
|
|
| return groups |
|
|
|
|
| def _get_label_differentiators(labels: list[str]) -> dict[str, str]: |
| """Given a list of labels, extracts the unique differentiators by stripping common prefix/suffix.""" |
| if not labels: |
| return {} |
| if len(labels) == 1: |
| return {labels[0]: labels[0]} |
|
|
| |
| valid_labels = [l for l in labels if isinstance(l, str) and l.strip()] |
| if len(valid_labels) <= 1: |
| return {l: l for l in labels} |
|
|
| |
| first = valid_labels[0] |
| prefix = "" |
| for i in range(1, len(first) + 1): |
| candidate = first[:i] |
| if all(l.startswith(candidate) for l in valid_labels): |
| prefix = candidate |
| else: |
| break |
|
|
| |
| if prefix: |
| ends_at_boundary = all( |
| l[len(prefix):].startswith((" ", ",", "-", "(", ")", "/", "[", "]", "{", "}")) or |
| prefix.endswith((" ", ",", "-", "(", ")", "/", "[", "]", "{", "}")) |
| for l in valid_labels |
| ) |
| if not ends_at_boundary: |
| while prefix and not prefix[-1].isspace() and prefix[-1] not in (",", "-", "(", ")", "/", "[", "]", "{", "}"): |
| prefix = prefix[:-1] |
|
|
| |
| reversed_first = first[::-1] |
| suffix = "" |
| for i in range(1, len(reversed_first) + 1): |
| candidate = reversed_first[:i][::-1] |
| if all(l.endswith(candidate) for l in valid_labels): |
| suffix = candidate |
| else: |
| break |
|
|
| |
| if suffix: |
| starts_at_boundary = all( |
| l[:-len(suffix)].endswith((" ", ",", "-", "(", ")", "/", "[", "]", "{", "}")) or |
| suffix.startswith((" ", ",", "-", "(", ")", "/", "[", "]", "{", "}")) |
| for l in valid_labels |
| ) |
| if not starts_at_boundary: |
| while suffix and not suffix[0].isspace() and suffix[0] not in (",", "-", "(", ")", "/", "[", "]", "{", "}"): |
| suffix = suffix[1:] |
|
|
| |
| mapping = {} |
| for l in labels: |
| if not isinstance(l, str): |
| mapping[l] = l |
| continue |
|
|
| shortened = l |
| if prefix: |
| shortened = shortened[len(prefix):] |
| if suffix: |
| shortened = shortened[:-len(suffix)] |
|
|
| |
| shortened = shortened.strip(",;.:-()[]{} ") |
|
|
| |
| if len(shortened) < 2: |
| mapping[l] = l |
| else: |
| |
| if shortened[0].islower(): |
| shortened = shortened[0].upper() + shortened[1:] |
| mapping[l] = shortened |
|
|
| return mapping |
|
|
|
|
| def _truncate_panel_title(title_text: str | list[str], max_len: int = 35) -> str | list[str]: |
| if isinstance(title_text, list): |
| return [_truncate_panel_title(t, max_len) for t in title_text] |
| if not isinstance(title_text, str): |
| return title_text |
| if len(title_text) > max_len: |
| return title_text[:max_len-1].strip() + "\u2026" |
| return title_text |
|
|
|
|
| def _determine_small_multiples_columns(n_panels: int, year_count: int) -> int | None: |
| if n_panels <= 1: |
| return None |
| |
| if year_count > 10: |
| return 1 if n_panels <= 3 else 2 |
| return 2 |
|
|
|
|
| def _build_scale_split_vconcat( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| grouping_threshold: float = 0.75, |
| force_separate_panels: bool = False, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Vconcat layout for scale-incompatible custom breakdowns. |
| |
| Produces one full-width panel per **scale-compatible group** of breakdown |
| values. Groups are discovered automatically by |
| :func:`_group_breakdowns_by_scale` (data-driven, no hard-coded metadata). |
| |
| Panel layout |
| ------------ |
| - **Single-member group** (680×140): one fixed-color line, panel title = |
| series label. Identical to the pre-grouping behaviour. |
| - **Multi-member group** (680×180): layered multi-series lines with Vega-Lite |
| ``color`` encoding, a right-side legend, and a panel title listing the |
| member labels (truncated at 80 characters). |
| |
| All panels share the X (year) axis via ``resolve.scale.x = 'shared'``. |
| X-axis labels are suppressed on all but the bottom panel. |
| |
| The default *grouping_threshold* of 0.75 (≈5.6×) is tighter than the |
| detection threshold of 1.5 (≈30×), so within-group series are always |
| visually comparable on a shared Y-axis. |
| """ |
| facet_dim = result.facet_dim or "comp_breakdown_1" |
| lab = indicator_labels or {} |
|
|
| |
| df, original_n = _cap_cardinality( |
| df, facet_dim, HIGH_CARDINALITY_THRESHOLDS["small_multiples_max_facets"] |
| ) |
|
|
| breakdown_vals = sorted(df[facet_dim].dropna().unique(), key=str) |
| |
| if facet_dim != "country": |
| raw_bd_labels = [lab.get(v, v) for v in breakdown_vals] |
| diff_map = _get_label_differentiators(raw_bd_labels) |
| for v in breakdown_vals: |
| orig = lab.get(v, v) |
| if orig in diff_map: |
| lab[v] = diff_map[orig] |
|
|
| rows = df.to_dict(orient="records") |
| label_expr = _value_label_expr(unit_measure) |
|
|
| |
| |
| if original_n is not None and facet_dim == "country" and isinstance(title, dict): |
| existing_sub = title.get("subtitle", "") |
| if isinstance(existing_sub, str): |
| existing_sub = [p.strip() for p in existing_sub.split(" · ") if p.strip()] |
|
|
| shown_countries = sorted(df["country"].unique().tolist(), key=str.casefold) |
| |
| from data360.viz_config import _year_range_label |
| year_lbl = _year_range_label(df["year"]) if "year" in df.columns else None |
|
|
| new_geo = ", ".join(shown_countries) |
| if year_lbl: |
| new_geo = f"{new_geo}, {year_lbl}" |
| title = {**title, "subtitle": [new_geo] + existing_sub[1:]} |
|
|
| annotated_title = _append_trim_note( |
| title, facet_dim, |
| len(breakdown_vals), |
| original_n, |
| ) |
|
|
| |
| if force_separate_panels: |
| |
| groups = [[v] for v in breakdown_vals] |
| else: |
| |
| is_pct_unit = False |
| if unit_measure and any(x in str(unit_measure).lower() for x in ["percent", "pct", "%"]): |
| is_pct_unit = True |
| elif y_label and any(x in str(y_label).lower() for x in ["percent", "pct", "%"]): |
| is_pct_unit = True |
|
|
| val_series = pd.to_numeric(df["value"], errors="coerce").dropna() |
| is_pct = is_pct_unit and not val_series.empty and val_series.max() <= 100.0 and val_series.min() >= 0.0 |
|
|
| |
| same_numeric_scale = True |
| if is_pct: |
| maxes = [] |
| for v in breakdown_vals: |
| s = pd.to_numeric(df.loc[df[facet_dim] == v, "value"], errors="coerce").dropna() |
| if not s.empty: |
| maxes.append(s.max()) |
| if maxes: |
| any_gt_1 = any(m > 1.0 for m in maxes) |
| all_lte_1 = all(m <= 1.0 for m in maxes) |
| same_numeric_scale = any_gt_1 or all_lte_1 |
|
|
| actual_threshold = float("inf") if (is_pct and same_numeric_scale) else grouping_threshold |
| groups = _group_breakdowns_by_scale(df, facet_dim, actual_threshold) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| color_resolve = ( |
| "shared" |
| if result.color_dim and result.color_dim != facet_dim |
| else "independent" |
| ) |
|
|
| |
| |
| |
| |
| _color_dim_domain: list[str] | None = ( |
| sorted(df[result.color_dim].dropna().unique().tolist()) |
| if result.color_dim and result.color_dim in df.columns |
| else None |
| ) |
|
|
| |
| |
| |
| _n_panels = len(groups) |
| if _n_panels == 1: |
| group = groups[0] |
| is_bar_chart = (df["year"].nunique() <= 1 if "year" in df.columns else True) |
|
|
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
| if is_bar_chart: |
| x_enc["type"] = "nominal" |
| if "timeUnit" in x_enc: |
| del x_enc["timeUnit"] |
| if "axis" in x_enc and "format" in x_enc["axis"]: |
| del x_enc["axis"]["format"] |
|
|
| if len(group) == 1: |
| if result.color_dim and result.color_dim in df.columns: |
| x_enc["field"] = result.color_dim |
| else: |
| x_enc["field"] = facet_dim |
|
|
| if "axis" not in x_enc or x_enc["axis"] is None: |
| x_enc["axis"] = {} |
| x_enc["axis"].update({ |
| "labelAngle": -45, |
| "labelAlign": "right", |
| "labelBaseline": "middle", |
| "labelLimit": 150, |
| "labelOverlap": False, |
| }) |
|
|
| y_title = _resolve_axis_title(y_label, indicator_name) |
| y_axis = { |
| **_axis_style(), |
| "title": y_title, |
| "labelExpr": label_expr, |
| } |
|
|
| df_filtered = df[df[facet_dim].isin(group)].copy() |
|
|
| if len(group) == 1: |
| bd_val = group[0] |
| bd_label = lab.get(bd_val, bd_val) |
| bd_data = df_filtered[df_filtered[facet_dim] == bd_val] |
| max_abs = ( |
| float(bd_data["value"].abs().max()) |
| if "value" in bd_data.columns and not bd_data.empty |
| else None |
| ) |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| if is_bar_chart: |
| mark_spec = { |
| "type": "bar", |
| "cornerRadiusTopRight": 2, |
| "cornerRadiusTopLeft": 2, |
| "size": 40, |
| } |
| else: |
| mark_spec = { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| } |
|
|
| encoding = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": is_bar_chart}, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(bd_data.columns), |
| "bar" if is_bar_chart else "line", |
| indicator_labels={**lab, "value": bd_label}, |
| value_format=tt_fmt, |
| viz_data=bd_data, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
|
|
| extra_transforms = [] |
| if result.color_dim and result.secondary_color_dim: |
| _sec = result.secondary_color_dim |
| _pri = result.color_dim |
| sorted_secondary = sorted(df[_sec].dropna().unique().tolist()) if _sec in df.columns else [] |
| sorted_primary = sorted(bd_data[_pri].dropna().unique().tolist()) if _pri in bd_data.columns else [] |
| _s_combo_domain = [ |
| f"{s} | {lab.get(p, p)}" |
| for s in sorted_secondary |
| for p in sorted_primary |
| ] |
| _s_combo_range = [] |
| for si, _country in enumerate(sorted_secondary): |
| base_color = WB_CAT_COLORS[si % len(WB_CAT_COLORS)] |
| _s_combo_range.extend(_generate_color_shades(base_color, len(sorted_primary))) |
| _s_combo_field = "_s_combo_label" |
| _s_combo_calc = { |
| "calculate": f"datum['{_sec}'] + ' | ' + datum['{_pri}']", |
| "as": _s_combo_field, |
| } |
| extra_transforms.append(_s_combo_calc) |
| _pri_title = result.dim_name_labels.get(_pri) or _pri.replace("_", " ").title() |
| _sec_title = result.dim_name_labels.get(_sec) or _sec.replace("_", " ").title() |
|
|
| _all_combo_labels = _s_combo_domain |
| _legend_layout = _compute_legend_layout(_all_combo_labels) |
| encoding["color"] = { |
| "field": _s_combo_field, |
| "type": "nominal", |
| "scale": {"domain": _s_combo_domain, "range": _s_combo_range}, |
| "legend": {**_legend_layout, "title": f"{_sec_title} | {_pri_title}"}, |
| } |
| elif result.color_dim: |
| n_items = bd_data[result.color_dim].nunique() if result.color_dim in bd_data.columns else 0 |
| domain_labels = list(bd_data[result.color_dim].unique()) if result.color_dim in bd_data.columns else None |
| legend_title = result.dim_name_labels.get(result.color_dim) |
| encoding["color"] = _color_encoding( |
| result.color_dim, |
| mark_type="bar" if is_bar_chart else "line", |
| n_items=n_items, |
| legend_title=legend_title, |
| domain_labels=domain_labels, |
| domain=_color_dim_domain, |
| ) |
| else: |
| mark_spec["color"] = WB_CAT_COLORS[0] |
|
|
| else: |
| group_labels = [lab.get(v, v) for v in group] |
| max_abs = ( |
| float(df_filtered["value"].abs().max()) |
| if "value" in df_filtered.columns and not df_filtered.empty |
| else None |
| ) |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| mark_spec = { |
| "type": "bar" if is_bar_chart else "line", |
| **( |
| {"cornerRadiusTopRight": 2, "cornerRadiusTopLeft": 2, "size": 40} |
| if is_bar_chart |
| else {"strokeWidth": 3, "strokeCap": "round", "point": _LINE_HOVER_POINT} |
| ) |
| } |
|
|
| extra_transforms = [] |
| if is_bar_chart and result.color_dim: |
| n_items = df_filtered[result.color_dim].nunique() if result.color_dim in df_filtered.columns else 0 |
| domain_labels = list(df_filtered[result.color_dim].unique()) if result.color_dim in df_filtered.columns else None |
| legend_title = result.dim_name_labels.get(result.color_dim) |
| encoding = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": is_bar_chart}, |
| }, |
| "color": _color_encoding( |
| result.color_dim, |
| mark_type="bar", |
| n_items=n_items, |
| legend_title=legend_title, |
| domain_labels=domain_labels, |
| domain=_color_dim_domain, |
| ), |
| "tooltip": build_structured_tooltips( |
| list(df_filtered.columns), |
| "bar", |
| indicator_labels=lab, |
| value_format=tt_fmt, |
| viz_data=df_filtered, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
| elif result.color_dim and result.color_dim != facet_dim: |
| sorted_countries = sorted(df[result.color_dim].dropna().unique().tolist()) |
| combo_domain = [ |
| f"{c} | {lab.get(bd, bd)}" |
| for c in sorted_countries |
| for bd in group |
| ] |
| combo_range = [] |
| for ci, country in enumerate(sorted_countries): |
| base_color = WB_CAT_COLORS[ci % len(WB_CAT_COLORS)] |
| combo_range.extend(_generate_color_shades(base_color, len(group))) |
|
|
| _combo_field = "_combo_label" |
| _combo_calc = { |
| "calculate": f"datum['{result.color_dim}'] + ' | ' + datum['{facet_dim}']", |
| "as": _combo_field, |
| } |
| extra_transforms.append(_combo_calc) |
|
|
| _country_title = result.dim_name_labels.get(result.color_dim) or result.color_dim.title() |
| _bd_title = result.dim_name_labels.get(facet_dim) or facet_dim.title() |
|
|
| _all_combo_labels = combo_domain |
| _legend_layout = _compute_legend_layout(_all_combo_labels) |
| encoding = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": is_bar_chart}, |
| }, |
| "color": { |
| "field": _combo_field, |
| "type": "nominal", |
| "scale": {"domain": combo_domain, "range": combo_range}, |
| "legend": {**_legend_layout, "title": f"{_country_title} | {_bd_title}"}, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(df_filtered.columns), |
| "bar" if is_bar_chart else "line", |
| indicator_labels=lab, |
| value_format=tt_fmt, |
| viz_data=df_filtered, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
| else: |
| group_colors = [ |
| WB_CAT_COLORS[j % len(WB_CAT_COLORS)] |
| for j in range(len(group)) |
| ] |
| encoding = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": is_bar_chart}, |
| }, |
| "color": { |
| "field": facet_dim, |
| "type": "nominal", |
| "scale": { |
| "domain": group, |
| "range": group_colors, |
| }, |
| "legend": { |
| "orient": "bottom" if max((len(lbl) for lbl in group_labels), default=0) > 40 else "right", |
| "labelFontSize": 11, |
| "symbolSize": 80, |
| **( |
| {"labelLimit": 0, "direction": "vertical"} |
| if max((len(lbl) for lbl in group_labels), default=0) > 40 |
| else {"labelLimit": 200} |
| ), |
| }, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(df_filtered.columns), |
| "bar" if is_bar_chart else "line", |
| indicator_labels=lab, |
| value_format=tt_fmt, |
| viz_data=df_filtered, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
|
|
| effective_color_dim = facet_dim if len(group) > 1 else result.color_dim |
| n_series = df_filtered[effective_color_dim].nunique() if (effective_color_dim and effective_color_dim in df_filtered.columns) else 0 |
| needs_end_labels = ( |
| not is_bar_chart |
| and effective_color_dim |
| and 2 <= n_series <= MAX_END_LABEL_SERIES |
| and "year" in df_filtered.columns |
| ) |
| if needs_end_labels: |
| df_filtered = _adjust_end_label_y(df_filtered, effective_color_dim) |
|
|
| rows = df_filtered.to_dict(orient="records") |
|
|
| spec = { |
| "$schema": _vl_schema(), |
| "title": annotated_title, |
| "data": {"values": rows}, |
| "mark": mark_spec, |
| "encoding": encoding, |
| "width": 600, |
| "height": 350, |
| } |
|
|
| if extra_transforms: |
| spec["transform"] = extra_transforms |
|
|
| needs_zero_line = ( |
| not is_bar_chart |
| and "value" in df_filtered.columns |
| and df_filtered["value"].min() < 0 < df_filtered["value"].max() |
| ) |
|
|
| |
| if is_bar_chart and result.color_dim and result.color_dim in df_filtered.columns: |
| if encoding.get("x", {}).get("field") != result.color_dim: |
| encoding["xOffset"] = {"field": result.color_dim, "type": "nominal"} |
|
|
| if needs_zero_line or needs_end_labels: |
| main_layer = {"mark": spec.pop("mark"), "encoding": spec.pop("encoding")} |
| if needs_end_labels: |
| if "color" in main_layer["encoding"] and isinstance(main_layer["encoding"]["color"], dict): |
| main_layer["encoding"]["color"]["legend"] = None |
| layers = [main_layer] |
|
|
| if needs_zero_line: |
| zero_layer = { |
| "mark": { |
| "type": "rule", |
| "color": "#999999", |
| "strokeWidth": 1.0, |
| "strokeDash": [4, 3], |
| "opacity": 0.8, |
| "tooltip": False, |
| }, |
| "encoding": {"y": {"datum": 0}}, |
| } |
| layers.append(zero_layer) |
|
|
| if needs_end_labels: |
| year_type = "temporal" if df_filtered["year"].dtype == "datetime64[ns]" else "ordinal" |
| end_label_layer = { |
| "transform": [ |
| { |
| "aggregate": [ |
| {"op": "argmax", "field": "year", "as": "_last"} |
| ], |
| "groupby": [effective_color_dim], |
| }, |
| { |
| "calculate": "datum._last._label_y", |
| "as": "_end_value", |
| }, |
| { |
| "calculate": "datum._last.year", |
| "as": "_end_year", |
| }, |
| ], |
| "mark": { |
| "type": "text", |
| "align": "left", |
| "dx": 5, |
| "fontSize": 10, |
| "fontWeight": "normal", |
| "tooltip": False, |
| }, |
| "encoding": { |
| "x": {"field": "_end_year", "type": year_type}, |
| "y": {"field": "_end_value", "type": "quantitative"}, |
| "text": {"field": effective_color_dim, "type": "nominal"}, |
| "color": encoding.get("color", {}), |
| }, |
| } |
| layers.append(end_label_layer) |
|
|
| spec["layer"] = layers |
|
|
| return inject_wb_config(spec) |
| _legend_target_total_px = 850 |
| _base_single_px = 240 |
| _base_multi_px = 300 |
|
|
| if result.color_dim and (result.color_dim != facet_dim or result.secondary_color_dim): |
| |
| _sec_dim = result.secondary_color_dim |
| _pri_dim = result.color_dim |
| if _sec_dim: |
| |
| _sorted_sec = sorted(df[_sec_dim].dropna().unique().tolist()) |
| _sorted_pri = sorted(df[_pri_dim].dropna().unique().tolist()) |
| _all_combo_labels = [ |
| f"{s} | {lab.get(p, p)}" |
| for s in _sorted_sec |
| for p in _sorted_pri |
| ] |
| else: |
| |
| _sorted_countries = sorted(df[_pri_dim].dropna().unique().tolist()) |
| _max_group = max(groups, key=len) if groups else [] |
| _all_combo_labels = [ |
| f"{c} | {lab.get(bd, bd)}" |
| for c in _sorted_countries |
| for bd in _max_group |
| ] |
| _legend_layout = _compute_legend_layout(_all_combo_labels) |
| _legend_h = _estimate_legend_height(len(_all_combo_labels), _legend_layout, has_title=True) |
| _panel_h_single = max(80, (_legend_target_total_px - _legend_h) // _n_panels) |
| _panel_h_multi = max(100, (_legend_target_total_px - _legend_h) // _n_panels) |
| else: |
| |
| _legend_layout = _compute_legend_layout([]) |
| _legend_h = 0 |
| _panel_h_single = _base_single_px |
| _panel_h_multi = _base_multi_px |
|
|
| is_multi_panel = len(groups) > 1 |
| panel_width = 280 if is_multi_panel else 680 |
| if is_multi_panel: |
| _panel_h_single = 200 |
| _panel_h_multi = 220 |
|
|
| is_bar_chart = (df["year"].nunique() <= 1 if "year" in df.columns else True) |
|
|
| charts: list[dict] = [] |
| color_offset = 0 |
|
|
| for g_idx, group in enumerate(groups): |
| is_last_panel = g_idx == len(groups) - 1 |
|
|
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
| if is_bar_chart: |
| x_enc["type"] = "nominal" |
| if "timeUnit" in x_enc: |
| del x_enc["timeUnit"] |
| if "axis" in x_enc and "format" in x_enc["axis"]: |
| del x_enc["axis"]["format"] |
|
|
| |
| if len(group) == 1: |
| if result.color_dim and result.color_dim in df.columns: |
| x_enc["field"] = result.color_dim |
| else: |
| x_enc["field"] = facet_dim |
|
|
| if "axis" not in x_enc or x_enc["axis"] is None: |
| x_enc["axis"] = {} |
| x_enc["axis"].update({ |
| "labelAngle": -45, |
| "labelAlign": "right", |
| "labelBaseline": "middle", |
| "labelLimit": 150, |
| "labelOverlap": False, |
| }) |
|
|
| if not is_last_panel and not is_multi_panel: |
| x_enc = {**x_enc, "axis": {**x_enc.get("axis", {}), "labels": False, "title": None}} |
|
|
| y_axis = {**_axis_style(), "title": None, "labelExpr": label_expr} |
|
|
| if len(group) == 1: |
| |
| |
| |
| bd_val = group[0] |
| color = WB_CAT_COLORS[color_offset % len(WB_CAT_COLORS)] |
| color_offset += 1 |
| bd_label = lab.get(bd_val, bd_val) |
|
|
| bd_data = df[df[facet_dim] == bd_val] |
| max_abs = ( |
| float(bd_data["value"].abs().max()) |
| if "value" in bd_data.columns and not bd_data.empty |
| else None |
| ) |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| if is_bar_chart: |
| mark_spec = { |
| "type": "bar", |
| "cornerRadiusTopRight": 2, |
| "cornerRadiusTopLeft": 2, |
| "size": 40, |
| } |
| else: |
| mark_spec = { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| } |
|
|
| chart_enc = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": False}, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(bd_data.columns), |
| "line", |
| indicator_labels={**lab, "value": bd_label}, |
| value_format=tt_fmt, |
| viz_data=bd_data, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
|
|
| if result.color_dim and result.secondary_color_dim: |
| |
| |
| |
| |
| _sec = result.secondary_color_dim |
| _pri = result.color_dim |
| sorted_secondary = sorted( |
| df[_sec].dropna().unique().tolist() |
| ) if _sec in df.columns else [] |
| sorted_primary = sorted( |
| bd_data[_pri].dropna().unique().tolist() |
| ) if _pri in bd_data.columns else [] |
| |
| _s_combo_domain: list[str] = [ |
| f"{s} | {lab.get(p, p)}" |
| for s in sorted_secondary |
| for p in sorted_primary |
| ] |
| |
| _s_combo_range: list[str] = [] |
| for si, _country in enumerate(sorted_secondary): |
| base_color = WB_CAT_COLORS[si % len(WB_CAT_COLORS)] |
| _s_combo_range.extend(_generate_color_shades(base_color, len(sorted_primary))) |
| _s_combo_field = "_s_combo_label" |
| _s_combo_calc = { |
| "calculate": f"datum['{_sec}'] + ' | ' + datum['{_pri}']", |
| "as": _s_combo_field, |
| } |
| _pri_title = result.dim_name_labels.get(_pri) or _pri.replace("_", " ").title() |
| _sec_title = result.dim_name_labels.get(_sec) or _sec.replace("_", " ").title() |
| |
| _s_legend: dict | None = ( |
| {**_legend_layout, "title": f"{_sec_title} | {_pri_title}"} |
| if is_last_panel |
| else None |
| ) |
| chart_enc["color"] = { |
| "field": _s_combo_field, |
| "type": "nominal", |
| "scale": {"domain": _s_combo_domain, "range": _s_combo_range}, |
| "legend": _s_legend, |
| } |
| charts.append({ |
| "title": { |
| "text": _truncate_panel_title(bd_label), |
| "fontSize": 12, |
| "fontWeight": "bold", |
| "anchor": "start", |
| "offset": 4, |
| }, |
| "width": panel_width, |
| "height": _panel_h_single, |
| "transform": [ |
| _s_combo_calc, |
| {"filter": {"field": facet_dim, "equal": bd_val}}, |
| ], |
| "mark": mark_spec, |
| "encoding": chart_enc, |
| }) |
| continue |
|
|
| elif result.color_dim: |
| n_items = ( |
| bd_data[result.color_dim].nunique() |
| if result.color_dim in bd_data.columns |
| else 0 |
| ) |
| domain_labels_for_group = ( |
| list(bd_data[result.color_dim].unique()) |
| if result.color_dim and result.color_dim in bd_data.columns |
| else None |
| ) |
| legend_title_for_group = result.dim_name_labels.get(result.color_dim) if result.color_dim else None |
| chart_enc["color"] = _color_encoding( |
| result.color_dim, |
| mark_type="line", |
| n_items=n_items, |
| legend_title=legend_title_for_group, |
| domain_labels=domain_labels_for_group, |
| domain=_color_dim_domain, |
| ) |
| |
| |
| |
| |
| |
| else: |
| mark_spec["color"] = color |
|
|
| charts.append({ |
| "title": { |
| "text": _truncate_panel_title(bd_label), |
| "color": color if not result.color_dim else None, |
| "fontSize": 12, |
| "fontWeight": "bold", |
| "anchor": "start", |
| "offset": 4, |
| }, |
| "width": panel_width, |
| "height": _panel_h_single, |
| "transform": [{"filter": {"field": facet_dim, "equal": bd_val}}], |
| "mark": mark_spec, |
| "encoding": chart_enc, |
| }) |
|
|
| else: |
| |
| |
| |
| group_labels = [lab.get(v, v) for v in group] |
| import textwrap as _textwrap |
| joined = ", ".join(group_labels) |
| wrapped = _textwrap.wrap(joined, width=100) |
| if len(wrapped) > 2: |
| wrapped = wrapped[:2] |
| wrapped[-1] = wrapped[-1].rstrip(",") + "\u2026" |
| panel_title: str | list[str] = wrapped if len(wrapped) > 1 else (wrapped[0] if wrapped else joined) |
|
|
| group_data = df[df[facet_dim].isin(group)] |
| max_abs = ( |
| float(group_data["value"].abs().max()) |
| if "value" in group_data.columns and not group_data.empty |
| else None |
| ) |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| _panel_extra_transforms: list[dict] = [] |
| if result.color_dim and result.color_dim != facet_dim: |
| |
| |
| |
| sorted_countries = sorted( |
| df[result.color_dim].dropna().unique().tolist() |
| ) |
| |
| combo_domain: list[str] = [ |
| f"{c} | {lab.get(bd, bd)}" |
| for c in sorted_countries |
| for bd in group |
| ] |
| |
| combo_range: list[str] = [] |
| for ci, country in enumerate(sorted_countries): |
| base_color = WB_CAT_COLORS[ci % len(WB_CAT_COLORS)] |
| combo_range.extend(_generate_color_shades(base_color, len(group))) |
|
|
| |
| _combo_field = "_combo_label" |
| _combo_calc = { |
| "calculate": ( |
| f"datum['{result.color_dim}'] + ' | ' + datum['{facet_dim}']" |
| ), |
| "as": _combo_field, |
| } |
|
|
| _country_title = ( |
| result.dim_name_labels.get(result.color_dim, result.color_dim.title()) |
| ) |
| _bd_title = ( |
| result.dim_name_labels.get(facet_dim, facet_dim.title()) |
| ) |
| |
| combo_legend: dict | None = ( |
| {**_legend_layout, "title": f"{_country_title} | {_bd_title}"} |
| if is_last_panel |
| else None |
| ) |
|
|
| panel_encoding: dict = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": False}, |
| }, |
| "color": { |
| "field": _combo_field, |
| "type": "nominal", |
| "scale": {"domain": combo_domain, "range": combo_range}, |
| "legend": combo_legend, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(group_data.columns), |
| "line", |
| indicator_labels=lab, |
| value_format=tt_fmt, |
| viz_data=group_data, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
| |
| _panel_extra_transforms = [_combo_calc] |
| else: |
| |
| group_colors = [ |
| WB_CAT_COLORS[(color_offset + j) % len(WB_CAT_COLORS)] |
| for j in range(len(group)) |
| ] |
| color_offset += len(group) |
| panel_encoding = { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": False}, |
| }, |
| "color": { |
| "field": facet_dim, |
| "type": "nominal", |
| "scale": { |
| "domain": group, |
| "range": group_colors, |
| }, |
| "legend": { |
| "orient": "bottom" if max((len(lbl) for lbl in group_labels), default=0) > 40 else "right", |
| "labelFontSize": 11, |
| "symbolSize": 80, |
| **( |
| {"labelLimit": 0, "direction": "vertical"} |
| if max((len(lbl) for lbl in group_labels), default=0) > 40 |
| else {"labelLimit": 200} |
| ), |
| }, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(group_data.columns), |
| "bar" if is_bar_chart else "line", |
| indicator_labels=lab, |
| value_format=tt_fmt, |
| viz_data=group_data, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
|
|
| charts.append({ |
| "title": { |
| "text": _truncate_panel_title(panel_title), |
| "fontSize": 12, |
| "fontWeight": "bold", |
| "anchor": "start", |
| "offset": 4, |
| }, |
| "width": panel_width, |
| "height": _panel_h_multi, |
| |
| |
| "transform": _panel_extra_transforms + [{"filter": {"field": facet_dim, "oneOf": group}}], |
| "mark": { |
| "type": "bar" if is_bar_chart else "line", |
| **( |
| {"cornerRadiusTopRight": 2, "cornerRadiusTopLeft": 2, "size": 40} |
| if is_bar_chart |
| else {"strokeWidth": 3, "strokeCap": "round", "point": _LINE_HOVER_POINT} |
| ) |
| }, |
| "encoding": panel_encoding, |
| }) |
|
|
| x_resolve = "independent" if is_bar_chart else "shared" |
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": annotated_title, |
| "data": {"values": rows}, |
| "resolve": {"scale": {"x": x_resolve, "color": color_resolve}}, |
| "autosize": {"type": "fit", "contains": "padding"}, |
| } |
| if is_bar_chart: |
| spec["padding"] = {"top": 10, "left": 10, "bottom": 110, "right": 10} |
| |
| |
| spec["concat"] = charts |
| if is_multi_panel: |
| year_cnt = df["year"].nunique() if "year" in df.columns else 1 |
| spec["columns"] = _determine_small_multiples_columns(len(groups), year_cnt) |
|
|
| |
| color_dim = result.color_dim |
| if color_dim and color_dim in df.columns: |
| n_items = df[color_dim].nunique() |
| if n_items > 4: |
| cols = min(5, (n_items + 1) // 2) |
| spec["config"] = { |
| "legend": { |
| "orient": "bottom", |
| "direction": "horizontal", |
| "columns": cols, |
| "title": result.dim_name_labels.get(color_dim, color_dim.replace("_", " ").title()) if result.dim_name_labels else color_dim.replace("_", " ").title(), |
| "labelLimit": 200, |
| } |
| } |
| else: |
| spec["columns"] = 1 |
| return inject_wb_config(spec) |
|
|
|
|
| def build_small_multiples_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Faceted small multiples: 1 indicator, 2+ breakdowns or breakdown+many countries. |
| |
| Facet panels are capped at HIGH_CARDINALITY_THRESHOLDS["small_multiples_max_facets"] |
| via the shared :func:`_cap_cardinality` utility. When trimmed, the top-N facet |
| values by most-recent data point are retained and a subtitle note is injected via |
| :func:`_append_trim_note` — the same logic used by build_cross_sectional_spec. |
| |
| When ``result.scale_incompatible`` is True, delegates to |
| :func:`_build_scale_split_vconcat` which produces one full-width vertically |
| stacked panel per breakdown value, each with an independent Y-axis. This |
| prevents scale-dominant series from compressing smaller ones on a shared axis. |
| """ |
| facet_dim = result.facet_dim or "comp_breakdown_1" |
| color_dim = result.color_dim |
|
|
| |
| indicator_cols = result.indicator_cols or [] |
| unique_inds = df["indicator"].dropna().unique().tolist() if "indicator" in df.columns else [] |
| if not indicator_cols and len(unique_inds) == 2: |
| indicator_cols = unique_inds |
|
|
| country_cnt = df["country"].nunique() if "country" in df.columns else 0 |
| is_bar = (result.mark_hint == "bar") or (df["year"].nunique() <= 1 if "year" in df.columns else True) |
|
|
| if len(indicator_cols) == 2 and country_cnt <= 1 and not is_bar: |
| try: |
| clean_labels = {} |
| if indicator_labels: |
| clean_labels = { |
| _clean_label_generic(k): _clean_label_generic(v) |
| for k, v in indicator_labels.items() |
| } |
|
|
| maxes = {} |
| for ind in indicator_cols: |
| clean_ind = _clean_label_generic(ind) |
| pretty_val = clean_labels.get(clean_ind, clean_ind) |
|
|
| if ind in df.columns: |
| vals = df[ind].dropna().abs() |
| elif pretty_val in df.columns: |
| vals = df[pretty_val].dropna().abs() |
| elif "indicator" in df.columns: |
| vals = df[df["indicator"].map(lambda x: _clean_label_generic(x) if isinstance(x, str) else x).isin([clean_ind, pretty_val])]["value"].dropna().abs() |
| else: |
| vals = pd.Series(dtype=float) |
| maxes[ind] = vals.max() if not vals.empty else 0.0 |
|
|
| nz_maxes = [m for m in maxes.values() if m > 0] |
| is_incompatible = False |
| if len(nz_maxes) == 2: |
| ratio = max(nz_maxes) / min(nz_maxes) |
| if ratio > 10.0: |
| is_incompatible = True |
| except Exception: |
| is_incompatible = False |
|
|
| if is_incompatible: |
| if "indicator" in df.columns and "value" in df.columns: |
| index_cols = [c for c in df.columns if c not in ("indicator", "value")] |
| wide_df = df.pivot_table(index=index_cols, columns="indicator", values="value", aggfunc="mean").reset_index() |
| wide_df.columns.name = None |
| temp_indicator_cols = [c for c in wide_df.columns if c not in index_cols] |
| else: |
| wide_df = df |
| temp_indicator_cols = indicator_cols |
|
|
| from data360.viz_config import build_temporal_multi_indicator_spec, StrategyResult, ChartStrategy |
| temp_result = StrategyResult( |
| strategy=ChartStrategy.TEMPORAL_MULTI_IND, |
| reason=result.reason, |
| indicator_cols=temp_indicator_cols, |
| color_dim=result.color_dim, |
| facet_dim=result.facet_dim, |
| mark_hint=result.mark_hint, |
| scale_incompatible=result.scale_incompatible, |
| ) |
| return build_temporal_multi_indicator_spec( |
| wide_df, |
| title, |
| temp_result, |
| indicator_labels=indicator_labels, |
| y_label=y_label, |
| unit_measure=unit_measure, |
| indicator_name=indicator_name, |
| ) |
|
|
| |
| if (facet_dim == "indicator" or color_dim == "indicator") and "indicator" not in df.columns: |
| if result.indicator_cols: |
| df = df.copy() |
| id_cols = [c for c in df.columns if c not in result.indicator_cols] |
| df = df.melt( |
| id_vars=id_cols, |
| value_vars=result.indicator_cols, |
| var_name="indicator", |
| value_name="value", |
| ) |
| |
| if indicator_labels: |
| df["indicator"] = df["indicator"].map(lambda x: indicator_labels.get(x, x)) |
| df = df.dropna(subset=["value"]) |
|
|
| if facet_dim in df.columns: |
| df = df.copy() |
| df[facet_dim] = df[facet_dim].map(lambda x: _clean_label_generic(x) if isinstance(x, str) else x) |
|
|
| if indicator_labels: |
| indicator_labels = { |
| _clean_label_generic(k): _clean_label_generic(v) |
| for k, v in indicator_labels.items() |
| } |
|
|
| force_sep = result.scale_incompatible |
| |
| if facet_dim == "indicator" and color_dim != "indicator" and not is_bar: |
| force_sep = True |
| |
| elif facet_dim == "country": |
| force_sep = True |
| return _build_scale_split_vconcat( |
| df, title, result, indicator_labels, y_label, unit_measure, |
| force_separate_panels=force_sep, |
| indicator_name=indicator_name |
| ) |
|
|
|
|
|
|
| def build_heatmap_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Heatmap chart: >8 countries, multi-year matrix.""" |
| df, original_n = _cap_cardinality(df, "country", 50) |
| title = _append_trim_note(title, "country", df["country"].nunique() if "country" in df.columns else 0, original_n) |
|
|
| rows = df.to_dict(orient="records") |
| tt_fmt = _compute_tooltip_format(float(df["value"].abs().max()) if "value" in df.columns else None, unit_measure) |
|
|
| |
| has_negative = df["value"].min() < 0 if "value" in df.columns else False |
| scheme = "redblue" if has_negative else "yellowgreenblue" |
|
|
| |
| |
| |
| val_series = pd.to_numeric(df["value"], errors="coerce").dropna() if "value" in df.columns else pd.Series(dtype=float) |
| if not val_series.empty: |
| val_min = float(val_series.min()) |
| val_max = float(val_series.max()) |
| if has_negative: |
| |
| abs_max = max(abs(val_min), abs(val_max)) |
| color_domain = [-abs_max, abs_max] |
| else: |
| color_domain = [val_min, val_max] |
| else: |
| color_domain = [0, 1] |
|
|
| y_enc = { |
| "field": "country", |
| "type": "nominal", |
| "axis": {"title": None, "labelFontWeight": "bold"} |
| } |
|
|
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
|
|
| color_enc = { |
| "field": "value", |
| "type": "quantitative", |
| |
| |
| "scale": {"scheme": scheme, "domain": color_domain}, |
| "legend": { |
| "type": "gradient", |
| "title": _resolve_axis_title(y_label, indicator_name), |
| "orient": "top", |
| "direction": "horizontal", |
| "gradientLength": 200, |
| |
| |
| "gradientThickness": 12, |
| "labelFontSize": 10, |
| }, |
| } |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": rows}, |
| "mark": {"type": "rect", "tooltip": True}, |
| "encoding": { |
| "x": x_enc, |
| "y": y_enc, |
| "color": color_enc, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "rect", |
| indicator_labels, |
| value_format=tt_fmt, |
| viz_data=df, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| indicator_name=indicator_name, |
| ) |
| }, |
| "width": 600, |
| "height": {"step": 15}, |
| |
| |
| |
| "config": { |
| "legend": { |
| "symbolType": "square", |
| } |
| }, |
| } |
|
|
| if result.facet_dim: |
| spec["facet"] = { |
| "field": result.facet_dim, |
| "type": "nominal", |
| "columns": 2, |
| "header": {"title": None, "labelFontWeight": "bold"} |
| } |
| spec["spec"] = { |
| "mark": {"type": "rect", "tooltip": True}, |
| "encoding": spec.pop("encoding"), |
| "width": 250, |
| "height": {"step": 15} |
| } |
| del spec["mark"] |
| del spec["width"] |
| del spec["height"] |
|
|
| return inject_wb_config(spec) |
|
|
|
|
| def build_choropleth_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Choropleth map chart: single indicator, geographic display.""" |
| |
| year_count = df["year"].nunique() if "year" in df.columns else 0 |
| if year_count > 1: |
| latest_year = df["year"].max() |
| df = df[df["year"] == latest_year].copy() |
| |
| if isinstance(title, str): |
| title = {"text": title} |
| subs = title.get("subtitle", []) |
| if isinstance(subs, str): |
| subs = [subs] |
| subs.append(f"Note: Map displays data for the most recent year ({latest_year})") |
| title["subtitle"] = subs |
|
|
| |
| |
| |
| if "country" in df.columns: |
| _WORLD_ATLAS_NAME_MAP = { |
| |
| "United States": "United States of America", |
| "Venezuela, RB": "Venezuela", |
| "Bolivia": "Bolivia", |
| "Trinidad and Tobago": "Trinidad and Tobago", |
| "Bahamas, The": "Bahamas", |
| "Gambia, The": "Gambia", |
| |
| "Russian Federation": "Russia", |
| "Slovak Republic": "Slovakia", |
| "Czech Republic": "Czechia", |
| "Kyrgyz Republic": "Kyrgyzstan", |
| "Turkiye": "Turkey", |
| "North Macedonia": "Macedonia", |
| "Bosnia and Herzegovina": "Bosnia and Herz.", |
| "Serbia": "Serbia", |
| "Kosovo": "Kosovo", |
| |
| "Egypt, Arab Rep.": "Egypt", |
| "Iran, Islamic Rep.": "Iran", |
| "Yemen, Rep.": "Yemen", |
| "West Bank and Gaza": "Palestine", |
| "Syrian Arab Republic": "Syria", |
| |
| "Korea, Rep.": "South Korea", |
| "Korea, Dem. People's Rep.": "North Korea", |
| "Viet Nam": "Vietnam", |
| "Lao PDR": "Laos", |
| "Brunei Darussalam": "Brunei", |
| "Timor-Leste": "East Timor", |
| "Micronesia, Fed. Sts.": "Micronesia", |
| "Solomon Islands": "Solomon Islands", |
| |
| "Congo, Dem. Rep.": "Dem. Rep. Congo", |
| "Congo, Rep.": "Congo", |
| "Cote d'Ivoire": "Côte d'Ivoire", |
| "Eswatini": "eSwatini", |
| "Tanzania": "United Republic of Tanzania", |
| "Cabo Verde": "Cape Verde", |
| "Sao Tome and Principe": "Sao Tome and Principe", |
| |
| "Sri Lanka": "Sri Lanka", |
| } |
| df = df.copy() |
| df["country"] = df["country"].map(lambda x: _WORLD_ATLAS_NAME_MAP.get(x, x)) |
|
|
| |
| |
| if isinstance(title, str): |
| title = {"text": title, "subtitle": []} |
| subs = title.get("subtitle", []) |
| if isinstance(subs, str): |
| subs = [subs] |
| if not any("Gray" in s for s in subs): |
| subs.append("Gray: countries not in selected set or no data for this year.") |
| title["subtitle"] = subs |
|
|
| rows = df.to_dict(orient="records") |
| max_abs = float(df["value"].abs().max()) if "value" in df.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
|
|
| has_negative = df["value"].min() < 0 if "value" in df.columns else False |
| scheme = "redblue" if has_negative else "blues" |
|
|
| |
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "width": 800, |
| "height": 450, |
| "data": { |
| "url": "https://unpkg.com/world-atlas@2.0.2/countries-110m.json", |
| "format": {"type": "topojson", "feature": "countries"} |
| }, |
| "transform": [ |
| { |
| "lookup": "properties.name", |
| "from": { |
| "data": {"values": rows}, |
| "key": "country", |
| "fields": list(df.columns) |
| } |
| } |
| ], |
| "projection": {"type": "equalEarth"}, |
| "layer": [ |
| { |
| |
| "mark": {"type": "geoshape", "fill": "#eee", "stroke": "white", "strokeWidth": 0.5} |
| }, |
| { |
| |
| "mark": {"type": "geoshape", "stroke": "white", "strokeWidth": 0.5}, |
| "transform": [{"filter": "isValid(datum.value)"}], |
| "encoding": { |
| "color": { |
| "field": "value", |
| "type": "quantitative", |
| "scale": {"scheme": scheme}, |
| "legend": { |
| "title": _resolve_axis_title(y_label, indicator_name), |
| "orient": "bottom", |
| "direction": "horizontal", |
| "gradientLength": 300 |
| } |
| }, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "geoshape", |
| indicator_labels, |
| value_format=tt_fmt, |
| dim_name_labels=result.dim_name_labels, |
| indicator_name=indicator_name, |
| ) |
| } |
| } |
| ] |
| } |
|
|
| return inject_wb_config(spec) |
|
|
|
|
| def build_stacked_area_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Stacked Area chart: multi-year, multiple series (part-to-whole).""" |
| color_dim = result.color_dim or "country" |
|
|
| |
| has_negative = df["value"].min() < 0 if "value" in df.columns else False |
| if has_negative: |
| return build_temporal_single_spec(df, title, result, indicator_labels, y_label, unit_measure) |
|
|
| df, original_n = _cap_cardinality(df, color_dim, HIGH_CARDINALITY_THRESHOLDS["line_max_series"]) |
|
|
| annotated_title = _append_breakdown_note(title, df, color_dim) |
| annotated_title = _append_trim_note(annotated_title, color_dim, df[color_dim].nunique() if color_dim in df.columns else 0, original_n) |
|
|
| rows = df.to_dict(orient="records") |
| tt_fmt = _compute_tooltip_format(float(df["value"].abs().max()) if "value" in df.columns else None, unit_measure) |
|
|
| legend_title = _TOOLTIP_SPECS.get(color_dim, {}).get("title") or color_dim.replace("_", " ").title() |
| domain = sorted(df[color_dim].dropna().unique().tolist(), key=str) if color_dim in df.columns else None |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": annotated_title, |
| "data": {"values": rows}, |
| "mark": {"type": "area", "tooltip": True, "line": True, "opacity": 0.8}, |
| "encoding": { |
| "x": _x_temporal_encoding(result.temporal_frequency), |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "stack": "zero", |
| "axis": {**_axis_style(), "title": _resolve_axis_title(y_label, indicator_name), "labelExpr": _value_label_expr(unit_measure)} |
| }, |
| "color": _color_encoding(color_dim, domain=domain, mark_type="area", legend_title=legend_title), |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "area", |
| indicator_labels, |
| value_format=tt_fmt, |
| viz_data=df, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| indicator_name=indicator_name, |
| ) |
| }, |
| "width": 600, |
| "height": 350 |
| } |
|
|
| return inject_wb_config(spec) |
|
|
|
|
| def build_stacked_bar_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Stacked Bar chart (part-to-whole / snapshot comparisons).""" |
| color_dim = result.color_dim or "sex" |
|
|
| |
| |
| |
| |
| |
| has_negative = df["value"].min() < 0 if "value" in df.columns else False |
| if has_negative and color_dim != "country": |
| return build_breakdown_comparison_spec(df, title, result, indicator_labels, y_label, unit_measure, indicator_name) |
|
|
| df, original_n = _cap_cardinality(df, color_dim, HIGH_CARDINALITY_THRESHOLDS["line_max_series"]) |
|
|
| annotated_title = _append_breakdown_note(title, df, color_dim) |
| annotated_title = _append_trim_note(annotated_title, color_dim, df[color_dim].nunique() if color_dim in df.columns else 0, original_n) |
|
|
| rows = df.to_dict(orient="records") |
| tt_fmt = _compute_tooltip_format(float(df["value"].abs().max()) if "value" in df.columns else None, unit_measure) |
|
|
| legend_title = _TOOLTIP_SPECS.get(color_dim, {}).get("title") or color_dim.replace("_", " ").title() |
| domain = sorted(df[color_dim].dropna().unique().tolist(), key=str) if color_dim in df.columns else None |
|
|
| if color_dim == "sex": |
| domain_colors = [k for k in WB_GENDER_COLORS if k in df[color_dim].unique()] |
| color_range = [WB_GENDER_COLORS[k] for k in domain_colors] |
| color_scale = {"domain": domain_colors, "range": color_range} |
| else: |
| color_scale = {"range": WB_CAT_COLORS} |
| if domain: |
| color_scale["domain"] = domain |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| if df.get("year", pd.Series()).nunique() <= 1 and df.get("country", pd.Series()).nunique() <= 1: |
| x_field = color_dim |
| elif color_dim == "country": |
| x_field = "year" if df.get("year", pd.Series()).nunique() > 1 else "country" |
| elif color_dim == "year": |
| x_field = "country" |
| else: |
| x_field = "country" if df.get("country", pd.Series()).nunique() > 1 else "year" |
| if x_field == "year" and df.get("year", pd.Series()).nunique() <= 1: |
| x_field = "country" |
|
|
| x_enc: dict |
| if x_field == "year": |
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
| else: |
| x_enc = { |
| "field": x_field, |
| "type": "nominal", |
| "axis": {"title": None, "labelFontWeight": "bold"}, |
| } |
|
|
| n_categories = df[x_field].nunique() if x_field in df.columns else 1 |
| mark_spec: dict = {"type": "bar", "tooltip": True} |
| if n_categories == 1: |
| mark_spec["size"] = 40 |
| elif n_categories == 2: |
| mark_spec["size"] = 45 |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": annotated_title, |
| "data": {"values": rows}, |
| "mark": mark_spec, |
| "encoding": { |
| "x": x_enc, |
| "y": { |
| "field": "value", |
| "type": "quantitative", |
| "stack": "zero", |
| "axis": {**_axis_style(), "title": _resolve_axis_title(y_label, indicator_name), "labelExpr": _value_label_expr(unit_measure)} |
| }, |
| "color": { |
| "field": color_dim, |
| "type": "nominal", |
| "scale": color_scale, |
| "legend": { |
| "orient": "top", |
| "title": legend_title, |
| "labelLimit": 100, |
| "columns": 3, |
| }, |
| }, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "bar", |
| indicator_labels, |
| value_format=tt_fmt, |
| viz_data=df, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| indicator_name=indicator_name, |
| ) |
| }, |
| "width": max(300, df[x_field].nunique() * 80), |
| "height": 320 |
| } |
|
|
| return inject_wb_config(spec) |
|
|
|
|
|
|
| def build_correlation_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Scatterplot: 2 indicators, single year, multi-country.""" |
| ind_cols = result.indicator_cols |
| if len(ind_cols) < 2: |
| raise ValueError("correlation spec requires exactly 2 indicator columns") |
|
|
| x_col, y_col = ind_cols[0], ind_cols[1] |
| lab = indicator_labels or {} |
| x_label = lab.get(x_col, x_col.replace("_", " ").title()) |
| y_label = lab.get(y_col, y_col.replace("_", " ").title()) |
|
|
| df_clean = df.dropna(subset=[x_col, y_col]).copy() |
| if not df_clean.empty: |
| x_min, x_max = df_clean[x_col].min(), df_clean[x_col].max() |
| y_min, y_max = df_clean[y_col].min(), df_clean[y_col].max() |
| x_range = x_max - x_min if x_max != x_min else 1.0 |
| y_range = y_max - y_min if y_max != y_min else 1.0 |
|
|
| df_clean["_x_norm"] = (df_clean[x_col] - x_min) / x_range |
| df_clean["_y_norm"] = (df_clean[y_col] - y_min) / y_range |
|
|
| |
| df_clean["_dist_from_center"] = (df_clean["_x_norm"] - 0.5) ** 2 + (df_clean["_y_norm"] - 0.5) ** 2 |
|
|
| |
| df_sorted = df_clean.sort_values(by="_dist_from_center", ascending=False) |
|
|
| labeled_points = [] |
| show_labels = {} |
|
|
| |
| EPSILON_X = 0.06 |
| EPSILON_Y = 0.04 |
|
|
| for idx, row in df_sorted.iterrows(): |
| x_n = row["_x_norm"] |
| y_n = row["_y_norm"] |
|
|
| collision = False |
| for lx, ly in labeled_points: |
| if abs(x_n - lx) < EPSILON_X and abs(y_n - ly) < EPSILON_Y: |
| collision = True |
| break |
|
|
| if not collision: |
| show_labels[idx] = True |
| labeled_points.append((x_n, y_n)) |
| else: |
| show_labels[idx] = False |
|
|
| df_clean["_show_label"] = df_clean.index.map(show_labels) |
| else: |
| df_clean["_show_label"] = True |
|
|
| rows = df_clean.to_dict(orient="records") |
|
|
| color_dim = result.color_dim or "country" |
| color_enc = _color_encoding(color_dim) |
|
|
| |
| |
| if isinstance(color_enc, dict) and "legend" in color_enc: |
| color_enc["legend"] = None |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": rows}, |
| "layer": [ |
| { |
| "mark": { |
| "type": "circle", |
| "opacity": 0.85, |
| "stroke": WB_WHITE, |
| "strokeWidth": 1, |
| "size": 70 |
| }, |
| "encoding": { |
| "x": { |
| "field": x_col, |
| "type": "quantitative", |
| "axis": _axis_style(x_label), |
| "scale": {"zero": False}, |
| }, |
| "y": { |
| "field": y_col, |
| "type": "quantitative", |
| "axis": _axis_style(_resolve_axis_title(y_label, indicator_name)), |
| "scale": {"zero": False}, |
| }, |
| "color": color_enc, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), "point", lab, viz_data=df |
| ), |
| } |
| }, |
| { |
| "mark": { |
| "type": "text", |
| "dy": -10, |
| "fontSize": 10, |
| "fontWeight": "bold" |
| }, |
| "encoding": { |
| "x": { |
| "field": x_col, |
| "type": "quantitative" |
| }, |
| "y": { |
| "field": y_col, |
| "type": "quantitative" |
| }, |
| "text": { |
| "field": color_dim, |
| "type": "nominal" |
| }, |
| "color": color_enc, |
| "opacity": { |
| "condition": {"test": "datum._show_label", "value": 1}, |
| "value": 0 |
| } |
| } |
| } |
| ], |
| "width": 600, |
| "height": 450, |
| } |
|
|
| return inject_wb_config(spec) |
|
|
|
|
| def build_correlation_temporal_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Connected scatterplot: 2 indicators, multi-country, multi-year.""" |
| ind_cols = result.indicator_cols |
| if len(ind_cols) < 2: |
| raise ValueError("correlation_temporal spec requires 2 indicator columns") |
|
|
| x_col, y_col = ind_cols[0], ind_cols[1] |
| lab = indicator_labels or {} |
| x_label = lab.get(x_col, x_col.replace("_", " ").title()) |
| y_label = lab.get(y_col, y_col.replace("_", " ").title()) |
|
|
| rows = df.dropna(subset=[x_col, y_col]).to_dict(orient="records") |
| color_dim = result.color_dim or "country" |
|
|
| |
| base_enc: dict = { |
| "x": { |
| "field": x_col, |
| "type": "quantitative", |
| "axis": _axis_style(x_label), |
| "scale": {"zero": False}, |
| }, |
| "y": { |
| "field": y_col, |
| "type": "quantitative", |
| "axis": _axis_style(y_label), |
| "scale": {"zero": False}, |
| }, |
| "color": _color_encoding(color_dim), |
| "order": {"field": "year", "type": "temporal"}, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), "line", lab, viz_data=df |
| ), |
| } |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": rows}, |
| "layer": [ |
| { |
| "mark": {"type": "line", "strokeWidth": 2, "opacity": 0.6}, |
| "encoding": {k: v for k, v in base_enc.items() if k != "tooltip"}, |
| }, |
| { |
| "mark": {"type": "circle", "size": 40, "opacity": 0.9}, |
| "encoding": base_enc, |
| }, |
| ], |
| "width": 550, |
| "height": 450, |
| } |
| return inject_wb_config(spec) |
|
|
|
|
| def build_temporal_multi_indicator_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Faceted (Small Multiples) multi-axis line chart: 2-4 indicators, multi-year. |
| |
| Complies with Grammar of Graphics by strictly avoiding dual-axis overlapping. |
| Uses Vega-Lite `vconcat` to create vertically stacked charts that share a |
| common X-axis (time), giving each indicator its own isolated Y-axis plane. |
| """ |
| ind_cols = result.indicator_cols |
| if not ind_cols: |
| raise ValueError("temporal_multi_indicator spec requires indicator_cols") |
|
|
| lab = indicator_labels or {} |
| df_copy = df.copy() |
|
|
| |
| should_layer = False |
| fold_cols = [col for col in ind_cols if col in df_copy.columns] |
| if len(fold_cols) >= 2: |
| try: |
| max_vals = [] |
| for col in fold_cols: |
| max_vals.append(df_copy[col].abs().max()) |
| if len(max_vals) == len(fold_cols) and all(v is not None and not pd.isna(v) for v in max_vals): |
| max_val = max(max_vals) |
| min_val = min(max_vals) |
| |
| if min_val > 0 and (max_val / min_val) <= 10.0: |
| should_layer = True |
| except Exception: |
| pass |
|
|
| if should_layer: |
| |
| raw_titles = [lab.get(col, col) for col in fold_cols] |
| col_to_legend = dict(zip(fold_cols, raw_titles)) |
| if len(raw_titles) >= 2: |
| try: |
| import os |
| common_pref = os.path.commonprefix(raw_titles) |
| common_pref_stripped = common_pref.strip() |
| while common_pref_stripped and common_pref_stripped[-1] in (",", "-", " ", "|", ":", ";", "("): |
| common_pref_stripped = common_pref_stripped[:-1].strip() |
|
|
| temp_cleaned = [] |
| for t in raw_titles: |
| part = t[len(common_pref):].strip() |
| while part and part[0] in (",", "-", " ", "|", ":", ";", "("): |
| part = part[1:].strip() |
| first_segment = part |
| for sep in (",", "(", " - ", " | "): |
| if sep in part: |
| candidate = part.split(sep)[0].strip() |
| if len(candidate) >= 3: |
| first_segment = candidate |
| break |
| temp_cleaned.append(first_segment.title()) |
|
|
| if len(set(temp_cleaned)) == len(temp_cleaned) and all(len(x) >= 2 for x in temp_cleaned): |
| col_to_legend = dict(zip(fold_cols, temp_cleaned)) |
| except Exception: |
| pass |
|
|
| |
| final_y_label = y_label |
| if final_y_label in ("Value", "obs_value", None) or len(final_y_label) > 60: |
| try: |
| import os |
| common_pref = os.path.commonprefix(raw_titles) |
| common_pref_stripped = common_pref.strip() |
| while common_pref_stripped and common_pref_stripped[-1] in (",", "-", " ", "|", ":", ";", "("): |
| common_pref_stripped = common_pref_stripped[:-1].strip() |
| if len(common_pref_stripped) >= 8: |
| is_pct = any("%" in u.lower() or "percent" in u.lower() or "rate" in u.lower() or "share" in u.lower() for u in (unit_measure or "", *raw_titles)) |
| if is_pct and not common_pref_stripped.endswith("(%)") and "%" not in common_pref_stripped: |
| final_y_label = f"{common_pref_stripped} (%)" |
| else: |
| final_y_label = common_pref_stripped |
| except Exception: |
| pass |
|
|
| |
| melted_df = df_copy.melt( |
| id_vars=[c for c in df_copy.columns if c not in fold_cols], |
| value_vars=fold_cols, |
| var_name="indicator_name_melted", |
| value_name="indicator_value_melted", |
| ) |
| |
| melted_df["indicator_name_melted"] = melted_df["indicator_name_melted"].map(col_to_legend) |
|
|
| |
| if "year" in melted_df.columns: |
| try: |
| parsed_years = pd.to_datetime(melted_df["year"].astype(str), errors="coerce") |
| if parsed_years.notna().any(): |
| if result.temporal_frequency == "monthly": |
| melted_df["year"] = parsed_years.dt.strftime("%Y-%m") |
| else: |
| melted_df["year"] = parsed_years.dt.strftime("%Y") |
| else: |
| melted_df["year"] = melted_df["year"].astype(str) |
| except Exception: |
| melted_df["year"] = melted_df["year"].astype(str) |
|
|
| rows = melted_df.to_dict(orient="records") |
| label_expr = _value_label_expr(unit_measure) |
|
|
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
| max_abs = float(melted_df["indicator_value_melted"].abs().max()) |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
| color_scale_range = WB_CAT_COLORS[:len(fold_cols)] |
|
|
| y_axis = { |
| **_axis_style(), |
| "title": final_y_label or "Value", |
| "labelExpr": label_expr, |
| } |
|
|
| tooltip_cols = ["year", "indicator_name_melted", "indicator_value_melted"] |
| if "country" in melted_df.columns: |
| tooltip_cols.append("country") |
|
|
| encoding = { |
| "x": x_enc, |
| "y": { |
| "field": "indicator_value_melted", |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": False}, |
| }, |
| "color": { |
| "field": "indicator_name_melted", |
| "type": "nominal", |
| "scale": { |
| "range": color_scale_range |
| }, |
| "legend": { |
| "title": None, |
| "orient": "bottom", |
| "offset": 12, |
| } |
| }, |
| "tooltip": build_structured_tooltips( |
| tooltip_cols, "line", lab, value_format=tt_fmt, viz_data=melted_df, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels={ |
| "indicator_name_melted": "Indicator", |
| "indicator_value_melted": "Value", |
| }, |
| ), |
| } |
|
|
| mark_spec = { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| "tooltip": True |
| } |
|
|
| data_spec = {"values": rows} |
| if melted_df["year"].nunique() <= 1 if "year" in melted_df.columns else True: |
| data_spec["format"] = {"parse": {"year": "string"}} |
|
|
| |
| if "country" in melted_df.columns and melted_df["country"].nunique() > 1: |
| countries = sorted(melted_df["country"].dropna().unique().tolist()) |
| panels = [] |
| for i, c in enumerate(countries): |
| |
| x_enc_panel = x_enc.copy() |
| if i < len(countries) - 1: |
| x_enc_panel = {**x_enc_panel, "axis": {**x_enc_panel.get("axis", {}), "labels": False, "title": None}} |
|
|
| panel = { |
| "title": { |
| "text": str(c), |
| "fontSize": 12, |
| "fontWeight": "bold", |
| "anchor": "start", |
| }, |
| "transform": [{"filter": f"datum.country == '{c}'"}], |
| "width": 680, |
| "height": 180, |
| "mark": mark_spec, |
| "encoding": { |
| **encoding, |
| "x": x_enc_panel, |
| } |
| } |
| panels.append(panel) |
|
|
| spec = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": data_spec, |
| "vconcat": panels, |
| "resolve": { |
| "scale": {"x": "shared"}, |
| "axis": {"x": "independent"} |
| } |
| } |
| return inject_wb_config(spec) |
|
|
| return { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": data_spec, |
| "width": 680, |
| "height": 350, |
| "mark": mark_spec, |
| "encoding": encoding, |
| } |
|
|
|
|
| |
| |
| |
| if "year" in df_copy.columns: |
| try: |
| parsed_years = pd.to_datetime(df_copy["year"].astype(str), errors="coerce") |
| if parsed_years.notna().any(): |
| if result.temporal_frequency == "monthly": |
| df_copy["year"] = parsed_years.dt.strftime("%Y-%m") |
| else: |
| df_copy["year"] = parsed_years.dt.strftime("%Y") |
| else: |
| df_copy["year"] = df_copy["year"].astype(str) |
| except Exception: |
| df_copy["year"] = df_copy["year"].astype(str) |
|
|
| rows = df_copy.to_dict(orient="records") |
|
|
| label_expr = _value_label_expr(unit_measure) |
| is_bar_chart = (result.mark_hint == "bar") or (df_copy["year"].nunique() <= 1 if "year" in df_copy.columns else True) |
|
|
| |
| country_cnt = df_copy["country"].nunique() if "country" in df_copy.columns else 0 |
| if len(ind_cols) == 2 and not should_layer and country_cnt <= 1 and not is_bar_chart: |
| chart0_color = WB_CAT_COLORS[0] |
| chart1_color = WB_CAT_COLORS[1] |
|
|
| col0 = ind_cols[0] |
| col1 = ind_cols[1] |
|
|
| col0_label = lab.get(col0, col0.replace("_", " ").title()) |
| col1_label = lab.get(col1, col1.replace("_", " ").title()) |
|
|
| max_abs0 = float(df_copy[col0].abs().max()) if col0 in df_copy.columns else None |
| max_abs1 = float(df_copy[col1].abs().max()) if col1 in df_copy.columns else None |
|
|
| tt_fmt0 = _compute_tooltip_format(max_abs0, unit_measure) |
| tt_fmt1 = _compute_tooltip_format(max_abs1, unit_measure) |
|
|
| if df_copy["year"].nunique() <= 1 if "year" in df_copy.columns else True: |
| x_enc = { |
| "field": "year" if "year" in df_copy.columns else "TIME_PERIOD", |
| "type": "nominal", |
| "axis": {"title": None} |
| } |
| else: |
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
|
|
| layer0 = { |
| "mark": { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| "color": chart0_color, |
| "tooltip": True |
| }, |
| "encoding": { |
| "x": x_enc, |
| "y": { |
| "field": col0, |
| "type": "quantitative", |
| "axis": { |
| **_axis_style(), |
| "title": col0_label, |
| "titleColor": chart0_color, |
| "labelColor": chart0_color, |
| "labelExpr": label_expr, |
| }, |
| "scale": {"zero": False} |
| }, |
| "tooltip": build_structured_tooltips( |
| _multi_indicator_tooltip_columns(list(df_copy.columns), col0), |
| "line", lab, value_format=tt_fmt0, viz_data=df_copy, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ) |
| } |
| } |
|
|
| layer1 = { |
| "mark": { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| "color": chart1_color, |
| "tooltip": True |
| }, |
| "encoding": { |
| "x": x_enc, |
| "y": { |
| "field": col1, |
| "type": "quantitative", |
| "axis": { |
| **_axis_style(), |
| "title": col1_label, |
| "titleColor": chart1_color, |
| "labelColor": chart1_color, |
| "labelExpr": label_expr, |
| }, |
| "scale": {"zero": False} |
| }, |
| "tooltip": build_structured_tooltips( |
| _multi_indicator_tooltip_columns(list(df_copy.columns), col1), |
| "line", lab, value_format=tt_fmt1, viz_data=df_copy, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ) |
| } |
| } |
|
|
| data_spec = {"values": rows} |
| if df_copy["year"].nunique() <= 1 if "year" in df_copy.columns else True: |
| data_spec["format"] = {"parse": {"year": "string"}} |
|
|
| spec = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": data_spec, |
| "width": 680, |
| "height": 350, |
| "layer": [layer0, layer1], |
| "resolve": { |
| "scale": {"y": "independent"} |
| } |
| } |
| return inject_wb_config(spec) |
|
|
| charts = [] |
|
|
| for i, col in enumerate(ind_cols): |
| color = WB_CAT_COLORS[i % len(WB_CAT_COLORS)] |
| col_label = lab.get(col, col.replace("_", " ").title()) |
| max_abs = float(df_copy[col].abs().max()) if col in df_copy.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
| y_axis = { |
| **_axis_style(), |
| "title": None, |
| "labelExpr": label_expr, |
| } |
|
|
|
|
| |
| if df_copy["year"].nunique() <= 1 if "year" in df_copy.columns else True: |
| if df_copy["country"].nunique() > 1 if "country" in df_copy.columns else False: |
| x_enc = { |
| "field": "country", |
| "type": "nominal", |
| "axis": {"title": "Country"} |
| } |
| else: |
| x_enc = { |
| "field": "year" if "year" in df_copy.columns else "TIME_PERIOD", |
| "type": "nominal", |
| "axis": {"title": None} |
| } |
| else: |
| x_enc = _x_temporal_encoding(result.temporal_frequency) |
|
|
| |
| if i < len(ind_cols) - 1: |
| x_enc = {**x_enc, "axis": {**x_enc.get("axis", {}), "labels": False, "title": None}} |
|
|
| tooltip_cols = _multi_indicator_tooltip_columns(list(df_copy.columns), col) |
|
|
| |
| if not is_bar_chart and "country" in df_copy.columns and df_copy["country"].nunique() > 1: |
| color_enc = { |
| "field": "country", |
| "type": "nominal", |
| "scale": {"range": WB_CAT_COLORS}, |
| "legend": { |
| "title": None, |
| "orient": "bottom", |
| "offset": 12 |
| } if i == len(ind_cols) - 1 else None |
| } |
| else: |
| color_enc = {"value": color} |
|
|
| layer_enc: dict = { |
| "x": x_enc, |
| "y": { |
| "field": col, |
| "type": "quantitative", |
| "axis": y_axis, |
| "scale": {"zero": is_bar_chart}, |
| }, |
| "color": color_enc, |
| "tooltip": build_structured_tooltips( |
| tooltip_cols, "bar" if is_bar_chart else "line", lab, value_format=tt_fmt, viz_data=df_copy, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
|
|
| |
| if is_bar_chart: |
| mark_spec = { |
| "type": "bar", |
| "color": color, |
| "size": 40, |
| "tooltip": True |
| } |
| else: |
| mark_spec = { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| "tooltip": True |
| } |
| |
| if not ("country" in df_copy.columns and df_copy["country"].nunique() > 1): |
| mark_spec["color"] = color |
|
|
|
|
| charts.append( |
| { |
| "title": { |
| "text": col_label, |
| "color": color, |
| "fontSize": 12, |
| "fontWeight": "bold", |
| "anchor": "start", |
| "offset": 4 |
| }, |
| "width": 680, |
| "height": 140, |
| "mark": mark_spec, |
| "encoding": layer_enc, |
| } |
| ) |
|
|
| data_spec: dict = {"values": rows} |
| if df_copy["year"].nunique() <= 1 if "year" in df_copy.columns else True: |
| data_spec["format"] = {"parse": {"year": "string"}} |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": data_spec, |
| "vconcat": charts, |
| "resolve": { |
| "scale": {"x": "shared"}, |
| "axis": {"x": "independent"} |
| }, |
| } |
| return inject_wb_config(spec) |
|
|
|
|
| def build_fallback_line_spec( |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Fallback: best-effort line chart for unclassified data shapes.""" |
| cols = set(df.columns) |
| x_col = ( |
| "year" |
| if "year" in cols |
| else ("time_period" if "time_period" in cols else df.columns[0]) |
| ) |
| y_col = ( |
| "value" |
| if "value" in cols |
| else ("obs_value" if "obs_value" in cols else df.columns[-1]) |
| ) |
| max_abs = float(df[y_col].abs().max()) if y_col in df.columns else None |
| tt_fmt = _compute_tooltip_format(max_abs, unit_measure) |
| y_ax = { |
| **_axis_style(), |
| "title": None, |
| "labelExpr": _value_label_expr(unit_measure), |
| } |
|
|
| encoding: dict = { |
| "x": _x_temporal_encoding(result.temporal_frequency) if "year" in x_col else { |
| "field": x_col, |
| "type": "ordinal", |
| "axis": _axis_style(), |
| }, |
| "y": {"field": y_col, "type": "quantitative", "axis": y_ax}, |
| "tooltip": build_structured_tooltips( |
| list(df.columns), |
| "line", |
| indicator_labels, |
| value_format=tt_fmt, |
| viz_data=df, |
| temporal_freq=result.temporal_frequency, |
| dim_name_labels=result.dim_name_labels, |
| ), |
| } |
| if result.color_dim and result.color_dim in cols: |
| n_items = df[result.color_dim].nunique() |
| _domain_labels = list(df[result.color_dim].unique()) |
| encoding["color"] = _color_encoding( |
| result.color_dim, |
| mark_type="line", |
| n_items=n_items, |
| legend_title=result.dim_name_labels.get(result.color_dim), |
| domain_labels=_domain_labels, |
| ) |
|
|
| spec: dict = { |
| "$schema": _vl_schema(), |
| "title": title, |
| "data": {"values": df.to_dict(orient="records")}, |
| "mark": { |
| "type": "line", |
| "strokeWidth": 3, |
| "strokeCap": "round", |
| "point": _LINE_HOVER_POINT, |
| }, |
| "encoding": encoding, |
| "width": 600, |
| "height": 350, |
| } |
| return inject_wb_config(spec) |
|
|
|
|
|
|
|
|
|
|
| STRATEGY_BUILDERS: dict[ChartStrategy, callable] = { |
| ChartStrategy.TEMPORAL_SINGLE: build_temporal_single_spec, |
| ChartStrategy.CROSS_SECTIONAL: build_cross_sectional_spec, |
| ChartStrategy.DISTRIBUTION: build_distribution_spec, |
| ChartStrategy.BREAKDOWN_COMPARISON: build_breakdown_comparison_spec, |
| ChartStrategy.SMALL_MULTIPLES: build_small_multiples_spec, |
| ChartStrategy.HEATMAP: build_heatmap_spec, |
| ChartStrategy.STACKED_AREA: build_stacked_area_spec, |
| ChartStrategy.STACKED_BAR: build_stacked_bar_spec, |
| ChartStrategy.CHOROPLETH: build_choropleth_spec, |
| ChartStrategy.CORRELATION: build_correlation_spec, |
| ChartStrategy.CORRELATION_TEMPORAL: build_correlation_temporal_spec, |
| ChartStrategy.TEMPORAL_MULTI_IND: build_temporal_multi_indicator_spec, |
| ChartStrategy.FALLBACK_LINE: build_fallback_line_spec, |
| } |
|
|
|
|
| def resolve_unmapped_disaggregations(df: pd.DataFrame, result: StrategyResult) -> pd.DataFrame: |
| """Filter or collapse disaggregation dimensions that are not mapped in the strategy.""" |
| if df.empty: |
| return df |
|
|
| mapped_dims = { |
| "year", |
| "value", |
| "country", |
| result.color_dim, |
| result.facet_dim, |
| result.secondary_color_dim, |
| result.x_dim, |
| result.y_dim, |
| } |
|
|
| |
| disagg_dims = [ |
| "sex", |
| "age", |
| "urbanisation", |
| "residence", |
| "comp_breakdown_1", |
| "comp_breakdown_2", |
| "comp_breakdown_3", |
| "unit_measure", |
| ] |
| unmapped_dims = [ |
| c for c in df.columns if c in disagg_dims and c not in mapped_dims |
| ] |
|
|
| if not unmapped_dims: |
| return df |
|
|
| import logging |
| logger = logging.getLogger("data360.viz_config") |
|
|
| df_resolved = df.copy() |
| key_cols = [c for c in df_resolved.columns if c not in unmapped_dims and c != "value"] |
|
|
| for dim in unmapped_dims: |
| if df_resolved[dim].nunique() <= 1: |
| continue |
|
|
| |
| sentinels = { |
| "_t", |
| "_z", |
| "total", |
| "all", |
| "overall", |
| "u", |
| "not applicable", |
| "not_applicable", |
| "notapplicable", |
| } |
| unique_vals = df_resolved[dim].dropna().unique() |
| sentinel_found = None |
| for val in unique_vals: |
| val_str = str(val).strip().lower() |
| if val_str in sentinels: |
| sentinel_found = val |
| break |
|
|
| if sentinel_found is not None: |
| df_resolved = df_resolved[df_resolved[dim] == sentinel_found] |
| logger.info("Filtered unmapped dimension '%s' to sentinel '%s'", dim, sentinel_found) |
|
|
| |
| if df_resolved.duplicated(subset=key_cols).any(): |
| n_before = len(df_resolved) |
| df_resolved = ( |
| df_resolved.groupby(key_cols, sort=False)["value"].mean().reset_index() |
| ) |
| n_collapsed = n_before - len(df_resolved) |
| logger.warning( |
| "Collapsed %d duplicate rows on unmapped dimensions %s", |
| n_collapsed, |
| unmapped_dims, |
| ) |
|
|
| return df_resolved |
|
|
|
|
| def dispatch_spec( |
| strategy: ChartStrategy, |
| df: pd.DataFrame, |
| title: str | dict, |
| result: StrategyResult, |
| indicator_labels: dict[str, str] | None = None, |
| y_label: str = "Value", |
| x_label: str = "Value", |
| unit_measure: str | None = None, |
| indicator_name: str | None = None, |
| ) -> dict: |
| """Call the right spec builder for the given strategy.""" |
| df = resolve_unmapped_disaggregations(df, result) |
| builder = STRATEGY_BUILDERS[strategy] |
| if strategy in ( |
| ChartStrategy.TEMPORAL_SINGLE, |
| ChartStrategy.TEMPORAL_MULTI_IND, |
| ChartStrategy.BREAKDOWN_COMPARISON, |
| ChartStrategy.SMALL_MULTIPLES, |
| ChartStrategy.STACKED_AREA, |
| ChartStrategy.STACKED_BAR, |
| ChartStrategy.CHOROPLETH, |
| ChartStrategy.FALLBACK_LINE, |
| ChartStrategy.HEATMAP, |
| ): |
| return builder(df, title, result, indicator_labels, y_label, unit_measure, indicator_name=indicator_name) |
| elif strategy in (ChartStrategy.CROSS_SECTIONAL, ChartStrategy.DISTRIBUTION): |
| return builder(df, title, result, indicator_labels, x_label, unit_measure, indicator_name=indicator_name) |
| else: |
| return builder(df, title, result, indicator_labels, indicator_name=indicator_name) |
|
|
|
|
| |
| |
| |
|
|
| HIGH_CARDINALITY_THRESHOLDS: dict[str, int] = { |
| |
| |
| "line_max_series": 12, |
| |
| "beeswarm_threshold": 20, |
| |
| "facet_threshold": 4, |
| |
| "top_n_series": 12, |
| |
| |
| |
| "small_multiples_max_facets": 6, |
| |
| |
| "cross_sectional_max_items": 20, |
| |
| "heatmap_threshold": 12, |
| } |
|
|
| |
| SMALL_MULTIPLES_MAX_FACETS: int = HIGH_CARDINALITY_THRESHOLDS["small_multiples_max_facets"] |
|
|
| |
| |
| |
| MAX_END_LABEL_SERIES: int = 10 |
|
|
| |
| |
| |
| |
| |
| CORRELATION_TEMPORAL_AUTO_MAX_COUNTRIES: int = 8 |
| CORRELATION_TEMPORAL_AUTO_MAX_YEARS: int = 8 |
|
|
|
|
|
|
| |
| def should_use_beeswarm( |
| viz_data: pd.DataFrame, |
| chart_type: str | None = None, |
| color_dim: str | None = None, |
| ) -> bool: |
| if color_dim is None or color_dim not in viz_data.columns: |
| return False |
| if chart_type and chart_type not in (None, "line", "area"): |
| return False |
| series_count = viz_data[color_dim].nunique() |
| year_count = viz_data["year"].nunique() if "year" in viz_data.columns else 0 |
| return ( |
| series_count > HIGH_CARDINALITY_THRESHOLDS["beeswarm_threshold"] |
| and year_count <= 1 |
| ) |
|
|
|
|
| def build_beeswarm_spec( |
| viz_data: pd.DataFrame, |
| title: str, |
| value_col: str = "value", |
| color_col: str = "country", |
| ) -> dict: |
| """Legacy alias → delegates to build_distribution_spec.""" |
| r = StrategyResult(ChartStrategy.DISTRIBUTION, "beeswarm", color_dim=color_col) |
| |
| df = viz_data.copy() |
| if value_col != "value" and value_col in df.columns: |
| df = df.rename(columns={value_col: "value"}) |
| return build_distribution_spec(df, title, r) |
|
|
|
|
|
|
|
|
|
|
| |
| |
| |
|
|
| FREQUENCY_TO_TIMEUNIT: dict[str, str] = { |
| "A": "utcyear", |
| "M": "utcyearmonth", |
| "Q": "utcyearquarter", |
| } |
|
|
| PERIODICITY_KEYWORDS: dict[str, list[str]] = { |
| "A": ["annual", "yearly"], |
| "M": ["month", "monthly"], |
| "Q": ["quarter", "quarterly"], |
| } |
|
|
| CHART_TYPE_KEYWORDS: dict[str, list[str]] = { |
| "line": ["line", "trend", "time series", "over time"], |
| "stacked_bar": ["stacked_bar", "stacked_column", "stacked bar", "stacked column"], |
| "bar": ["bar", "column", "ranking", "compare", "histogram"], |
| "point": ["scatter", "point", "dot", "correlation", "bubble"], |
| "area": ["area", "filled", "cumulative", "stacked"], |
| "tick": ["tick", "strip", "beeswarm", "distribution"], |
| |
| |
| "heatmap": ["heatmap", "heat map", "heat"], |
| "map": ["map", "choropleth", "geoshape", "geographic"], |
| "small_multiples": ["facet", "small multiples", "small_multiples", "grid"], |
| } |
| DEFAULT_CHART_TYPE: str = "line" |
|
|
|
|
| def parse_chart_type_hint(chart_type: str | None) -> str: |
| if not chart_type: |
| return DEFAULT_CHART_TYPE |
| hint = chart_type.lower().strip() |
| for mark_type, keywords in CHART_TYPE_KEYWORDS.items(): |
| if any(keyword in hint for keyword in keywords): |
| return mark_type |
| return DEFAULT_CHART_TYPE |
|
|
|
|
| _REASON_CHART_PHRASES: dict[str, str] = { |
| "line": "line chart", |
| "bar": "bar chart", |
| "stacked_bar": "stacked bar chart", |
| "area": "area chart", |
| "point": "point chart", |
| "tick": "strip chart", |
| } |
|
|
|
|
| def chart_type_phrase_for_reason(mark_type: str | None) -> str: |
| """Human phrase for strategy / tool ``reason`` (aligned with mark type hint or render).""" |
| if not mark_type: |
| return _REASON_CHART_PHRASES["line"] |
| normalized = mark_type.lower().strip() |
| return _REASON_CHART_PHRASES.get(normalized, f"{normalized} chart") |
|
|
|
|
| def patch_strategy_reason_chart_phrase(reason: str, mark_type: str) -> str: |
| """Replace the trailing ``→ …`` segment so it reflects the given mark type.""" |
| sep = " → " |
| if sep not in reason: |
| return reason |
| prefix, _old = reason.rsplit(sep, 1) |
| return f"{prefix}{sep}{chart_type_phrase_for_reason(mark_type)}" |
|
|
|
|
| def extract_top_level_mark_type(vl_spec: dict) -> str | None: |
| """Best-effort mark ``type`` from a single-view Vega-Lite spec.""" |
| mark = vl_spec.get("mark") |
| if isinstance(mark, str): |
| return mark |
| if isinstance(mark, dict): |
| t = mark.get("type") |
| return t if isinstance(t, str) else None |
| return None |
|
|
|
|
| def get_main_data_layer(spec: dict) -> dict: |
| """Return the primary data-bearing layer from a flat or layered Vega-Lite spec. |
| |
| Phase 6 and Phase 7 may convert a flat ``{mark, encoding}`` spec into a |
| layered spec ``{layer: [{mark, encoding}, ...]}`` by appending a zero-line |
| rule and/or a text end-label layer. Tests and post-processing code that |
| inspect ``spec["mark"]`` or ``spec["encoding"]`` must call this helper to |
| get the correct sub-spec regardless of whether the wrapping occurred. |
| |
| Convention: the first layer ``layer[0]`` is always the primary data mark. |
| Decoration layers (rule, text) are appended after it. |
| """ |
| if "layer" in spec: |
| layers = spec["layer"] |
| if layers: |
| return layers[0] |
| return spec |
|
|
|
|
| def infer_frequency_from_periodicity(periodicity: str) -> str | None: |
| pl = periodicity.lower() |
| for code, kws in PERIODICITY_KEYWORDS.items(): |
| if any(kw in pl for kw in kws): |
| return code |
| return None |
|
|
|
|
| def should_use_temporal_x_axis( |
| viz_data: pd.DataFrame, chart_type: str | None, available_dimensions: list[str] |
| ) -> tuple[bool, str | None]: |
| if "year" not in available_dimensions: |
| return False, _select_categorical_dimension(available_dimensions) |
| year_count = viz_data["year"].nunique() if "year" in viz_data.columns else 0 |
| if year_count > 1: |
| return True, None |
| mark_type = parse_chart_type_hint(chart_type) if chart_type else "line" |
| pref = {"tick": 1.0, "point": 0.7, "bar": 0.5, "line": 0.2, "area": 0.1} |
| cat_field = _select_categorical_dimension(available_dimensions) |
| if cat_field is None: |
| return True, None |
| if pref.get(mark_type, 0.5) >= 0.5: |
| return False, cat_field |
| return True, None |
|
|
|
|
| def _select_categorical_dimension(available_dimensions: list[str]) -> str | None: |
| for dim in ["country", "sex", "age", "urbanisation", "residence", "education", "income_group"]: |
| if dim in available_dimensions: |
| return dim |
| for dim in available_dimensions: |
| if dim not in ["year", "value", "time_period", "obs_value"]: |
| return dim |
| return None |
|
|
|
|
| |
| |
| |
|
|
|
|
| @dataclass |
| class DataPreparationRule: |
| chart_type: str |
| frequency: str | None |
| action: Literal["year_strings", "datetime"] |
| description: str |
|
|
|
|
| DATA_PREPARATION_RULES: list[DataPreparationRule] = [ |
| DataPreparationRule( |
| "bar", "A", "year_strings", "Bar charts with annual data use year strings" |
| ), |
| DataPreparationRule( |
| "bar", |
| None, |
| "year_strings", |
| "Bar charts when API frequency is unknown — assume annual WDI-style years", |
| ), |
| DataPreparationRule("*", "*", "datetime", "Default: datetime"), |
| ] |
|
|
|
|
| def get_data_preparation_action( |
| chart_type: str, frequency: str | None |
| ) -> Literal["year_strings", "datetime"]: |
| for rule in DATA_PREPARATION_RULES: |
| if (rule.chart_type == "*" or rule.chart_type == chart_type) and ( |
| rule.frequency == "*" or rule.frequency == frequency |
| ): |
| return rule.action |
| return "datetime" |
|
|
|
|
| _YEAR_GAP_FILL_MAX_SPAN = 400 |
|
|
|
|
| def frequency_allows_annual_year_gap_fill(data_frequency: str | None) -> bool: |
| """True when data are treated as annual so missing calendar years can be inserted.""" |
| if data_frequency is None or not str(data_frequency).strip(): |
| return True |
| code = str(data_frequency).strip().upper() |
| if code in FREQUENCY_TO_TIMEUNIT and code != "A": |
| return False |
| return code in ("A", "ANNUAL", "Y", "YEAR", "YA") |
|
|
|
|
| def _clone_year_field(y_int: int, sample: Any) -> Any: |
| """Match ``year`` dtype/shape used in the source group (string, datetime, int).""" |
| if isinstance(sample, str): |
| stripped = sample.strip() |
| if len(stripped) == 4 and stripped.isdigit(): |
| return str(y_int) |
| return pd.Timestamp(year=y_int, month=1, day=1) |
| if isinstance(sample, pd.Timestamp): |
| return pd.Timestamp(year=y_int, month=1, day=1) |
| if isinstance(sample, Integral) and not isinstance(sample, bool): |
| return int(y_int) |
| if isinstance(sample, float) and not pd.isna(sample) and sample == int(sample): |
| return int(y_int) |
| return pd.Timestamp(year=y_int, month=1, day=1) |
|
|
|
|
| def _coerce_year_column_to_int(series: pd.Series) -> pd.Series: |
| if pd.api.types.is_datetime64_any_dtype(series): |
| return series.dt.year.astype("Int64") |
| parsed = pd.to_datetime(series.astype(str), errors="coerce") |
| if parsed.notna().mean() >= 0.99 and parsed.notna().any(): |
| return parsed.dt.year.astype("Int64") |
| num = pd.to_numeric(series, errors="coerce") |
| return num.round().astype("Int64") |
|
|
|
|
| def fill_missing_calendar_years_annual( |
| df: pd.DataFrame, |
| data_frequency: str | None, |
| ) -> pd.DataFrame: |
| """Insert NaN rows for missing integer calendar years within each series' span. |
| |
| Each *series* is defined by every column except ``year`` and ``value`` (e.g. one |
| country). For that series, all calendar years from min(year) to max(year) appear |
| exactly once; gaps in the source (e.g. no 2010) become explicit rows with null |
| ``value``. Skipped when frequency is not annual, years cannot be coerced, any |
| (series, year) duplicates exist, or the span exceeds ``_YEAR_GAP_FILL_MAX_SPAN``. |
| """ |
| if df.empty or "year" not in df.columns or "value" not in df.columns: |
| return df |
| if not frequency_allows_annual_year_gap_fill(data_frequency): |
| return df |
|
|
| y_int = _coerce_year_column_to_int(df["year"]) |
| if y_int.isna().all(): |
| return df |
|
|
| work = df.copy() |
| work["_yi"] = y_int |
| work = work.loc[~work["_yi"].isna()].copy() |
| work["_yi"] = work["_yi"].astype(int) |
|
|
| gcols = [c for c in work.columns if c not in ("year", "value", "_yi")] |
| dup_check = work.groupby(gcols + ["_yi"], dropna=False).size() |
| if (dup_check > 1).any(): |
| return df |
|
|
| out_rows: list[pd.Series] = [] |
| grouped = ( |
| work.groupby(gcols, dropna=False) |
| if gcols |
| else [(tuple(), work)] |
| ) |
| for _gkey, g in grouped: |
| lo = int(g["_yi"].min()) |
| hi = int(g["_yi"].max()) |
| if hi - lo > _YEAR_GAP_FILL_MAX_SPAN: |
| return df |
| sample_year = g["year"].iloc[0] |
| existing = set(int(x) for x in g["_yi"].tolist()) |
| for yi in range(lo, hi + 1): |
| match = g[g["_yi"] == yi] |
| if len(match) > 0: |
| out_rows.append(match.iloc[0].drop(labels=["_yi"])) |
| else: |
| proto = g.iloc[0].drop(labels=["_yi"]).to_dict() |
| proto["year"] = _clone_year_field(yi, sample_year) |
| proto["value"] = float("nan") |
| out_rows.append(pd.Series(proto)) |
|
|
| out = pd.DataFrame(out_rows) |
| out = out.reindex(columns=df.columns) |
| meta_cols = [c for c in df.columns if c not in ("year", "value")] |
| out["_sy"] = _coerce_year_column_to_int(out["year"]) |
| sort_keys = [k for k in (*meta_cols, "_sy") if k in out.columns] |
| out = out.sort_values(by=sort_keys, na_position="last").drop(columns=["_sy"]) |
| return out |
|
|
|
|
| def should_prepare_as_datetime( |
| viz_data: pd.DataFrame, chart_type: str, frequency: str | None |
| ) -> bool: |
| return get_data_preparation_action(chart_type, frequency) == "datetime" |
|
|
|
|
| |
| |
| |
|
|
|
|
| @dataclass |
| class PostProcessingRule: |
| name: str |
| applies_to_mark_types: list[str] |
| description: str |
|
|
| def should_apply(self, mark_type: str, encoding: dict, data: dict) -> bool: |
| raise NotImplementedError |
|
|
| def apply( |
| self, |
| spec: dict, |
| data_frequency: str | None = None, |
| unit_measure: str | None = None, |
| ) -> dict: |
| raise NotImplementedError |
|
|
|
|
| def _first_non_null_dataset_value(dataset: list, field: str) -> object: |
| for row in dataset: |
| if field in row: |
| v = row[field] |
| if v is not None: |
| return v |
| return None |
|
|
|
|
| |
| |
| _YEAR_ORDINAL_EPOCH_MS_THRESHOLD = 1e12 |
|
|
|
|
| def _year_ordinal_value_needs_temporal_encoding(value: object) -> bool: |
| """True when x is ordinal but values are ISO datetimes or epoch ms (Vega-Lite).""" |
| if value is None or isinstance(value, bool): |
| return False |
| if isinstance(value, str): |
| return "T" in value |
| if isinstance(value, (int, float)): |
| fv = float(value) |
| |
| return abs(fv) >= _YEAR_ORDINAL_EPOCH_MS_THRESHOLD |
| return False |
|
|
|
|
| def _extract_spec_dataset_rows(spec: dict) -> list[dict] | None: |
| """Return embedded chart rows from ``data.values`` or ``datasets[name]``.""" |
| data = spec.get("data") |
| if isinstance(data, dict) and "values" in data: |
| v = data.get("values") |
| return v if isinstance(v, list) else None |
| ds_name = data.get("name") if isinstance(data, dict) else None |
| if ds_name and isinstance(spec.get("datasets"), dict): |
| rows = spec["datasets"].get(ds_name) |
| return rows if isinstance(rows, list) else None |
| return None |
|
|
|
|
| def _single_obs_year_to_int(value: object) -> int | None: |
| """Parse one observation's year field to a calendar year, or None.""" |
| if value is None or (isinstance(value, float) and pd.isna(value)): |
| return None |
| if isinstance(value, bool): |
| return None |
| if isinstance(value, Integral): |
| return int(value) |
| if isinstance(value, float) and value == int(value): |
| return int(value) |
| if isinstance(value, str): |
| s = value.strip() |
| if len(s) >= 4 and s[:4].isdigit(): |
| return int(s[:4]) |
| ts = pd.to_datetime(s, errors="coerce") |
| if pd.notna(ts): |
| return int(ts.year) |
| return None |
| if isinstance(value, pd.Timestamp): |
| return int(value.year) |
| return None |
|
|
|
|
| class ContiguousCalendarYearDomainRule(PostProcessingRule): |
| """Ordinal/nominal ``year`` / ``time_period`` on x: full calendar year range on the scale. |
| |
| Applies to marks that commonly use a discrete year axis (bar, line, area, point, tick). |
| Does **not** apply when ``x`` is ``temporal`` (handled separately; continuous time ≠ discrete domain). |
| """ |
|
|
| def __init__(self): |
| super().__init__( |
| "contiguous_calendar_year_domain", |
| ["bar", "line", "area", "point", "tick"], |
| "Ordinal/nominal year x: explicit scale.domain for every calendar year in min–max span", |
| ) |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if not frequency_allows_annual_year_gap_fill(data_frequency): |
| return spec |
| if "encoding" not in spec or "x" not in spec["encoding"]: |
| return spec |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| if mark_type not in self.applies_to_mark_types: |
| return spec |
| x = spec["encoding"]["x"] |
| xf = x.get("field") |
| if xf not in ("year", "time_period"): |
| return spec |
| if x.get("type") not in ("ordinal", "nominal"): |
| return spec |
| if x.get("scale", {}).get("domain") is not None: |
| return spec |
| rows = _extract_spec_dataset_rows(spec) |
| if not rows or len(rows) < 2: |
| return spec |
| years: list[int] = [] |
| template: object | None = None |
| for row in rows: |
| if not isinstance(row, dict) or xf not in row: |
| continue |
| raw = row.get(xf) |
| if raw is None: |
| continue |
| if template is None: |
| template = raw |
| yi = _single_obs_year_to_int(raw) |
| if yi is not None: |
| years.append(yi) |
| if len(years) < 2: |
| return spec |
| lo, hi = min(years), max(years) |
| if hi - lo > _YEAR_GAP_FILL_MAX_SPAN: |
| return spec |
| if template is None: |
| return spec |
| domain = [_clone_year_field(y, template) for y in range(lo, hi + 1)] |
| x.setdefault("scale", {})["domain"] = domain |
| |
| x["sort"] = domain |
| return spec |
|
|
|
|
| class OrdinalToTemporalRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "ordinal_to_temporal", |
| ["line", "area", "point", "tick", "bar"], |
| "Fix ordinal→temporal for time fields (ISO or epoch ms), including bars", |
| ) |
|
|
| def should_apply(self, mark_type, x_enc, dataset): |
| if mark_type not in self.applies_to_mark_types: |
| return False |
| if x_enc.get("type") != "ordinal": |
| return False |
| x_field = x_enc.get("field") |
| if x_field not in ["year", "time_period"]: |
| return False |
| if not dataset: |
| return False |
| sample = _first_non_null_dataset_value(dataset, x_field) |
| return _year_ordinal_value_needs_temporal_encoding(sample) |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| if "encoding" not in spec or "x" not in spec["encoding"]: |
| return spec |
| x_enc = spec["encoding"]["x"] |
| ds_name = spec.get("data", {}).get("name") |
| if not ds_name or "datasets" not in spec: |
| return spec |
| dataset = spec["datasets"].get(ds_name, []) |
| if self.should_apply(mark_type, x_enc, dataset): |
| x_enc["type"] = "temporal" |
| return spec |
|
|
|
|
| class ApplyTimeUnitRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "apply_timeunit", |
| ["line", "area", "point", "tick"], |
| "Add timeUnit from frequency", |
| ) |
|
|
| def should_apply(self, x_enc, freq): |
| return freq in FREQUENCY_TO_TIMEUNIT and x_enc.get("type") == "temporal" |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if "encoding" not in spec or "x" not in spec["encoding"]: |
| return spec |
| x_enc = spec["encoding"]["x"] |
| if self.should_apply(x_enc, data_frequency): |
| x_enc["timeUnit"] = FREQUENCY_TO_TIMEUNIT[data_frequency] |
| return spec |
|
|
|
|
| class FixValueAxisEncodingRule(PostProcessingRule): |
| """Altair can infer ordinal for `value` after Draco strips types; fix for line/area/point.""" |
|
|
| def __init__(self): |
| super().__init__( |
| "fix_value_axis_encodings", |
| ["point", "line", "area"], |
| "Fix ordinal y on value for line/area/point; point-only size cleanup", |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None): |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| return mark_type in self.applies_to_mark_types |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if not self.should_apply(spec, data_frequency): |
| return spec |
| if "encoding" not in spec: |
| return spec |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| y = spec["encoding"].get("y", {}) |
| if y.get("type") == "ordinal" and y.get("field") == "value": |
| y["type"] = "quantitative" |
| y.setdefault("scale", {})["type"] = "linear" |
| if mark_type == "point": |
| sz = spec["encoding"].get("size", {}) |
| if sz.get("aggregate") == "count" and "field" not in sz: |
| del spec["encoding"]["size"] |
| return spec |
|
|
|
|
| class TemporalAxisCleanupRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "temporal_axis_cleanup", |
| ["line", "area", "point", "bar"], |
| "Remove title from temporal x-axis", |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None): |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| if mark_type not in self.applies_to_mark_types: |
| return False |
| return spec.get("encoding", {}).get("x", {}).get("type") == "temporal" |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if not self.should_apply(spec): |
| return spec |
| x = spec["encoding"]["x"] |
| x.setdefault("axis", {}) |
| x["axis"]["title"] = None |
| x["axis"]["labelAngle"] = 0 |
| x["axis"].setdefault("format", "%Y") |
| x["axis"].setdefault("tickCount", 5) |
| return spec |
|
|
|
|
| class DiscreteYearBarXAxisRule(PostProcessingRule): |
| """Vega-Lite defaults often rotate discrete x labels on bars; force horizontal years.""" |
|
|
| def __init__(self): |
| super().__init__( |
| "discrete_year_bar_x_axis", |
| ["bar"], |
| "Horizontal labels for ordinal/nominal year on column/bar x-axis", |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None): |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| if mark_type not in self.applies_to_mark_types: |
| return False |
| x = spec.get("encoding", {}).get("x", {}) |
| if x.get("field") not in ("year", "time_period"): |
| return False |
| return x.get("type") in ("ordinal", "nominal") |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if not self.should_apply(spec): |
| return spec |
| x = spec["encoding"]["x"] |
| x.setdefault("axis", {}) |
| x["axis"]["labelAngle"] = 0 |
| return spec |
|
|
|
|
| class ValueAxisLabelFormatRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "value_axis_label_format", |
| ["bar", "line", "area", "point", "tick"], |
| "Apply compact/value-aware y-axis label formatting", |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None): |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| if mark_type not in self.applies_to_mark_types: |
| return False |
| y = spec.get("encoding", {}).get("y", {}) |
| return y.get("type") == "quantitative" and y.get("field") in { |
| "value", |
| "obs_value", |
| } |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, scale_type=None, df=None, **kwargs): |
| if not self.should_apply(spec, data_frequency): |
| return spec |
|
|
| resolved_scale_type = scale_type |
| if df is not None and "value" in df.columns: |
| import pandas as pd |
| vals = df["value"].dropna() |
| if not vals.empty: |
| max_val = float(vals.abs().max()) |
| if max_val <= 1.0 and _is_proportion_indicator(df, unit_measure, scale_type): |
| resolved_scale_type = "proportion" |
|
|
| y = spec["encoding"]["y"] |
| y.setdefault("axis", {}) |
| y["axis"]["labelExpr"] = _value_label_expr(unit_measure, scale_type=resolved_scale_type) |
| return spec |
|
|
|
|
| |
| _LINE_GAP_SEG_DETAIL = "_d360_lseg" |
| _LINE_GAP_STROKE_FLAG = "_d360_ygap" |
|
|
|
|
| class LineYearGapStrokeDashRule(PostProcessingRule): |
| """Temporal / discrete-year line charts: dashed stroke across multi-year gaps. |
| |
| Vega-Lite draws one continuous polyline per color series. We split each |
| consecutive observation pair into its own ``detail`` group and use |
| ``strokeDash`` so segments that skip one or more calendar years render dashed. |
| """ |
|
|
| def __init__(self): |
| super().__init__( |
| "line_year_gap_stroke_dash", |
| ["line"], |
| "Dashed line segments where consecutive points differ by >1 calendar year", |
| ) |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, is_composite=False, **kwargs): |
| if not isinstance(spec, dict): |
| return spec |
| if "layer" in spec and isinstance(spec["layer"], list): |
| self._apply_to_layer_root(spec, is_composite=is_composite) |
| return spec |
| if "spec" in spec and isinstance(spec.get("spec"), dict): |
| inner = spec["spec"] |
| |
| if "layer" in inner and isinstance(inner["layer"], list): |
| self._apply_to_layer_root(inner, data_root=spec, is_composite=is_composite) |
| return spec |
| if self._is_candidate_line_spec(inner): |
| self._maybe_transform_line_spec(inner, spec, is_composite=is_composite) |
| return spec |
| if self._is_candidate_line_spec(spec): |
| self._maybe_transform_line_spec(spec, spec, is_composite=is_composite) |
| return spec |
|
|
| def _apply_to_layer_root( |
| self, layer_parent: dict, data_root: dict | None = None, is_composite: bool = False |
| ) -> None: |
| root = data_root if data_root is not None else layer_parent |
| line_layers = [ |
| layer |
| for layer in layer_parent["layer"] |
| if isinstance(layer, dict) and self._is_candidate_line_spec(layer) |
| ] |
| if len(line_layers) != 1: |
| return |
| self._maybe_transform_line_spec(line_layers[0], root, is_composite=is_composite) |
|
|
|
|
| def _mark_type(self, enc_spec: dict) -> str | None: |
| m = enc_spec.get("mark") |
| if isinstance(m, str): |
| return m |
| if isinstance(m, dict): |
| return m.get("type") |
| return None |
|
|
| def _is_candidate_line_spec(self, enc_spec: dict) -> bool: |
| mt = self._mark_type(enc_spec) |
| enc = enc_spec.get("encoding") |
| if mt != "line": |
| return False |
| if not isinstance(enc, dict): |
| return False |
| x = enc.get("x", {}) |
| if not isinstance(x, dict): |
| return False |
| xf = x.get("field") |
| if xf not in ("year", "time_period"): |
| return False |
| |
| |
| |
| if not self._x_timeunit_allows_year_gap_segments(x): |
| return False |
| if x.get("type") not in ("temporal", "ordinal", "nominal"): |
| return False |
| y = enc.get("y", {}) |
| if not isinstance(y, dict) or y.get("type") != "quantitative": |
| return False |
| if not y.get("field"): |
| return False |
| if enc.get("detail") is not None: |
| return False |
| if enc.get("strokeDash") is not None: |
| return False |
| return True |
|
|
| @staticmethod |
| def _x_timeunit_allows_year_gap_segments(x: dict) -> bool: |
| tu = x.get("timeUnit") |
| if tu is None: |
| return True |
| if isinstance(tu, str): |
| return tu in ("year", "utcyear") |
| if isinstance(tu, dict): |
| return tu.get("unit") in ("year", "utcyear") |
| return False |
|
|
| def _find_inline_values_holder(self, line_spec: dict, data_root: dict) -> dict | None: |
| for candidate in (line_spec, data_root): |
| data = candidate.get("data") |
| if isinstance(data, dict) and isinstance(data.get("values"), list): |
| return candidate |
| return None |
|
|
| def _write_inline_values(self, holder: dict, rows: list[dict]) -> None: |
| data = holder.get("data") |
| if isinstance(data, dict) and "values" in data: |
| data["values"] = rows |
|
|
| def _named_dataset_rows( |
| self, line_spec: dict, data_root: dict |
| ) -> tuple[str, list] | None: |
| for candidate in (line_spec, data_root): |
| data = candidate.get("data") |
| if not isinstance(data, dict): |
| continue |
| name = data.get("name") |
| if ( |
| isinstance(name, str) |
| and isinstance(data_root.get("datasets"), dict) |
| and isinstance(data_root["datasets"].get(name), list) |
| ): |
| return name, data_root["datasets"][name] |
| return None |
|
|
| def _set_named_dataset_rows(self, data_root: dict, name: str, rows: list[dict]) -> None: |
| data_root.setdefault("datasets", {})[name] = rows |
|
|
| def _series_keys(self, encoding: dict) -> list[str]: |
| c = encoding.get("color") |
| if isinstance(c, dict) and isinstance(c.get("field"), str): |
| return [c["field"]] |
| return [] |
|
|
| def _facet_field_keys(self, data_root: dict) -> list[str]: |
| """Facet / row / column fields so multi-panel specs split series per panel.""" |
| keys: list[str] = [] |
| for name in ("facet", "row", "column"): |
| node = data_root.get(name) |
| if not isinstance(node, dict): |
| continue |
| f = node.get("field") |
| if isinstance(f, str): |
| keys.append(f) |
| return keys |
|
|
| def _dedupe_sort_group( |
| self, rows: list[dict], x_field: str |
| ) -> list[tuple[int, dict]]: |
| by_year: dict[int, dict] = {} |
| order: list[int] = [] |
| for row in rows: |
| if not isinstance(row, dict): |
| continue |
| yi = _single_obs_year_to_int(row.get(x_field)) |
| if yi is None: |
| continue |
| if yi not in by_year: |
| order.append(yi) |
| by_year[yi] = dict(row) |
| order.sort() |
| return [(y, by_year[y]) for y in order] |
|
|
| def _group_rows_by_series( |
| self, rows: list[dict], series_keys: list[str] |
| ) -> dict[tuple, list[dict]]: |
| groups: defaultdict[tuple, list[dict]] = defaultdict(list) |
| for row in rows: |
| if not isinstance(row, dict): |
| continue |
| key = tuple(row.get(k) for k in series_keys) if series_keys else (None,) |
| groups[key].append(dict(row)) |
| return dict(groups) |
|
|
| def _any_calendar_year_gap( |
| self, groups: dict[tuple, list[dict]], x_field: str |
| ) -> bool: |
| for grp_rows in groups.values(): |
| chain = self._dedupe_sort_group(grp_rows, x_field) |
| if len(chain) < 2: |
| continue |
| for i in range(len(chain) - 1): |
| y0 = chain[i][0] |
| y1 = chain[i + 1][0] |
| if y1 - y0 > 1: |
| return True |
| return False |
|
|
| def _build_segment_rows( |
| self, groups: dict[tuple, list[dict]], x_field: str, series_keys: list[str] |
| ) -> list[dict]: |
| out: list[dict] = [] |
| seg_i = 0 |
| for key, grp_rows in groups.items(): |
| chain = self._dedupe_sort_group(grp_rows, x_field) |
| n = len(chain) |
| if n < 2: |
| |
| prefix = "_".join("" if v is None else str(v) for v in key) |
| sid = f"{prefix}_{seg_i}" if prefix else str(seg_i) |
| seg_i += 1 |
| for _, r in chain: |
| out.append({**r, _LINE_GAP_SEG_DETAIL: sid, _LINE_GAP_STROKE_FLAG: 0}) |
| continue |
|
|
| prefix = "_".join("" if v is None else str(v) for v in key) |
| i = 0 |
| while i < n: |
| |
| block = [chain[i]] |
| j = i + 1 |
| while j < n and (chain[j][0] - chain[j - 1][0]) == 1: |
| block.append(chain[j]) |
| j += 1 |
|
|
| |
| |
| if len(block) >= 2: |
| sid = f"{prefix}_{seg_i}" if prefix else str(seg_i) |
| seg_i += 1 |
| for _, r in block: |
| out.append({**r, _LINE_GAP_SEG_DETAIL: sid, _LINE_GAP_STROKE_FLAG: 0}) |
|
|
| |
| if j < n: |
| gap_sid = f"{prefix}_{seg_i}" if prefix else str(seg_i) |
| seg_i += 1 |
| _, r0 = block[-1] |
| _, r1 = chain[j] |
| out.append({**r0, _LINE_GAP_SEG_DETAIL: gap_sid, _LINE_GAP_STROKE_FLAG: 1}) |
| out.append({**r1, _LINE_GAP_SEG_DETAIL: gap_sid, _LINE_GAP_STROKE_FLAG: 1}) |
|
|
| i = j |
| return out |
|
|
| def _strip_internal_tooltip_channels(self, encoding: dict) -> None: |
| tips = encoding.get("tooltip") |
| if not isinstance(tips, list): |
| return |
| internal = {_LINE_GAP_SEG_DETAIL, _LINE_GAP_STROKE_FLAG} |
| encoding["tooltip"] = [ |
| t |
| for t in tips |
| if not (isinstance(t, dict) and t.get("field") in internal) |
| ] |
|
|
| def _maybe_transform_line_spec(self, line_spec: dict, data_root: dict, is_composite: bool = False) -> None: |
| |
| |
| |
| if is_composite or "vconcat" in data_root or "hconcat" in data_root or "concat" in data_root: |
| holder = self._find_inline_values_holder(line_spec, data_root) |
| rows: list[dict] | None = None |
| if holder is not None: |
| v = holder["data"]["values"] |
| rows = v if isinstance(v, list) else None |
| else: |
| named = self._named_dataset_rows(line_spec, data_root) |
| if named is not None: |
| _, rows_list = named |
| rows = rows_list if isinstance(rows_list, list) else None |
| if not rows or len(rows) < 2: |
| return |
|
|
| |
| |
| filters = [] |
| for candidate in (data_root, line_spec): |
| if "transform" in candidate and isinstance(candidate["transform"], list): |
| for t in candidate["transform"]: |
| if isinstance(t, dict) and "filter" in t: |
| filters.append(t["filter"]) |
|
|
| panel_rows = list(rows) |
| for filt in filters: |
| if isinstance(filt, dict): |
| field = filt.get("field") |
| equal_val = filt.get("equal") |
| one_of_val = filt.get("oneOf") |
| if field: |
| if equal_val is not None: |
| panel_rows = [r for r in panel_rows if r.get(field) == equal_val] |
| elif isinstance(one_of_val, list): |
| panel_rows = [r for r in panel_rows if r.get(field) in one_of_val] |
|
|
| enc = line_spec.get("encoding") |
| if not isinstance(enc, dict): |
| return |
| x = enc.get("x", {}) |
| x_field = x.get("field") if isinstance(x, dict) else None |
| if x_field not in ("year", "time_period"): |
| return |
|
|
| facet_keys = self._facet_field_keys(data_root) |
| series_keys = list( |
| dict.fromkeys([*self._series_keys(enc), *facet_keys]), |
| ) |
| groups = self._group_rows_by_series(panel_rows, series_keys) |
| if not self._any_calendar_year_gap(groups, x_field): |
| return |
|
|
| new_rows = self._build_segment_rows(groups, x_field, series_keys) |
| if len(new_rows) < 2: |
| return |
|
|
| |
| line_spec["data"] = {"values": new_rows} |
|
|
| enc["detail"] = {"field": _LINE_GAP_SEG_DETAIL, "type": "nominal"} |
| enc["strokeDash"] = { |
| "condition": { |
| "test": f"datum.{_LINE_GAP_STROKE_FLAG} == 1", |
| "value": [6, 4], |
| }, |
| "value": [], |
| } |
| self._strip_internal_tooltip_channels(enc) |
| return |
|
|
| enc = line_spec.get("encoding") |
| if not isinstance(enc, dict): |
| return |
| x = enc.get("x", {}) |
| x_field = x.get("field") if isinstance(x, dict) else None |
| if x_field not in ("year", "time_period"): |
| return |
|
|
|
|
|
|
| holder = self._find_inline_values_holder(line_spec, data_root) |
| rows: list[dict] | None = None |
| if holder is not None: |
| v = holder["data"]["values"] |
| rows = v if isinstance(v, list) else None |
| else: |
| named = self._named_dataset_rows(line_spec, data_root) |
| if named is not None: |
| _, rows_list = named |
| rows = rows_list if isinstance(rows_list, list) else None |
| if not rows or len(rows) < 2: |
| return |
|
|
| facet_keys = self._facet_field_keys(data_root) |
| series_keys = list( |
| dict.fromkeys([*self._series_keys(enc), *facet_keys]), |
| ) |
| groups = self._group_rows_by_series(rows, series_keys) |
| if not self._any_calendar_year_gap(groups, x_field): |
| return |
|
|
| new_rows = self._build_segment_rows(groups, x_field, series_keys) |
| if len(new_rows) < 2: |
| return |
|
|
| if holder is not None: |
| self._write_inline_values(holder, new_rows) |
| else: |
| named = self._named_dataset_rows(line_spec, data_root) |
| if named is None: |
| return |
| name, _ = named |
| self._set_named_dataset_rows(data_root, name, new_rows) |
|
|
| enc["detail"] = {"field": _LINE_GAP_SEG_DETAIL, "type": "nominal"} |
| enc["strokeDash"] = { |
| "condition": { |
| "test": f"datum.{_LINE_GAP_STROKE_FLAG} == 1", |
| "value": [6, 4], |
| }, |
| "value": [], |
| } |
| self._strip_internal_tooltip_channels(enc) |
|
|
|
|
| class LineChartPointHoverRule(PostProcessingRule): |
| """Add / enlarge line points so tooltips are easier to trigger (thin line geometry).""" |
|
|
| def __init__(self): |
| super().__init__( |
| "line_chart_point_hover", |
| ["line"], |
| "Widen line tooltip hit target with point marks", |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None, unit_measure=None): |
| m = spec.get("mark") |
| if isinstance(m, str): |
| return m == "line" |
| if isinstance(m, dict): |
| return m.get("type") == "line" |
| return False |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if not self.should_apply(spec): |
| return spec |
| m = spec["mark"] |
| if isinstance(m, str): |
| spec["mark"] = { |
| "type": "line", |
| "point": _LINE_HOVER_POINT, |
| } |
| return spec |
| pt = m.get("point") |
| if pt is False or pt is None or pt is True: |
| m["point"] = dict(_LINE_HOVER_POINT) |
| elif isinstance(pt, dict): |
| sz = pt.get("size", 0) |
| if not isinstance(sz, (int, float)) or sz < 40: |
| m["point"] = {**pt, **_LINE_HOVER_POINT} |
| return spec |
|
|
|
|
| class ZeroLineRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "zero_line", ["bar", "line", "area"], "Bar charts start at zero" |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None): |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| return mark_type in self.applies_to_mark_types |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None): |
| if not self.should_apply(spec): |
| return spec |
| y = spec.get("encoding", {}).get("y", {}) |
| if y.get("type") == "quantitative": |
| y.setdefault("scale", {}) |
| mark_type = ( |
| spec.get("mark", {}).get("type") |
| if isinstance(spec.get("mark"), dict) |
| else spec.get("mark") |
| ) |
| if mark_type == "bar": |
| y["scale"]["zero"] = True |
| return spec |
|
|
|
|
| class SkewnessLogScaleRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "skewness_log_scale", |
| ["line", "bar", "point", "area", "circle", "text"], |
| "Automatically apply log scale if values are highly positive skewed and positive" |
| ) |
|
|
| def should_apply(self, spec, data_frequency=None, df=None, raw_hint=None, **kwargs): |
| return True |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, df=None, raw_hint=None, **kwargs): |
| def _apply_log_scale_to_spec(subspec): |
| if not isinstance(subspec, dict): |
| return subspec |
|
|
| |
| if "spec" in subspec: |
| _apply_log_scale_to_spec(subspec["spec"]) |
| if "layer" in subspec and isinstance(subspec["layer"], list): |
| for child in subspec["layer"]: |
| _apply_log_scale_to_spec(child) |
| if "vconcat" in subspec and isinstance(subspec["vconcat"], list): |
| for child in subspec["vconcat"]: |
| _apply_log_scale_to_spec(child) |
| if "hconcat" in subspec and isinstance(subspec["hconcat"], list): |
| for child in subspec["hconcat"]: |
| _apply_log_scale_to_spec(child) |
|
|
| |
| enc = subspec.get("encoding", {}) |
| if not isinstance(enc, dict): |
| return subspec |
|
|
| applied_log = False |
| for channel in ("x", "y"): |
| ch_enc = enc.get(channel) |
| if not isinstance(ch_enc, dict) or ch_enc.get("type") != "quantitative": |
| continue |
| if "stack" in ch_enc: |
| continue |
|
|
| field_name = ch_enc.get("field") |
| if not field_name or df is None or field_name not in df.columns: |
| continue |
|
|
| |
| hint_str = (raw_hint or "").lower().strip() |
| is_explicit = "log" in hint_str or "logarithmic" in hint_str |
|
|
| |
| val_series = pd.to_numeric(df[field_name], errors="coerce").dropna() |
| if val_series.empty or val_series.min() <= 0: |
| continue |
|
|
| should_log = False |
| if is_explicit: |
| should_log = True |
| else: |
| |
| is_percentage = False |
| unit_str = str(unit_measure or "").upper().strip() |
| if any(x in unit_str for x in ["%", "PERCENT", "PROP", "SHARE", "RATIO", "INDEX"]): |
| is_percentage = True |
|
|
| if not is_percentage and df is not None: |
| if "unit_measure" in df.columns: |
| units = df["unit_measure"].dropna().astype(str).str.upper().unique() |
| if any(any(x in u for x in ["%", "PERCENT", "PROP", "SHARE", "RATIO", "INDEX"]) for u in units): |
| is_percentage = True |
| if not is_percentage and val_series.min() >= 0 and val_series.max() <= 1.0: |
| is_percentage = True |
|
|
| if not is_percentage: |
| skewness = val_series.skew() |
| if not pd.isna(skewness): |
| median = val_series.median() |
| val_max = val_series.max() |
| val_min = val_series.min() |
| ratio = val_max / (median or 1) |
| min_max_ratio = val_max / (val_min or 1) |
| |
| if (skewness > 0.5 and ratio > 7.5) or min_max_ratio > 50.0: |
| should_log = True |
|
|
| if should_log: |
| ch_enc.setdefault("scale", {}) |
| ch_enc["scale"]["type"] = "log" |
| ch_enc["scale"]["zero"] = False |
| if not val_series.empty and val_series.min() > 0: |
| min_val = float(val_series.min()) |
| domain_min = 1.0 if min_val > 1.0 else float(min_val * 0.9) |
| ch_enc["scale"]["domain"] = [domain_min, float(val_series.max() * 1.1)] |
| applied_log = True |
|
|
| if applied_log: |
| |
| |
| |
| |
| mark = subspec.get("mark") |
| is_bar = (isinstance(mark, dict) and mark.get("type") == "bar") or mark == "bar" |
| if is_bar and val_series is not None and not val_series.empty and val_series.min() > 0: |
| min_val = float(val_series.min()) |
| domain_min = 1.0 if min_val > 1.0 else float(min_val * 0.9) |
| enc = subspec.setdefault("encoding", {}) |
| |
| |
| |
| if "x" in enc and enc["x"].get("field") == "value": |
| enc["x2"] = {"datum": domain_min} |
| elif "y" in enc and enc["y"].get("field") == "value": |
| enc["y2"] = {"datum": domain_min} |
|
|
|
|
|
|
| if applied_log: |
| has_log[0] = True |
|
|
| return subspec |
|
|
| has_log = [False] |
| _apply_log_scale_to_spec(spec) |
| if has_log[0]: |
| title_obj = spec.get("title") |
| if isinstance(title_obj, dict): |
| subtitle = title_obj.get("subtitle", []) |
| if isinstance(subtitle, str): |
| subtitle = [subtitle] |
| elif not isinstance(subtitle, list): |
| subtitle = [] |
|
|
| note_str = "Note: Value axis is on a logarithmic scale due to wide dynamic range." |
| if note_str not in subtitle: |
| subtitle.append(note_str) |
| title_obj["subtitle"] = subtitle |
| elif isinstance(title_obj, str): |
| spec["title"] = { |
| "text": title_obj, |
| "subtitle": ["Note: Value axis is on a logarithmic scale due to wide dynamic range."] |
| } |
| return spec |
|
|
|
|
| class PercentageBoundaryClampingRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "percentage_boundary_clamping", |
| ["line", "bar", "area"], |
| "Clamps percentage axes to 0-100 or 0-1 standard bounds" |
| ) |
|
|
| def should_apply(self, spec, scale_type=None, df=None, unit_measure=None, **kwargs): |
| import pandas as pd |
| if df is not None and "value" in df.columns: |
| vals = df["value"].dropna() |
| if not vals.empty: |
| max_val = float(vals.max()) |
| min_val = float(vals.min()) |
| if min_val < 0: |
| return False |
| if max_val > 100: |
| return False |
|
|
| |
| normalized_unit = (unit_measure or "").upper() |
| unbounded_keywords = {"GROWTH", "INFLATION", "INTEREST", "YIELD", "INDEX"} |
| if any(k in normalized_unit for k in unbounded_keywords): |
| return False |
| if "indicator" in df.columns: |
| ind_names = df["indicator"].dropna().unique() |
| for name in ind_names: |
| name_upper = str(name).upper() |
| if any(k in name_upper for k in unbounded_keywords): |
| return False |
|
|
| |
| is_proportion = max_val <= 1.0 and _is_proportion_indicator(df, unit_measure, scale_type) |
| if is_proportion and min_val > 0.3: |
| return False |
| if not is_proportion and min_val > 30: |
| return False |
|
|
| if scale_type == "percentage": |
| return True |
|
|
| |
| if df is not None and "value" in df.columns: |
| import pandas as pd |
| vals = df["value"].dropna() |
| if not vals.empty: |
| max_val = float(vals.abs().max()) |
| if max_val <= 1.0 and _is_proportion_indicator(df, unit_measure, scale_type): |
| return True |
| return False |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, scale_type=None, df=None, **kwargs): |
| if not self.should_apply(spec, scale_type=scale_type, df=df, unit_measure=unit_measure): |
| return spec |
|
|
| |
| is_proportion = False |
| max_val = 100.0 |
| if df is not None and "value" in df.columns: |
| import pandas as pd |
| vals = df["value"].dropna() |
| if not vals.empty: |
| max_abs = float(vals.abs().max()) |
| max_val = float(vals.max()) |
| if max_abs <= 1.0 and _is_proportion_indicator(df, unit_measure, scale_type): |
| is_proportion = True |
|
|
| |
| if is_proportion: |
| if max_val <= 0.05: |
| upper = 0.05 |
| elif max_val <= 0.10: |
| upper = 0.10 |
| elif max_val <= 0.25: |
| upper = 0.25 |
| elif max_val <= 0.50: |
| upper = 0.50 |
| else: |
| upper = 1.0 |
| domain = [0, upper] |
| else: |
| if max_val <= 5: |
| upper = 5 |
| elif max_val <= 10: |
| upper = 10 |
| elif max_val <= 25: |
| upper = 25 |
| elif max_val <= 50: |
| upper = 50 |
| else: |
| upper = 100 |
| domain = [0, upper] |
|
|
| enc = spec.get("encoding", {}) |
| for channel in ("x", "y"): |
| ch_enc = enc.get(channel) |
| if isinstance(ch_enc, dict) and ch_enc.get("type") == "quantitative": |
| ch_enc.setdefault("scale", {}) |
| if ch_enc["scale"].get("type") == "log": |
| continue |
| ch_enc["scale"]["domain"] = domain |
|
|
| |
| if is_proportion: |
| axis_spec = ch_enc.setdefault("axis", {}) |
| if "labelExpr" in axis_spec: |
| expr = axis_spec["labelExpr"] |
| if "datum.value" in expr and "* 100" not in expr and "%" in expr: |
| axis_spec["labelExpr"] = expr.replace("datum.value", "datum.value * 100") |
|
|
| if "spec" in spec: |
| self.apply(spec["spec"], data_frequency, unit_measure, scale_type=scale_type, df=df, **kwargs) |
|
|
| if "layer" in spec: |
| for subspec in spec["layer"]: |
| self.apply(subspec, data_frequency, unit_measure, scale_type=scale_type, df=df, **kwargs) |
|
|
| return spec |
|
|
|
|
| def _filter_df_for_error_band(df: pd.DataFrame) -> pd.DataFrame: |
| """Pre-filters the dataframe to keep only the best errorband triplet/pair if present.""" |
| if "comp_breakdown_1" not in df.columns or "value" not in df.columns: |
| return df |
|
|
| cb1_vals = df["comp_breakdown_1"].dropna().unique().tolist() |
|
|
| |
| lower_keywords = {"lower", "lb", "min", "minimum", "low"} |
| upper_keywords = {"upper", "ub", "max", "maximum", "high"} |
| se_keywords = {"se", "std_err", "stderr", "error"} |
| est_keywords = {"estimate", "score", "value", "val", "est"} |
|
|
| classified = {} |
| for v in cb1_vals: |
| v_lower = str(v).lower().strip() |
| parts = set(v_lower.replace("-", "_").replace(" ", "_").split("_")) |
|
|
| if parts.intersection(lower_keywords): |
| classified[v] = "lower" |
| elif parts.intersection(upper_keywords): |
| classified[v] = "upper" |
| elif parts.intersection(se_keywords) or "standard error" in v_lower: |
| classified[v] = "se" |
| elif parts.intersection(est_keywords): |
| classified[v] = "est" |
|
|
| |
| lower_vals = [k for k, val in classified.items() if val == "lower"] |
| upper_vals = [k for k, val in classified.items() if val == "upper"] |
| est_vals = [k for k, val in classified.items() if val == "est"] |
| se_vals = [k for k, val in classified.items() if val == "se"] |
|
|
| best_triplet = None |
| for est in est_vals: |
| |
| est_clean = re.sub(r'\(.*?\)', '', str(est)) |
| est_base = est_clean.lower().replace("_sc", "").replace("_est", "").strip() |
|
|
| matching_lower = [] |
| for l in lower_vals: |
| l_clean = re.sub(r'\(.*?\)', '', str(l)).lower() |
| if est_base in l_clean or est_base == l_clean.replace("_lb", "").replace("_lower", "").strip(): |
| matching_lower.append(l) |
|
|
| matching_upper = [] |
| for u in upper_vals: |
| u_clean = re.sub(r'\(.*?\)', '', str(u)).lower() |
| if est_base in u_clean or est_base == u_clean.replace("_ub", "").replace("_upper", "").strip(): |
| matching_upper.append(u) |
|
|
| if matching_lower and matching_upper: |
| best_triplet = (est, matching_lower[0], matching_upper[0]) |
| break |
|
|
| if best_triplet: |
| return df[df["comp_breakdown_1"].isin(best_triplet)].copy() |
|
|
| |
| best_pair = None |
| for est in est_vals: |
| est_clean = re.sub(r'\(.*?\)', '', str(est)) |
| est_base = est_clean.lower().replace("_est", "").strip() |
|
|
| matching_se = [] |
| for s in se_vals: |
| s_clean = re.sub(r'\(.*?\)', '', str(s)).lower() |
| |
| is_generic = s_clean.strip() in ("se", "standard error", "std_err", "stderr", "std error", "error") |
| if est_base in s_clean or is_generic: |
| matching_se.append(s) |
|
|
| if matching_se: |
| best_pair = (est, matching_se[0]) |
| break |
|
|
| if best_pair: |
| return df[df["comp_breakdown_1"].isin(best_pair)].copy() |
|
|
| |
| if est_vals and lower_vals and upper_vals: |
| return df[df["comp_breakdown_1"].isin([est_vals[0], lower_vals[0], upper_vals[0]])].copy() |
| if est_vals and se_vals: |
| return df[df["comp_breakdown_1"].isin([est_vals[0], se_vals[0]])].copy() |
|
|
| return df |
|
|
|
|
| class GeneralErrorBandRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "general_error_band", |
| ["line"], |
| "Layer estimate lines with lower/upper bound area error bands when confidence intervals or standard errors are detected" |
| ) |
|
|
| def _classify_values(self, unique_vals: list[str]) -> tuple[dict[str, str], set[str]]: |
| """Maps each unique category label to a normalized key: 'est', 'lower', 'upper', or 'se'.""" |
| mapping = {} |
| unclassified = [] |
|
|
| lower_keywords = {"lower", "lb", "min", "minimum", "low"} |
| upper_keywords = {"upper", "ub", "max", "maximum", "high"} |
| se_keywords = {"se", "std_err", "stderr", "error"} |
| est_keywords = {"estimate", "score", "value", "val", "est"} |
|
|
| for v in unique_vals: |
| v_lower = str(v).lower().strip() |
| parts = set(v_lower.replace("-", "_").replace(" ", "_").split("_")) |
|
|
| |
| if parts.intersection(lower_keywords): |
| mapping[v] = "lower" |
| |
| elif parts.intersection(upper_keywords): |
| mapping[v] = "upper" |
| |
| elif parts.intersection(se_keywords) or "standard error" in v_lower: |
| mapping[v] = "se" |
| |
| elif parts.intersection(est_keywords): |
| mapping[v] = "est" |
| else: |
| unclassified.append(v) |
|
|
| |
| classified_keys = set(mapping.values()) |
| if len(unclassified) == 1: |
| if "lower" in classified_keys and "upper" in classified_keys and "est" not in classified_keys: |
| mapping[unclassified[0]] = "est" |
| elif "se" in classified_keys and "est" not in classified_keys: |
| mapping[unclassified[0]] = "est" |
|
|
| return mapping, set(mapping.values()) |
|
|
| def should_apply(self, spec, df=None, **kwargs): |
| if df is None or "comp_breakdown_1" not in df.columns or "value" not in df.columns: |
| return False |
| cb1_vals = df["comp_breakdown_1"].dropna().unique().tolist() |
| _, mapped_keys = self._classify_values(cb1_vals) |
| has_triplet = {"est", "lower", "upper"}.issubset(mapped_keys) |
| has_pair = {"est", "se"}.issubset(mapped_keys) |
| return has_triplet or has_pair |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, df=None, **kwargs): |
| if not self.should_apply(spec, df=df): |
| return spec |
|
|
| df_copy = df.copy() |
| cb1_vals = df_copy["comp_breakdown_1"].dropna().unique().tolist() |
| mapping, _ = self._classify_values(cb1_vals) |
|
|
| |
| df_copy["comp_breakdown_1"] = df_copy["comp_breakdown_1"].map(mapping).fillna(df_copy["comp_breakdown_1"]) |
| unique_vals = set(df_copy["comp_breakdown_1"].dropna().unique()) |
|
|
| groupby_cols = [c for c in ["year", "time_period", "country", "ref_area", "ref_area_name"] if c in df_copy.columns] |
|
|
| if {"est", "lower", "upper"}.issubset(unique_vals): |
| transforms = [ |
| { |
| "pivot": "comp_breakdown_1", |
| "value": "value", |
| "groupby": groupby_cols |
| } |
| ] |
| lb_field = "lower" |
| ub_field = "upper" |
| elif {"est", "se"}.issubset(unique_vals): |
| transforms = [ |
| { |
| "pivot": "comp_breakdown_1", |
| "value": "value", |
| "groupby": groupby_cols |
| }, |
| { |
| "calculate": "datum.est - 1.645 * datum.se", |
| "as": "lower" |
| }, |
| { |
| "calculate": "datum.est + 1.645 * datum.se", |
| "as": "upper" |
| } |
| ] |
| lb_field = "lower" |
| ub_field = "upper" |
| else: |
| return spec |
|
|
| is_layered = "layer" in spec |
| is_vconcat = ("vconcat" in spec or "concat" in spec) and not is_layered |
| concat_key = "concat" if "concat" in spec else "vconcat" |
|
|
| |
| |
| |
| |
| if is_vconcat: |
| new_panels = [] |
| for panel in spec[concat_key]: |
| panel_enc = panel.get("encoding", {}) |
| if not panel_enc or "y" not in panel_enc: |
| new_panels.append(panel) |
| continue |
|
|
| color_enc = panel_enc.get("color") |
| keep_color = False |
| if isinstance(color_enc, dict): |
| color_field = color_enc.get("field") |
| if color_field and color_field != "comp_breakdown_1": |
| keep_color = True |
|
|
| |
| |
| |
| country_filter_transforms = [ |
| t for t in panel.get("transform", []) |
| if "filter" in t and isinstance(t["filter"], dict) and t["filter"].get("field") == "country" |
| ] |
| indicator_filter_transforms = [ |
| t for t in panel.get("transform", []) |
| if "filter" in t and isinstance(t["filter"], dict) and t["filter"].get("field") == "indicator" |
| ] |
| panel_x_enc = panel_enc.get("x", {}) |
| panel_y_enc = panel_enc.get("y", {}) |
| panel_mark = panel.get("mark", {"type": "line", "strokeWidth": 3}) |
| panel_title = panel.get("title") |
| panel_width = panel.get("width", spec.get("width", 600)) |
| panel_height = panel.get("height", spec.get("height", 200)) |
|
|
| |
| |
| |
| |
| scoping_filters = country_filter_transforms + indicator_filter_transforms |
| full_transforms = scoping_filters + transforms |
|
|
| is_nominal_x = panel_x_enc.get("type") == "nominal" |
| is_nominal_y = panel_y_enc.get("type") == "nominal" |
|
|
| if is_nominal_x or is_nominal_y: |
| errorband_mark = {"type": "rule", "color": "#666666", "strokeWidth": 1.5} |
| else: |
| errorband_mark = {"type": "area", "opacity": 0.2} |
| if not keep_color: |
| errorband_mark["color"] = "#34A7F2" |
|
|
| if is_nominal_y: |
| errorband_encoding = { |
| "y": panel_y_enc, |
| "x": { |
| "field": lb_field, |
| "type": "quantitative", |
| "scale": {"zero": False} |
| }, |
| "x2": { |
| "field": ub_field |
| } |
| } |
| if "yOffset" in panel_enc: |
| errorband_encoding["yOffset"] = panel_enc["yOffset"] |
| else: |
| errorband_encoding = { |
| "x": panel_x_enc, |
| "y": { |
| "field": lb_field, |
| "type": "quantitative", |
| "scale": {"zero": False} |
| }, |
| "y2": { |
| "field": ub_field |
| } |
| } |
| if "xOffset" in panel_enc: |
| errorband_encoding["xOffset"] = panel_enc["xOffset"] |
|
|
| if keep_color: |
| errorband_encoding["color"] = color_enc |
|
|
| errorband = { |
| "transform": full_transforms, |
| "mark": errorband_mark, |
| "encoding": errorband_encoding |
| } |
|
|
| |
| if keep_color: |
| line_enc = panel_enc |
| else: |
| line_enc = {k: v for k, v in panel_enc.items() if k != "color"} |
|
|
| line = { |
| "transform": scoping_filters + [ |
| {"filter": "datum.comp_breakdown_1 == 'est'"} |
| ], |
| "mark": panel_mark, |
| "encoding": line_enc |
| } |
|
|
| new_panel = { |
| "title": panel_title, |
| "width": panel_width, |
| "height": panel_height, |
| "layer": [errorband, line] |
| } |
| new_panels.append(new_panel) |
|
|
| spec[concat_key] = new_panels |
| spec["data"] = {"values": df_copy.to_dict(orient="records")} |
| return spec |
|
|
|
|
| if is_layered: |
| |
| main_layer = None |
| for layer in spec["layer"]: |
| if "encoding" in layer and "y" in layer["encoding"]: |
| main_layer = layer |
| break |
| if main_layer is None: |
| return spec |
| enc = main_layer["encoding"] |
| mark_spec = main_layer.get("mark", "line") |
| else: |
| enc = spec.get("encoding", {}) |
| mark_spec = spec.get("mark", "line") |
|
|
| if not enc: |
| return spec |
|
|
| x_enc = enc.get("x", {}) |
| y_enc = enc.get("y", {}) |
| color_enc = enc.get("color", {}) |
|
|
| keep_color = False |
| if isinstance(color_enc, dict): |
| color_field = color_enc.get("field") |
| if color_field and color_field != "comp_breakdown_1": |
| keep_color = True |
|
|
| is_nominal_x = x_enc.get("type") == "nominal" |
| is_nominal_y = y_enc.get("type") == "nominal" |
|
|
| if is_nominal_x or is_nominal_y: |
| errorband_mark = {"type": "rule", "color": "#666666", "strokeWidth": 1.5} |
| else: |
| errorband_mark = {"type": "area", "opacity": 0.2} |
| if not keep_color: |
| errorband_mark["color"] = "#34A7F2" |
|
|
| if is_nominal_y: |
| errorband_encoding = { |
| "y": y_enc, |
| "x": { |
| "field": lb_field, |
| "type": "quantitative", |
| "scale": {"zero": False} |
| }, |
| "x2": { |
| "field": ub_field |
| } |
| } |
| if "yOffset" in enc: |
| errorband_encoding["yOffset"] = enc["yOffset"] |
| else: |
| errorband_encoding = { |
| "x": x_enc, |
| "y": { |
| "field": lb_field, |
| "type": "quantitative", |
| "scale": {"zero": False} |
| }, |
| "y2": { |
| "field": ub_field |
| } |
| } |
| if "xOffset" in enc: |
| errorband_encoding["xOffset"] = enc["xOffset"] |
|
|
| if keep_color: |
| errorband_encoding["color"] = color_enc |
|
|
| errorband_layer = { |
| "transform": transforms, |
| "mark": errorband_mark, |
| "encoding": errorband_encoding |
| } |
|
|
| |
| line_enc = enc.copy() |
| if not keep_color: |
| if "color" in line_enc: |
| del line_enc["color"] |
|
|
| line_layer = { |
| "transform": [ |
| { |
| "filter": "datum.comp_breakdown_1 == 'est'" |
| } |
| ], |
| "mark": mark_spec, |
| "encoding": line_enc |
| } |
|
|
| if is_layered: |
| new_layers = [errorband_layer] |
| for layer in spec["layer"]: |
| |
| is_text = False |
| mark = layer.get("mark", {}) |
| mark_type = mark.get("type") if isinstance(mark, dict) else mark |
| if mark_type == "text": |
| is_text = True |
|
|
| if layer is main_layer: |
| new_layers.append(line_layer) |
| elif is_text: |
| continue |
| else: |
| new_layers.append(layer) |
| spec["layer"] = new_layers |
| spec["data"] = {"values": df_copy.to_dict(orient="records")} |
| return spec |
| else: |
| layered_spec = { |
| "$schema": spec.get("$schema", "https://vega.github.io/schema/vega-lite/v5.json"), |
| "title": spec.get("title"), |
| "width": spec.get("width", 600), |
| "height": spec.get("height", 350), |
| "config": spec.get("config", {}), |
| "data": {"values": df_copy.to_dict(orient="records")}, |
| "layer": [errorband_layer, line_layer] |
| } |
| return layered_spec |
|
|
|
|
| class PopulationPyramidRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__( |
| "population_pyramid", |
| ["*"], |
| "Transform bar chart into a diverging population pyramid when age and sex dimensions are present" |
| ) |
|
|
| def should_apply(self, spec, df=None, raw_hint=None, **kwargs): |
| if df is None: |
| return False |
| cols = {c.lower() for c in df.columns} |
| has_dims = "sex" in cols and "age" in cols |
| if not has_dims: |
| return False |
|
|
| hint_str = (raw_hint or "").lower().strip() |
| is_pyramid_hint = "pyramid" in hint_str or "population_pyramid" in hint_str |
| if is_pyramid_hint: |
| return True |
|
|
| |
| year_count = df["year"].nunique() if "year" in df.columns else 0 |
| country_count = df["country"].nunique() if "country" in df.columns else 0 |
| if year_count == 1 and country_count == 1: |
| return True |
|
|
| return False |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, df=None, raw_hint=None, **kwargs): |
| if not self.should_apply(spec, df=df, raw_hint=raw_hint): |
| return spec |
|
|
| sex_col = [c for c in df.columns if c.lower() == "sex"][0] |
| age_col = [c for c in df.columns if c.lower() == "age"][0] |
|
|
| unique_ages = df[age_col].dropna().unique().tolist() |
| sort_order = _get_dimension_sort_order(age_col, unique_ages) |
|
|
| y_enc = { |
| "field": age_col, |
| "type": "nominal", |
| "axis": { |
| "title": "Age Group", |
| "grid": False |
| } |
| } |
| if sort_order: |
| y_enc["sort"] = sort_order |
|
|
| unique_sexes = df[sex_col].dropna().unique().tolist() |
| domain = [s for s in ["Male", "Female", "M", "F"] if s in unique_sexes] |
| range_colors = [] |
| for s in domain: |
| if s in ("Male", "M"): |
| range_colors.append("#34A7F2") |
| else: |
| range_colors.append("#F3578E") |
|
|
| pyramid_spec = { |
| "$schema": spec.get("$schema", "https://vega.github.io/schema/vega-lite/v5.json"), |
| "title": spec.get("title"), |
| "width": spec.get("width", 600), |
| "height": spec.get("height", 350), |
| "config": spec.get("config", {}), |
| "data": spec.get("data", {}), |
| "transform": [ |
| { |
| "calculate": f"datum.{sex_col} == 'Male' || datum.{sex_col} == 'M' ? -datum.value : datum.value", |
| "as": "signed_value" |
| } |
| ], |
| "mark": { |
| "type": "bar", |
| "tooltip": True |
| }, |
| "encoding": { |
| "y": y_enc, |
| "x": { |
| "field": "signed_value", |
| "type": "quantitative", |
| "axis": { |
| "title": "Population", |
| "labelExpr": "abs(datum.value)" |
| } |
| }, |
| "color": { |
| "field": sex_col, |
| "type": "nominal", |
| "scale": { |
| "domain": domain, |
| "range": range_colors |
| }, |
| "legend": { |
| "title": "Sex" |
| } |
| } |
| } |
| } |
| return pyramid_spec |
|
|
|
|
| class ApplyWBStyleRule(PostProcessingRule): |
| def __init__(self): |
| super().__init__("apply_wb_style", ["*"], "Inject WB style config") |
|
|
| def should_apply(self, spec, data_frequency=None): |
| return True |
|
|
| def apply(self, spec, data_frequency=None, unit_measure=None, **kwargs): |
| return inject_wb_config(spec) |
|
|
|
|
| POST_PROCESSING_RULES: list[PostProcessingRule] = [ |
| OrdinalToTemporalRule(), |
| ApplyTimeUnitRule(), |
| FixValueAxisEncodingRule(), |
| ContiguousCalendarYearDomainRule(), |
| TemporalAxisCleanupRule(), |
| DiscreteYearBarXAxisRule(), |
| SkewnessLogScaleRule(), |
| PercentageBoundaryClampingRule(), |
| ValueAxisLabelFormatRule(), |
| LineYearGapStrokeDashRule(), |
| LineChartPointHoverRule(), |
| GeneralErrorBandRule(), |
| PopulationPyramidRule(), |
| ZeroLineRule(), |
| ApplyWBStyleRule(), |
| ] |
|
|
|
|
| |
| |
| |
|
|
|
|
| @dataclass |
| class DracoConstraintConfig: |
| base_constraints: list[str] |
| nominal_color_fields: list[str] |
| color_dimension_priority: list[str] |
|
|
| def __init__(self): |
| self.base_constraints = ["entity(view,root,view).", "entity(mark,view,m)."] |
| self.nominal_color_fields = ["country", "sex", "urbanisation", "residence", "ref_area"] |
| self.color_dimension_priority = ["country", "sex", "age", "urbanisation", "residence"] |
|
|
|
|
| DEFAULT_DRACO_CONFIG = DracoConstraintConfig() |
|
|