Skip to content
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");
    }
}