use crate::db::ColumnInfo;
use super::types::{ChartType, ColKind};
/// Strip `Nullable(...)`, `LowCardinality(...)`, `SimpleAggregateFunction(..., T)`
/// wrappers so we get to the underlying primitive type. Case-insensitive.
fn unwrap_type(raw: &str) -> String {
let mut t = raw.trim().to_string();
loop {
let lower = t.to_lowercase();
if let Some(stripped) = strip_wrapper(&lower, &t, "nullable(") {
t = stripped;
} else if let Some(stripped) = strip_wrapper(&lower, &t, "lowcardinality(") {
t = stripped;
} else {
break;
}
}
t
}
fn strip_wrapper(lower: &str, original: &str, prefix: &str) -> Option<String> {
if lower.starts_with(prefix) && original.ends_with(')') {
// preserve original casing of the inner type
let start = prefix.len();
let end = original.len() - 1;
if end > start {
return Some(original[start..end].trim().to_string());
}
}
None
}
pub fn classify(raw: &str) -> ColKind {
let inner = unwrap_type(raw);
let lower = inner.to_lowercase();
if lower.starts_with("date") {
ColKind::DateTime
} else if lower.starts_with("uint")
|| lower.starts_with("int")
|| lower.starts_with("float")
|| lower.starts_with("decimal")
{
ColKind::Numeric
} else {
ColKind::StringLike
}
}
/// Heuristic: is this column likely an opaque identifier we should *not* pick
/// as a default axis?
fn is_idish(name: &str) -> bool {
let l = name.to_lowercase();
l == "id" || l == "uuid" || l.ends_with("_id") || l.ends_with("_uuid") || l.starts_with("uuid_")
}
/// Pick (chart_type, x_col, y_col) defaults given a table's columns.
///
/// Strategy:
/// 1. date + numeric → Line
/// 2. (non-id) string + numeric → Bar
/// 3. ≥ 2 numerics → Scatter (skip id-ish ones first)
/// 4. 1 numeric → Histogram
/// 5. fallback → Bar over first two columns
pub fn infer_chart(cols: &[ColumnInfo]) -> (ChartType, String, String) {
let date_col = cols
.iter()
.find(|c| classify(&c.column_type) == ColKind::DateTime);
// Prefer non-id-ish numerics, but fall back to id-ish ones if those are
// all we have.
let (numeric_pref, numeric_fallback): (Vec<_>, Vec<_>) = cols
.iter()
.filter(|c| classify(&c.column_type) == ColKind::Numeric)
.partition(|c| !is_idish(&c.name));
let num_cols: Vec<_> = numeric_pref.into_iter().chain(numeric_fallback).collect();
let (string_pref, string_fallback): (Vec<_>, Vec<_>) = cols
.iter()
.filter(|c| classify(&c.column_type) == ColKind::StringLike)
.partition(|c| !is_idish(&c.name));
let str_cols: Vec<_> = string_pref.into_iter().chain(string_fallback).collect();
if let (Some(d), Some(n)) = (date_col, num_cols.first()) {
return (ChartType::Line, d.name.clone(), n.name.clone());
}
if let (Some(s), Some(n)) = (str_cols.first(), num_cols.first()) {
return (ChartType::Bar, s.name.clone(), n.name.clone());
}
if num_cols.len() >= 2 {
return (
ChartType::Scatter,
num_cols[0].name.clone(),
num_cols[1].name.clone(),
);
}
if let Some(n) = num_cols.first() {
return (ChartType::Histogram, n.name.clone(), String::new());
}
let x = cols.first().map(|c| c.name.clone()).unwrap_or_default();
let y = cols.get(1).map(|c| c.name.clone()).unwrap_or_default();
(ChartType::Bar, x, y)
}
#[cfg(test)]
mod tests {
use super::*;
fn col(name: &str, ty: &str) -> ColumnInfo {
ColumnInfo {
name: name.to_string(),
column_type: ty.to_string(),
}
}
#[test]
fn classify_basic_types() {
assert_eq!(classify("Date"), ColKind::DateTime);
assert_eq!(classify("DateTime64(3)"), ColKind::DateTime);
assert_eq!(classify("UInt32"), ColKind::Numeric);
assert_eq!(classify("Float64"), ColKind::Numeric);
assert_eq!(classify("Decimal(18, 4)"), ColKind::Numeric);
assert_eq!(classify("String"), ColKind::StringLike);
assert_eq!(classify("FixedString(8)"), ColKind::StringLike);
}
#[test]
fn classify_unwraps_nullable_and_lowcardinality() {
assert_eq!(classify("Nullable(Int64)"), ColKind::Numeric);
assert_eq!(classify("LowCardinality(String)"), ColKind::StringLike);
assert_eq!(
classify("LowCardinality(Nullable(String))"),
ColKind::StringLike
);
assert_eq!(classify("Nullable(DateTime)"), ColKind::DateTime);
}
#[test]
fn infer_skips_idish_columns_for_bar() {
let cols = vec![
col("user_id", "UInt64"),
col("name", "String"),
col("amount", "Float64"),
];
let (kind, x, y) = infer_chart(&cols);
assert_eq!(kind, ChartType::Bar);
assert_eq!(x, "name");
assert_eq!(y, "amount");
}
}