Files
Lakehouse/config/superset/apply_palantir_rich_dashboard.py
T
Lakehouse Admin 0402a4fbae
Validate Lakehouse config / validate (push) Has been cancelled
update (bart)
2026-06-28 17:08:57 +02:00

454 lines
15 KiB
Python

#!/usr/bin/env python3
"""Apply Palantir styling and enrich ATC Lakehouse Superset dashboard."""
import json
import os
import requests
BASE = "http://127.0.0.1:8088"
DASH_ID = 1
CSS_PATH = "/tmp/palantir_dashboard.css"
NEW_CHARTS = [
(
"Trino · PostgreSQL Sales",
"public",
"sales_orders",
"PostgreSQL · Orders by Channel",
"pie",
{
"metric": {"expressionType": "SQL", "sqlExpression": "COUNT(*)", "label": "COUNT(*)"},
"groupby": ["sales_channel"],
"row_limit": 10,
},
),
(
"Trino · PostgreSQL Sales",
"public",
"sales_orders",
"PostgreSQL · Avg Order by Region",
"echarts_timeseries_bar",
{
"metrics": [
{"expressionType": "SQL", "sqlExpression": "AVG(amount)", "label": "Avg Amount"}
],
"groupby": ["region"],
"row_limit": 15,
},
),
(
"Trino · PostgreSQL Sales",
"public",
"sales_orders",
"PostgreSQL · Status Breakdown",
"pie",
{
"metric": {"expressionType": "SQL", "sqlExpression": "COUNT(*)", "label": "COUNT(*)"},
"groupby": ["order_status"],
"row_limit": 10,
},
),
(
"Trino · MySQL HR",
"hr",
"employee_events",
"MySQL HR · By Event Type",
"pie",
{
"metric": {"expressionType": "SQL", "sqlExpression": "COUNT(*)", "label": "COUNT(*)"},
"groupby": ["event_type"],
"row_limit": 12,
},
),
(
"Trino · MySQL HR",
"hr",
"employee_events",
"MySQL HR · Events per Month",
"echarts_timeseries_line",
{
"metrics": [
{"expressionType": "SQL", "sqlExpression": "COUNT(*)", "label": "Events"}
],
"groupby": [
{
"expressionType": "SQL",
"sqlExpression": "date_trunc('month', event_ts)",
"label": "Month",
}
],
"row_limit": 24,
},
),
(
"Trino · MongoDB Supply Chain",
"supplychain",
"events",
"MongoDB · Amount by Source",
"echarts_timeseries_bar",
{
"metrics": [
{"expressionType": "SQL", "sqlExpression": "SUM(amount)", "label": "Total Amount"}
],
"groupby": ["source"],
"row_limit": 10,
},
),
(
"Trino · MongoDB Supply Chain",
"supplychain",
"events",
"MongoDB · Events per Month",
"echarts_timeseries_line",
{
"metrics": [
{"expressionType": "SQL", "sqlExpression": "COUNT(*)", "label": "Events"}
],
"groupby": [
{
"expressionType": "SQL",
"sqlExpression": "date_trunc('month', ts)",
"label": "Month",
}
],
"row_limit": 24,
},
),
(
"Trino · Cassandra Telemetry",
"telemetry",
"device_metrics",
"Cassandra · Avg Metric Over Time",
"echarts_timeseries_line",
{
"metrics": [
{
"expressionType": "SQL",
"sqlExpression": "AVG(metric_value)",
"label": "Avg Value",
}
],
"groupby": [
{
"expressionType": "SQL",
"sqlExpression": "date_trunc('day', metric_ts)",
"label": "Day",
}
],
"row_limit": 30,
},
),
]
SECTIONS = [
("HEADER", "ATC Lakehouse · Federated Data Platform", "Trino · PostgreSQL · MySQL · MongoDB · Cassandra · Dell Technologies FDE"),
("PostgreSQL Sales", "30M orders · postgres_sales.public.sales_orders"),
("MySQL HR", "569K events · mysql_hr.hr.employee_events"),
("MongoDB Supply Chain", "3M events · mongodb_supplychain.supplychain.events"),
("Cassandra Telemetry", "Device metrics · cassandra_telemetry.telemetry.device_metrics"),
]
def session():
s = requests.Session()
r = s.post(
f"{BASE}/api/v1/security/login",
json={"username": "admin", "password": "admin", "provider": "db", "refresh": True},
)
r.raise_for_status()
h = {
"Authorization": "Bearer " + r.json()["access_token"],
"Content-Type": "application/json",
}
h["X-CSRFToken"] = s.get(f"{BASE}/api/v1/security/csrf_token/", headers=h).json()["result"]
h["Referer"] = BASE
return s, h
def get_db_map(s, h):
r = s.get(f"{BASE}/api/v1/database/", headers=h)
r.raise_for_status()
return {d["database_name"]: d["id"] for d in r.json().get("result", [])}
def get_or_create_dataset(s, h, db_id, schema, table):
r = s.get(f"{BASE}/api/v1/dataset/", headers=h)
for d in r.json().get("result", []):
if (
d.get("table_name") == table
and d.get("schema") == schema
and d.get("database", {}).get("id") == db_id
):
return d["id"]
r = s.post(
f"{BASE}/api/v1/dataset/",
headers=h,
json={"database": db_id, "schema": schema, "table_name": table},
)
r.raise_for_status()
return r.json()["id"]
def create_chart(s, h, name, ds_id, viz_type, params):
r = s.get(f"{BASE}/api/v1/chart/", headers=h)
for c in r.json().get("result", []):
if c.get("slice_name") == name:
return c["id"]
full = {"datasource": f"{ds_id}__table", "viz_type": viz_type, "row_limit": 1000, **params}
r = s.post(
f"{BASE}/api/v1/chart/",
headers=h,
json={
"slice_name": name,
"viz_type": viz_type,
"datasource_id": ds_id,
"datasource_type": "table",
"params": json.dumps(full),
"owners": [1, 2, 3],
},
)
if r.status_code not in (200, 201):
raise RuntimeError(f"chart {name}: {r.text[:300]}")
return r.json()["id"]
def build_layout(chart_items):
"""chart_items: list of (chart_id, name, section) or ('md', title, subtitle)."""
layout = {
"DASHBOARD_VERSION": "v2",
"ROOT_ID": {"type": "ROOT", "id": "ROOT_ID", "children": ["GRID_ID"]},
"GRID_ID": {"type": "GRID", "id": "GRID_ID", "children": [], "parents": ["ROOT_ID"]},
}
row_idx = 0
def add_row():
nonlocal row_idx
row_idx += 1
rid = f"ROW-{row_idx}"
layout["GRID_ID"]["children"].append(rid)
layout[rid] = {
"type": "ROW",
"id": rid,
"children": [],
"parents": ["ROOT_ID", "GRID_ID"],
"meta": {"background": "BACKGROUND_TRANSPARENT"},
}
return rid
for item in chart_items:
if item[0] == "md":
_, title, subtitle = item
rid = add_row()
mid = f"MARKDOWN-{row_idx}"
layout[rid]["children"].append(mid)
layout[mid] = {
"type": "MARKDOWN",
"id": mid,
"children": [],
"parents": ["ROOT_ID", "GRID_ID", rid],
"meta": {
"width": 12,
"height": 12,
"code": f"## {title}\n\n{subtitle}",
},
}
else:
cid, name, _section = item
rid = add_row()
# up to 3 charts per row
existing = [
k
for k in layout[rid]["children"]
if k.startswith("CHART-")
]
if len(existing) >= 3:
rid = add_row()
chart_key = f"CHART-explore-{cid}"
layout[rid]["children"].append(chart_key)
col = len([k for k in layout[rid]["children"] if k.startswith("CHART-")]) - 1
layout[chart_key] = {
"type": "CHART",
"id": chart_key,
"children": [],
"parents": ["ROOT_ID", "GRID_ID", rid],
"meta": {
"width": 4,
"height": 55 if "Total" in name or "Records" in name else 65,
"chartId": cid,
"sliceName": name,
},
}
return layout
def save_query_contexts():
app = __import__("superset.app", fromlist=["create_app"]).create_app()
with app.app_context():
from flask import g
from superset.extensions import db
from superset.models.slice import Slice
from superset.charts.schemas import ChartDataQueryContextSchema
from superset import security_manager
g.user = security_manager.find_user(username="admin")
for sl in db.session.query(Slice).all():
try:
fd = sl.form_data
metric = fd.get("metric")
metrics = fd.get("metrics") or ([metric] if metric else [])
if not metrics:
metrics = [
{
"expressionType": "SQL",
"sqlExpression": "COUNT(*)",
"label": "COUNT(*)",
}
]
groupby = fd.get("groupby") or []
payload = {
"datasource": {"id": sl.datasource_id, "type": sl.datasource_type},
"force": False,
"queries": [
{
"filters": [],
"extras": {"having": "", "where": ""},
"applied_time_extras": {},
"columns": groupby if isinstance(groupby, list) else [],
"metrics": metrics,
"orderby": [],
"annotation_layers": [],
"row_limit": int(fd.get("row_limit") or 1000),
"series_limit": 0,
"order_desc": True,
"url_params": {},
"custom_params": {},
"custom_form_data": {},
}
],
"form_data": fd,
"result_format": "json",
"result_type": "full",
}
ChartDataQueryContextSchema().load(payload)
sl.query_context = json.dumps(payload)
sl.query_context_generation = True
db.session.add(sl)
except Exception as e:
print("qc err", sl.id, e)
db.session.commit()
def main():
s, h = session()
db_map = get_db_map(s, h)
# Create new charts
new_ids = []
for db_name, schema, table, name, viz, params in NEW_CHARTS:
db_id = db_map.get(db_name)
if not db_id:
print("skip, no db:", db_name)
continue
ds_id = get_or_create_dataset(s, h, db_id, schema, table)
cid = create_chart(s, h, name, ds_id, viz, params)
new_ids.append((cid, name, db_name.split("·")[-1].strip()))
print("new chart", cid, name)
# All charts for dashboard
r = s.get(f"{BASE}/api/v1/chart/?q=(page:0,page_size:200)", headers=h)
all_charts = r.json().get("result", [])
def sort_key(c):
n = c.get("slice_name") or ""
if "Lakehouse" in n or "Records per Source" in n:
return (0, n)
if "PostgreSQL" in n:
return (1, n)
if "MySQL" in n:
return (2, n)
if "MongoDB" in n:
return (3, n)
if "Cassandra" in n:
return (4, n)
return (5, n)
all_charts.sort(key=sort_key)
chart_items = [
("md", "ATC Lakehouse · Federated Data Platform", "Real-time analytics across all Trino catalogs · Dell Technologies"),
]
current_section = None
for c in all_charts:
name = c.get("slice_name") or ""
if "PostgreSQL" in name and current_section != "pg":
chart_items.append(("md", "PostgreSQL Sales", "30M orders · CDC-enabled · atc-db02"))
current_section = "pg"
elif "MySQL" in name and current_section != "mysql":
chart_items.append(("md", "MySQL HR", "569K employee events · HR domain"))
current_section = "mysql"
elif "MongoDB" in name and current_section != "mongo":
chart_items.append(("md", "MongoDB Supply Chain", "3M supply chain events"))
current_section = "mongo"
elif "Cassandra" in name and current_section != "cass":
chart_items.append(("md", "Cassandra Telemetry", "IoT device metrics"))
current_section = "cass"
chart_items.append((c["id"], name, current_section))
chart_ids = [x[0] for x in chart_items if x[0] != "md"]
position = build_layout(chart_items)
css = ""
if os.path.exists(CSS_PATH):
css = open(CSS_PATH, encoding="utf-8").read()
chart_configuration = {
str(cid): {"id": cid, "crossFilters": {"scope": "global", "chartsInScope": chart_ids}}
for cid in chart_ids
}
payload = {
"dashboard_title": "ATC Lakehouse · Trino Federated",
"published": True,
"position_json": json.dumps(position),
"css": css,
"json_metadata": json.dumps(
{
"color_scheme": "palantir_ops",
"label_colors": {},
"refresh_frequency": 120,
"timed_refresh_immune_slices": [],
"expanded_slices": {},
"chart_configuration": chart_configuration,
"global_chart_configuration": {
"scope": {"rootPath": ["ROOT_ID"], "excluded": []},
"chartsInScope": chart_ids,
},
"native_filter_configuration": [],
"color_scheme_domain": [],
"shared_label_colors": {},
}
),
"owners": [1, 2, 3],
}
r = s.put(f"{BASE}/api/v1/dashboard/{DASH_ID}", headers=h, json=payload)
print("dashboard update", r.status_code)
if r.status_code >= 400:
print(r.text[:500])
return
for cid in chart_ids:
s.put(
f"{BASE}/api/v1/chart/{cid}",
headers=h,
json={"dashboards": [DASH_ID], "owners": [1, 2, 3]},
)
print("Saving query contexts...")
save_query_contexts()
print(f"Done — {len(chart_ids)} charts, Palantir theme applied to dashboard.")
if __name__ == "__main__":
main()