Skip to content
use serde::Deserialize;

use super::types::ChartType;

#[derive(Deserialize)]
pub struct DbQuery {
    pub database: Option<String>,
}

#[derive(Deserialize, Debug)]
pub struct ChartQuery {
    pub database: String,
    pub table: String,
    pub chart_type: Option<ChartType>,
    pub x_col: Option<String>,
    pub y_col: Option<String>,
    #[serde(default = "default_limit")]
    pub limit: usize,
}

pub fn default_limit() -> usize {
    500
}

/// Backtick-quote a ClickHouse identifier (database / table / column) to defeat
/// reserved-keyword collisions and basic SQL-injection vectors. Inner backticks
/// are doubled per ClickHouse's identifier escape rules.
pub fn q(ident: &str) -> String {
    format!("`{}`", ident.replace('`', "``"))
}

/// Cap user-provided limits so we don't accidentally pull millions of rows
/// for a chart.
fn clamp_limit(limit: usize) -> usize {
    limit.clamp(1, 100_000)
}

/// Build the SQL for the chosen chart. Identifiers are backtick-quoted; the
/// limit is range-clamped before interpolation.
fn build_query(
    database: &str,
    table: &str,
    chart_type: &ChartType,
    x_col: &str,
    y_col: &str,
    limit: usize,
) -> String {
    let db = q(database);
    let tbl = q(table);
    let x = q(x_col);
    let y = q(y_col);
    let limit = clamp_limit(limit);

    match chart_type {
        ChartType::Line => format!(
            "SELECT toString(toDate({x})) AS x_label, \
             toFloat64(sum({y})) AS y_value \
             FROM {db}.{tbl} \
             WHERE {x} IS NOT NULL \
             GROUP BY x_label ORDER BY x_label \
             LIMIT {limit} FORMAT JSONEachRow"
        ),
        ChartType::Bar => format!(
            "SELECT toString({x}) AS x_label, \
             toFloat64(sum({y})) AS y_value \
             FROM {db}.{tbl} \
             WHERE {x} IS NOT NULL \
             GROUP BY x_label ORDER BY y_value DESC \
             LIMIT {limit} FORMAT JSONEachRow"
        ),
        ChartType::Scatter => format!(
            "SELECT toString(toFloat64OrNull({x})) AS x_label, \
             toFloat64OrZero({y}) AS y_value \
             FROM {db}.{tbl} \
             WHERE {x} IS NOT NULL AND {y} IS NOT NULL \
             LIMIT {limit} FORMAT JSONEachRow"
        ),
        ChartType::Histogram => format!(
            "SELECT toString(toFloat64OrNull({x})) AS x_label, \
             0.0 AS y_value \
             FROM {db}.{tbl} \
             WHERE {x} IS NOT NULL \
             LIMIT {limit} FORMAT JSONEachRow"
        ),
    }
}

/// Returns `(x_label, y_value)` rows from ClickHouse.
pub async fn fetch_chart_data(
    conn: &crate::connections::Connection,
    database: &str,
    table: &str,
    chart_type: &ChartType,
    x_col: &str,
    y_col: &str,
    limit: usize,
) -> Result<Vec<(String, f64)>, String> {
    let query = build_query(database, table, chart_type, x_col, y_col, limit);

    let client = reqwest::Client::new();
    let mut req = client
        .post(conn.url())
        .query(&[("user", conn.user.as_str())])
        .body(query);
    if !conn.password.is_empty() {
        req = req.query(&[("password", conn.password.as_str())]);
    }
    let resp = req.send().await.map_err(|e| e.to_string())?;
    if !resp.status().is_success() {
        return Err(resp.text().await.unwrap_or_default());
    }
    let text = resp.text().await.map_err(|e| e.to_string())?;

    let rows = text
        .lines()
        .filter(|l| !l.trim().is_empty())
        .filter_map(|l| {
            let v: serde_json::Value = serde_json::from_str(l).ok()?;
            let x_raw = v["x_label"]
                .as_str()
                .map(ToString::to_string)
                .or_else(|| v["x_label"].as_f64().map(|f| f.to_string()))?;
            // Filter ClickHouse / serde-json renderings of NULL.
            let trimmed = x_raw.trim();
            if trimmed.is_empty() || trimmed == "\\N" || trimmed.eq_ignore_ascii_case("null") {
                return None;
            }
            let y = v["y_value"]
                .as_f64()
                .or_else(|| v["y_value"].as_str().and_then(|s| s.parse::<f64>().ok()))
                .unwrap_or(0.0);
            Some((x_raw, y))
        })
        .collect();
    Ok(rows)
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn quote_escapes_backticks() {
        assert_eq!(q("foo"), "`foo`");
        assert_eq!(q("we`ird"), "`we``ird`");
    }

    #[test]
    fn build_query_quotes_all_identifiers() {
        let sql = build_query("my db", "tbl`name", &ChartType::Bar, "col x", "col y", 10);
        assert!(sql.contains("`my db`.`tbl``name`"));
        assert!(sql.contains("`col x`"));
        assert!(sql.contains("`col y`"));
        assert!(sql.contains("LIMIT 10"));
    }

    #[test]
    fn limit_is_clamped() {
        let sql = build_query("d", "t", &ChartType::Histogram, "x", "", 10_000_000);
        assert!(sql.contains("LIMIT 100000"));
    }
}