from .charts_axis_band_axis import band_axis from .charts_axis_grid_lines import grid_lines from .charts_axis_number_axis import number_axis from .charts_bar_chart_types import BarChart, BarChartSpec, ChartBar from .charts_format_format_tick import format_tick from .charts_layout_legend_rows import legend_rows from .charts_layout_plot_area import plot_area from .charts_layout_types import LegendItem, Margins from .charts_palette import palette from .charts_scale_band_scale import band_scale from .charts_scale_linear_scale import linear_scale from .charts_scale_nice_domain import nice_domain from .charts_shape_bar_rects import bar_rects from .charts_shape_types import BarSpec from .charts_stack import stack from .charts_ticks_nice_ticks import nice_ticks from .charts_ticks_tick_step import tick_step from .math_round_float import round_float # Fixed layout constants; the README explains each. CHAR_WIDTH = 7 SWATCH = 10 LEGEND_GAP = 16 LEGEND_ROW = 18 EDGE = 8 TICK_SIZE = 6 LABEL_PADDING = 3 TOP_WITHOUT_LEGEND = 12 RIGHT = 20 BOTTOM = 30 CATEGORY_PADDING_INNER = 0.2 CATEGORY_PADDING_OUTER = 0.1 SERIES_PADDING = 0.05 def bar_chart(spec: BarChartSpec) -> BarChart: """A whole bar chart as geometry, by the fixed rules in the README.""" width, height, categories, series = spec.width, spec.height, spec.categories, spec.series if spec.mode not in ("grouped", "stacked"): raise ValueError("mode must be grouped or stacked, received %s" % (spec.mode,)) if len(categories) == 0: raise ValueError("a bar chart needs at least one category") if len(series) == 0: raise ValueError("a bar chart needs at least one series") for s in series: if len(s.values) != len(categories): raise ValueError( 'series "%s" has %d values but there are %d categories' % (s.name, len(s.values), len(categories)) ) names = [s.name for s in series] legend = [] top = TOP_WITHOUT_LEGEND if len(series) > 1: rows = legend_rows(names, width - 2 * EDGE, CHAR_WIDTH, SWATCH, LEGEND_GAP, LEGEND_ROW) legend = [ LegendItem( label=it.label, row=it.row, x=round_float(it.x + EDGE, 2), y=round_float(it.y + EDGE, 2), width=it.width, text_x=round_float(it.text_x + EDGE, 2), ) for it in rows ] top = EDGE + (rows[-1].row + 1) * LEGEND_ROW + EDGE stacked = stack([list(s.values) for s in series], "diverging") if spec.mode == "stacked" else None lo = 0.0 hi = 0.0 for i, s in enumerate(series): for j, v in enumerate(s.values): ends = [v] if stacked is None else [stacked[i][j].y0, stacked[i][j].y1] for e in ends: if e < lo: lo = e if e > hi: hi = e if lo == hi: hi = 1.0 y_count = max(2, int((height - top - BOTTOM) // 50)) y_domain = nice_domain([lo, hi], y_count).domain y_step = tick_step(y_domain[0], y_domain[1], y_count) widest = 0 for v in nice_ticks(y_domain[0], y_domain[1], y_count): widest = max(widest, len(format_tick(v, y_step))) left = TICK_SIZE + LABEL_PADDING + widest * CHAR_WIDTH + EDGE plot = plot_area(width, height, Margins(top=top, right=RIGHT, bottom=BOTTOM, left=left)) bottom_y = plot.y + plot.height y_range = [bottom_y, plot.y] def y(v: float) -> float: return linear_scale(y_domain, y_range, v, False) y_axis = number_axis(y_domain, y_range, y_count, "left", plot.x, TICK_SIZE) x_range = [plot.x, plot.x + plot.width] x_axis = band_axis( categories, x_range, CATEGORY_PADDING_INNER, CATEGORY_PADDING_OUTER, "bottom", bottom_y, TICK_SIZE ) colors = palette("okabe-ito", len(series), True) specs = [] meta = [] for i, s in enumerate(series): for j, c in enumerate(categories): outer = band_scale(categories, x_range, c, CATEGORY_PADDING_INNER, CATEGORY_PADDING_OUTER, 0.5) if stacked is not None: specs.append( BarSpec(band=outer.start, thickness=outer.width, base=y(stacked[i][j].y0), value=y(stacked[i][j].y1)) ) elif len(series) == 1: specs.append(BarSpec(band=outer.start, thickness=outer.width, base=y(0), value=y(s.values[j]))) else: inner = band_scale(names, [outer.start, outer.start + outer.width], s.name, SERIES_PADDING, 0, 0.5) specs.append(BarSpec(band=inner.start, thickness=inner.width, base=y(0), value=y(s.values[j]))) meta.append((s.name, c, s.values[j], colors[i])) rects = bar_rects(specs, "vertical") bars = [ ChartBar(series=m[0], category=m[1], value=m[2], color=m[3], rect=rects[k]) for k, m in enumerate(meta) ] return BarChart( width=round_float(width, 2), height=round_float(height, 2), plot=plot, x_axis=x_axis, y_axis=y_axis, grid_lines=grid_lines(y_axis, plot.width), bars=bars, colors=list(colors), legend=legend, )