use ratatui_core::layout::Rect;
use super::geometry::{Axis, Padding, Size};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Dimension {
#[default]
Auto,
Fixed(u16),
Percent(u16),
Flex(u16),
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Align {
Start,
Center,
End,
#[default]
Stretch,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Justify {
#[default]
Start,
Center,
End,
SpaceBetween,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Direction {
Row,
#[default]
Column,
}
impl Direction {
pub fn axis(self) -> Axis {
match self {
Direction::Row => Axis::Horizontal,
Direction::Column => Axis::Vertical,
}
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct LayoutStyle {
pub direction: Direction,
pub padding: Padding,
pub gap: u16,
pub align_items: Align,
pub justify: Justify,
}
impl LayoutStyle {
pub fn row() -> Self {
Self {
direction: Direction::Row,
..Self::default()
}
}
pub fn column() -> Self {
Self {
direction: Direction::Column,
..Self::default()
}
}
pub fn gap(mut self, gap: u16) -> Self {
self.gap = gap;
self
}
pub fn padding(mut self, padding: Padding) -> Self {
self.padding = padding;
self
}
pub fn align(mut self, align: Align) -> Self {
self.align_items = align;
self
}
pub fn justify(mut self, justify: Justify) -> Self {
self.justify = justify;
self
}
}
#[derive(Clone, Copy, Debug)]
pub struct Item {
pub dimension: Dimension,
pub intrinsic: Size,
}
impl Item {
pub fn new(dimension: Dimension, intrinsic: Size) -> Self {
Self {
dimension,
intrinsic,
}
}
}
pub fn solve(area: Rect, style: &LayoutStyle, items: &[Item]) -> Vec<Rect> {
if items.is_empty() {
return Vec::new();
}
let axis = style.direction.axis();
let inner = style.padding.inner(area);
let inner_size = Size::from(inner);
let main_avail = axis.main(inner_size);
let cross_avail = axis.cross(inner_size);
let total_gap = style
.gap
.saturating_mul(items.len().saturating_sub(1) as u16);
let space_for_children = main_avail.saturating_sub(total_gap);
let mut main_sizes: Vec<u16> = Vec::with_capacity(items.len());
let mut flex_weight_total: u32 = 0;
let mut consumed: u16 = 0;
for item in items {
let size = match item.dimension {
Dimension::Auto => axis.main(item.intrinsic).min(space_for_children),
Dimension::Fixed(n) => n.min(space_for_children),
Dimension::Percent(p) => {
let p = p.min(100) as u32;
((space_for_children as u32 * p) / 100) as u16
}
Dimension::Flex(weight) => {
flex_weight_total += weight.max(1) as u32;
0 }
};
main_sizes.push(size);
consumed = consumed.saturating_add(size);
}
let leftover = space_for_children.saturating_sub(consumed);
if flex_weight_total > 0 {
let mut distributed: u16 = 0;
let flex_indices: Vec<usize> = items
.iter()
.enumerate()
.filter(|(_, it)| matches!(it.dimension, Dimension::Flex(_)))
.map(|(i, _)| i)
.collect();
for (nth, &i) in flex_indices.iter().enumerate() {
let weight = match items[i].dimension {
Dimension::Flex(w) => w.max(1) as u32,
_ => unreachable!(),
};
let size = if nth + 1 == flex_indices.len() {
leftover.saturating_sub(distributed)
} else {
(leftover as u32 * weight)
.checked_div(flex_weight_total)
.unwrap_or(0) as u16
};
distributed = distributed.saturating_add(size);
main_sizes[i] = size;
}
}
let used_main: u16 = main_sizes.iter().copied().fold(0, u16::saturating_add);
let free = space_for_children.saturating_sub(used_main);
let (mut cursor, between_extra) = match style.justify {
Justify::Start => (0, 0),
Justify::Center => (free / 2, 0),
Justify::End => (free, 0),
Justify::SpaceBetween if items.len() > 1 => (0, free / (items.len() as u16 - 1)),
Justify::SpaceBetween => (0, 0),
};
let mut rects = Vec::with_capacity(items.len());
for (i, item) in items.iter().enumerate() {
let main_len = main_sizes[i];
let cross_len = match style.align_items {
Align::Stretch => cross_avail,
_ => axis.cross(item.intrinsic).min(cross_avail),
};
let cross_off = match style.align_items {
Align::Start | Align::Stretch => 0,
Align::Center => cross_avail.saturating_sub(cross_len) / 2,
Align::End => cross_avail.saturating_sub(cross_len),
};
let main_start = cursor.min(main_avail);
let main_len = main_len.min(main_avail.saturating_sub(main_start));
rects.push(axis.place(inner, main_start, cross_off, main_len, cross_len));
cursor = cursor
.saturating_add(main_sizes[i])
.saturating_add(style.gap)
.saturating_add(between_extra);
}
rects
}
#[cfg(test)]
mod tests {
use super::*;
use crate::geometry::{Padding, Size};
use ratatui_core::layout::Rect;
fn item(dim: Dimension, w: u16, h: u16) -> Item {
Item::new(dim, Size::new(w, h))
}
#[test]
fn flex_distributes_leftover_to_grow_children() {
let area = Rect::new(0, 0, 30, 1);
let style = LayoutStyle::row();
let items = [
item(Dimension::Fixed(10), 10, 1),
item(Dimension::Flex(1), 0, 1),
item(Dimension::Flex(1), 0, 1),
];
let rects = solve(area, &style, &items);
assert_eq!(rects[0].width, 10);
assert_eq!(rects[1].width, 10);
assert_eq!(rects[2].width, 10);
assert_eq!(rects[1].x, 10);
assert_eq!(rects[2].x, 20);
}
#[test]
fn flex_grow_weights_and_remainder_fill_exactly() {
let area = Rect::new(0, 0, 10, 1);
let style = LayoutStyle::row();
let items = [
item(Dimension::Flex(1), 0, 1),
item(Dimension::Flex(2), 0, 1),
];
let rects = solve(area, &style, &items);
assert_eq!(rects[0].width + rects[1].width, 10);
assert_eq!(rects[0].width, 3);
assert_eq!(rects[1].width, 7);
}
#[test]
fn flex_percent_and_gap() {
let area = Rect::new(0, 0, 20, 1);
let style = LayoutStyle::row().gap(2);
let items = [
item(Dimension::Percent(50), 0, 1),
item(Dimension::Auto, 4, 1),
];
let rects = solve(area, &style, &items);
assert_eq!(rects[0].width, 9);
assert_eq!(rects[1].x, rects[0].x + 9 + 2);
}
#[test]
fn column_stretch_fills_cross_axis() {
let area = Rect::new(0, 0, 12, 6);
let style = LayoutStyle::column().align(Align::Stretch);
let items = [
item(Dimension::Fixed(2), 3, 2),
item(Dimension::Fixed(2), 5, 2),
];
let rects = solve(area, &style, &items);
assert_eq!(rects[0].width, 12);
assert_eq!(rects[1].width, 12);
assert_eq!(rects[0].height, 2);
assert_eq!(rects[1].y, 2);
}
#[test]
fn justify_center_and_end_offset_main_axis() {
let area = Rect::new(0, 0, 20, 1);
let items = [item(Dimension::Fixed(4), 4, 1)];
let center = solve(area, &LayoutStyle::row().justify(Justify::Center), &items);
assert_eq!(center[0].x, 8); let end = solve(area, &LayoutStyle::row().justify(Justify::End), &items);
assert_eq!(end[0].x, 16);
}
#[test]
fn padding_shrinks_layout_area() {
let area = Rect::new(0, 0, 20, 5);
let style = LayoutStyle::column().padding(Padding::all(1));
let items = [item(Dimension::Flex(1), 0, 0)];
let rects = solve(area, &style, &items);
assert_eq!(rects[0].x, 1);
assert_eq!(rects[0].y, 1);
assert_eq!(rects[0].width, 18);
assert_eq!(rects[0].height, 3);
}
#[test]
fn flex_solver_survives_degenerate_areas() {
let items = [
item(Dimension::Flex(1), 0, 0),
item(Dimension::Fixed(5), 5, 1),
item(Dimension::Percent(50), 0, 0),
];
for (w, h) in [(0u16, 0u16), (1, 1), (2, 2), (3, 10), (4, 1), (60, 3)] {
let area = Rect::new(0, 0, w, h);
let style = LayoutStyle::row().gap(10);
let rects = solve(area, &style, &items);
assert_eq!(rects.len(), items.len());
for r in &rects {
assert!(r.right() <= area.right(), "{r:?} exceeds width of {area:?}");
assert!(
r.bottom() <= area.bottom(),
"{r:?} exceeds height of {area:?}"
);
}
}
}
}