diff --git a/src/db.py b/src/db.py index c4d0892..0235f6f 100644 --- a/src/db.py +++ b/src/db.py @@ -191,3 +191,30 @@ def count_unique_marches(where_sql: str = "TRUE", params: tuple | list = ()) -> logger.debug("count_unique_marches: " + sql.replace("?", "{}").format(*params)) result = get_cursor().execute(sql, list(params)).fetchone() return int(result[0]) if result else 0 + + +def aggregate_marches( + select_sql: str, + where_sql: str = "TRUE", + params: tuple | list = (), + group_by: str | None = None, + limit: int | None = None, + offset: int | None = None, +) -> pl.DataFrame: + """SELECT agrégé paramétré contre la table decp. + + `select_sql` et `group_by` sont des fragments SQL construits depuis des + noms de colonnes validés (jamais de valeur utilisateur libre). Les + valeurs de filtre passent par le binding `?` via `params`. + """ + sql = f"SELECT {select_sql} FROM decp WHERE {where_sql}" + if group_by: + sql += f" GROUP BY {group_by}" + if limit is not None: + sql += f" LIMIT {int(limit)}" + if offset is not None: + sql += f" OFFSET {int(offset)}" + + logger.debug("aggregate_marches: " + sql.replace("?", "{}").format(*params)) + + return get_cursor().execute(sql, list(params)).pl() diff --git a/tests/api/test_db_aggregate.py b/tests/api/test_db_aggregate.py new file mode 100644 index 0000000..191083b --- /dev/null +++ b/tests/api/test_db_aggregate.py @@ -0,0 +1,19 @@ +import polars as pl + +from src.db import aggregate_marches + + +def test_aggregate_groupby_count_returns_named_columns(): + df = aggregate_marches( + select_sql='"acheteur_departement_code", COUNT("uid") AS "uid__count"', + group_by='"acheteur_departement_code"', + ) + assert isinstance(df, pl.DataFrame) + assert df.columns == ["acheteur_departement_code", "uid__count"] + assert df["uid__count"].sum() > 0 + + +def test_aggregate_global_without_groupby_returns_one_row(): + df = aggregate_marches(select_sql='COUNT("uid") AS "uid__count"') + assert df.height == 1 + assert df["uid__count"][0] > 0