Files
atc-agents/api/sql_console.py
T
mo 5828113f53 feat(ui): Neo4j graph explorer with interactive relationship visualization
Add Graph sub-tab in Data Sources UI for Neo4j: force-directed SVG view of
Product-Supplier nodes and SUPPLIES/RELATED_TO/PART_OF/COMPATIBLE_WITH edges.
Backend GET /api/sql/graph/neo4j with rel_type filter and edge limit.
2026-06-27 15:48:56 +00:00

730 lines
34 KiB
Python

"""Live SQL console — PostgreSQL, MySQL, MongoDB, Cassandra, Neo4j + Trino."""
from __future__ import annotations
import os
import re
import time
from typing import Any
import httpx
import psycopg2
import pymysql
from cassandra.auth import PlainTextAuthProvider
from cassandra.cluster import Cluster
from fastapi import APIRouter, Query
from fastapi.responses import JSONResponse
from neo4j import GraphDatabase
from pydantic import BaseModel, Field
from pymongo import MongoClient
DB_HOST = os.getenv("DB_VAULT_HOST", "10.0.21.51")
PG_PORT = int(os.getenv("PG_PORT", "5432"))
PG_USER = os.getenv("PG_USER", "mo")
PG_PASS = os.getenv("PG_PASSWORD", "Dell2026!")
PG_DB = os.getenv("PG_DATABASE", "postgres")
MYSQL_PORT = int(os.getenv("MYSQL_PORT", "3306"))
MYSQL_USER = os.getenv("MYSQL_USER", "mo")
MYSQL_PASS = os.getenv("MYSQL_PASSWORD", "Dell2026!")
MYSQL_DB = os.getenv("MYSQL_DATABASE", "hr")
MONGO_HOST = os.getenv("MONGO_HOST", DB_HOST)
MONGO_PORT = int(os.getenv("MONGO_PORT", "27017"))
MONGO_DB = os.getenv("MONGO_DATABASE", "supplychain")
CASS_PORT = int(os.getenv("CASSANDRA_PORT", "9042"))
CASS_USER = os.getenv("CASSANDRA_USER", "cassandra")
CASS_PASS = os.getenv("CASSANDRA_PASSWORD", "cassandra")
CASS_KS = os.getenv("CASSANDRA_KEYSPACE", "telemetry")
NEO4J_URI = os.getenv("NEO4J_URI", f"bolt://{DB_HOST}:7687")
NEO4J_USER = os.getenv("NEO4J_USER", "neo4j")
NEO4J_PASS = os.getenv("NEO4J_PASSWORD", "testpwd")
TRINO_URL = os.getenv("TRINO_URL", "http://10.0.21.50:8089").rstrip("/")
TRINO_USER = os.getenv("TRINO_USER", "atc")
SOURCE_ENGINES = ("postgres", "mysql", "mongodb", "cassandra", "neo4j")
ENGINES = SOURCE_ENGINES + ("trino",)
router = APIRouter(prefix="/api/sql", tags=["sql"])
SAMPLES: dict[str, list[dict[str, str]]] = {
"postgres": [
{"id": "pg1", "label": "Server version", "sql": "SELECT version();"},
{"id": "pg2", "label": "Current session", "sql": "SELECT current_database(), current_user, now();"},
{"id": "pg3", "label": "Public tables", "sql": "SELECT tablename FROM pg_tables WHERE schemaname='public' ORDER BY 1 LIMIT 20;"},
{"id": "pg4", "label": "Table count", "sql": "SELECT count(*) AS public_tables FROM information_schema.tables WHERE table_schema='public';"},
{"id": "pg5", "label": "Active queries", "sql": "SELECT pid, usename, state, left(query, 100) AS query FROM pg_stat_activity WHERE state='active' AND pid <> pg_backend_pid() LIMIT 10;"},
{"id": "pg6", "label": "Database sizes", "sql": "SELECT datname, pg_size_pretty(pg_database_size(datname)) AS size FROM pg_database ORDER BY pg_database_size(datname) DESC LIMIT 10;"},
{"id": "pg7", "label": "Row counts (stats)", "sql": "SELECT relname, n_live_tup FROM pg_stat_user_tables ORDER BY n_live_tup DESC NULLS LAST LIMIT 10;"},
{"id": "pg8", "label": "Connections", "sql": "SELECT count(*) AS connections, state FROM pg_stat_activity GROUP BY state;"},
{"id": "pg9", "label": "Sales orders sample", "sql": "SELECT order_id, customer_id, region, total_amount FROM sales_orders ORDER BY order_id DESC LIMIT 10;"},
{"id": "pg10", "label": "Memory settings", "sql": "SELECT name, setting, unit FROM pg_settings WHERE name IN ('max_connections','shared_buffers','work_mem','effective_cache_size');"},
],
"mysql": [
{"id": "my1", "label": "Server version", "sql": "SELECT VERSION() AS version;"},
{"id": "my2", "label": "Current session", "sql": "SELECT DATABASE() AS db, USER() AS user, NOW() AS now_ts;"},
{"id": "my3", "label": "HR tables", "sql": "SHOW TABLES;"},
{"id": "my4", "label": "Table sizes", "sql": "SELECT table_name, table_rows, ROUND((data_length+index_length)/1024/1024,1) AS mb FROM information_schema.tables WHERE table_schema='hr' ORDER BY (data_length+index_length) DESC;"},
{"id": "my5", "label": "Employee events sample", "sql": "SELECT event_type, COUNT(*) AS cnt FROM employee_events GROUP BY event_type ORDER BY cnt DESC LIMIT 10;"},
{"id": "my6", "label": "Recent events", "sql": "SELECT employee_id, event_type, event_date FROM employee_events ORDER BY event_date DESC LIMIT 10;"},
{"id": "my7", "label": "Active connections", "sql": "SHOW STATUS LIKE 'Threads_connected';"},
{"id": "my8", "label": "InnoDB buffer pool", "sql": "SHOW STATUS LIKE 'Innodb_buffer_pool%';"},
{"id": "my9", "label": "Explain count scan", "sql": "EXPLAIN SELECT COUNT(*) FROM employee_events;"},
{"id": "my10", "label": "Events per employee", "sql": "SELECT employee_id, COUNT(*) AS events FROM employee_events GROUP BY employee_id ORDER BY events DESC LIMIT 10;"},
],
"mongodb": [
{"id": "mg1", "label": "List databases", "sql": "SHOW DATABASES"},
{"id": "mg2", "label": "Supplychain collections", "sql": "SHOW COLLECTIONS supplychain"},
{"id": "mg3", "label": "Count events", "sql": "COUNT supplychain.events"},
{"id": "mg4", "label": "Sample document", "sql": "FIND supplychain.events LIMIT 3"},
{"id": "mg5", "label": "Events by type", "sql": "AGGREGATE supplychain.events GROUP type TOP 10"},
{"id": "mg6", "label": "Events by region", "sql": "AGGREGATE supplychain.events GROUP region TOP 10"},
{"id": "mg7", "label": "Distinct event types", "sql": "DISTINCT supplychain.events type"},
{"id": "mg8", "label": "Server status", "sql": "SERVER STATUS"},
{"id": "mg9", "label": "Indexes on events", "sql": "INDEXES supplychain.events"},
{"id": "mg10", "label": "Recent events", "sql": "FIND supplychain.events SORT timestamp DESC LIMIT 5"},
],
"cassandra": [
{"id": "cs1", "label": "Cluster release", "sql": "SELECT release_version FROM system.local"},
{"id": "cs2", "label": "Keyspaces", "sql": "SELECT keyspace_name FROM system_schema.keyspaces"},
{"id": "cs3", "label": "Telemetry tables", "sql": "SELECT table_name FROM system_schema.tables WHERE keyspace_name='telemetry'"},
{"id": "cs4", "label": "Device metrics sample", "sql": "SELECT device_id, metric_name, metric_value, recorded_at FROM telemetry.device_metrics LIMIT 10"},
{"id": "cs5", "label": "Metrics by device", "sql": "SELECT device_id, count(*) AS cnt FROM telemetry.device_metrics GROUP BY device_id LIMIT 10"},
{"id": "cs6", "label": "Distinct metric names", "sql": "SELECT DISTINCT metric_name FROM telemetry.device_metrics LIMIT 20"},
{"id": "cs7", "label": "Table columns", "sql": "SELECT column_name, type FROM system_schema.columns WHERE keyspace_name='telemetry' AND table_name='device_metrics'"},
{"id": "cs8", "label": "Recent metrics", "sql": "SELECT device_id, metric_name, metric_value FROM telemetry.device_metrics LIMIT 15"},
{"id": "cs9", "label": "Partition token sample", "sql": "SELECT token(device_id) AS tok, device_id FROM telemetry.device_metrics LIMIT 5"},
{"id": "cs10", "label": "Cluster peers", "sql": "SELECT peer, data_center, rack FROM system.peers"},
],
"neo4j": [
{"id": "nj1", "label": "Node labels", "sql": "CALL db.labels() YIELD label RETURN label"},
{"id": "nj2", "label": "Relationship types", "sql": "CALL db.relationshipTypes() YIELD relationshipType RETURN relationshipType"},
{"id": "nj3", "label": "Total nodes", "sql": "MATCH (n) RETURN count(n) AS node_count"},
{"id": "nj4", "label": "Total relationships", "sql": "MATCH ()-[r]->() RETURN count(r) AS rel_count"},
{"id": "nj5", "label": "Products sample", "sql": "MATCH (p:Product) RETURN p.product_id AS id, p.name AS name, p.category AS category LIMIT 10"},
{"id": "nj6", "label": "Suppliers sample", "sql": "MATCH (s:Supplier) RETURN s.supplier_id AS id, s.name AS name, s.country AS country LIMIT 10"},
{"id": "nj7", "label": "Supply chain links", "sql": "MATCH (s:Supplier)-[r:SUPPLIES]->(p:Product) RETURN s.name AS supplier, p.name AS product LIMIT 10"},
{"id": "nj8", "label": "Products by category", "sql": "MATCH (p:Product) RETURN p.category AS category, count(*) AS cnt ORDER BY cnt DESC LIMIT 10"},
{"id": "nj9", "label": "Schema visualization", "sql": "CALL db.schema.visualization()"},
{"id": "nj10", "label": "Constraint info", "sql": "SHOW CONSTRAINTS"},
],
"trino": [
{"id": "tq1", "label": "Trino version", "sql": "SELECT version()"},
{"id": "tq2", "label": "Catalogs", "sql": "SHOW CATALOGS"},
{"id": "tq3", "label": "Iceberg schemas", "sql": "SHOW SCHEMAS FROM iceberg"},
{"id": "tq4", "label": "Sales tables (federated PG)", "sql": "SHOW TABLES FROM postgres_sales.public"},
{"id": "tq5", "label": "Federated PG catalog", "sql": "SHOW SCHEMAS FROM postgres_sales"},
{"id": "tq6", "label": "Cluster nodes", "sql": "SELECT node_id, http_uri, state FROM system.runtime.nodes"},
{"id": "tq7", "label": "Running queries", "sql": "SELECT query_id, state, user FROM system.runtime.queries WHERE state != 'FINISHED' LIMIT 10"},
{"id": "tq8", "label": "Count 1M rows (parallel)", "sql": "SELECT count(*) AS rows FROM unnest(sequence(1,1000)) a(x) CROSS JOIN unnest(sequence(1,1000)) b(y)"},
{"id": "tq9", "label": "Federated sales by region", "sql": "SELECT region, count(*) AS orders FROM postgres_sales.public.sales_orders WHERE order_id < 500000 GROUP BY region ORDER BY orders DESC"},
{"id": "tq10", "label": "Parallel sin() 1M rows", "sql": "SELECT sum(sin(a.x)) AS total FROM unnest(sequence(1,1000)) a(x) CROSS JOIN unnest(sequence(1,1000)) b(y)"},
],
}
BENCHMARK_SQL = {
"postgres": "SELECT sum(sin(i)) AS total FROM generate_series(1, 5000000) i",
"trino": "SELECT sum(sin(x + (y - 1) * 5000)) AS total FROM unnest(sequence(1,5000)) a(x) CROSS JOIN unnest(sequence(1,1000)) b(y)",
}
def _tabular(columns: list[str], rows: list[list[Any]], elapsed_ms: int, **extra: Any) -> dict[str, Any]:
return {
"ok": True,
"columns": columns,
"rows": rows,
"row_count": len(rows),
"elapsed_ms": elapsed_ms,
"truncated": extra.pop("truncated", False),
**extra,
}
def _fmt(val: Any) -> Any:
if val is None or isinstance(val, (str, int, float, bool)):
return val
return str(val)[:200]
def _run_postgres(sql: str, limit: int = 200) -> dict[str, Any]:
t0 = time.perf_counter()
conn = psycopg2.connect(
host=DB_HOST, port=PG_PORT, user=PG_USER, password=PG_PASS, dbname=PG_DB, connect_timeout=8,
)
try:
conn.set_session(readonly=True, autocommit=True)
cur = conn.cursor()
cur.execute(sql)
elapsed_ms = int((time.perf_counter() - t0) * 1000)
if cur.description:
columns = [d[0] for d in cur.description]
rows = [[_fmt(r) for r in row] for row in cur.fetchmany(limit)]
return _tabular(columns, rows, elapsed_ms, truncated=len(rows) >= limit)
return _tabular([], [], elapsed_ms, message="OK")
finally:
conn.close()
def _run_mysql(sql: str, limit: int = 200) -> dict[str, Any]:
t0 = time.perf_counter()
conn = pymysql.connect(
host=DB_HOST, port=MYSQL_PORT, user=MYSQL_USER, password=MYSQL_PASS,
database=MYSQL_DB, connect_timeout=8, read_timeout=30,
)
try:
cur = conn.cursor()
cur.execute(sql)
elapsed_ms = int((time.perf_counter() - t0) * 1000)
if cur.description:
columns = [d[0] for d in cur.description]
rows = [[_fmt(r) for r in row] for row in cur.fetchmany(limit)]
return _tabular(columns, rows, elapsed_ms, truncated=len(rows) >= limit)
return _tabular([], [], elapsed_ms, message="OK")
finally:
conn.close()
def _mongo_client() -> MongoClient:
return MongoClient(f"mongodb://{MONGO_HOST}:{MONGO_PORT}/", serverSelectionTimeoutMS=8000)
def _run_mongo(query: str, limit: int = 200) -> dict[str, Any]:
t0 = time.perf_counter()
q = query.strip()
client = _mongo_client()
try:
columns: list[str] = []
rows: list[list[Any]] = []
if q == "SHOW DATABASES":
columns = ["database"]
rows = [[name] for name in client.list_database_names()]
elif m := re.match(r"^SHOW COLLECTIONS(?:\s+(\w+))?$", q, re.I):
db_name = m.group(1) or MONGO_DB
columns = ["collection"]
rows = [[name] for name in client[db_name].list_collection_names()]
elif m := re.match(r"^COUNT\s+(\w+)\.(\w+)$", q, re.I):
db_name, coll = m.group(1), m.group(2)
columns = ["count"]
rows = [[client[db_name][coll].estimated_document_count()]]
elif m := re.match(r"^FIND\s+(\w+)\.(\w+)(?:\s+SORT\s+(\w+)\s+DESC)?(?:\s+LIMIT\s+(\d+))?$", q, re.I):
db_name, coll, sort_field, lim = m.group(1), m.group(2), m.group(3), m.group(4)
n = min(int(lim) if lim else 10, limit)
cursor = client[db_name][coll].find({})
if sort_field:
cursor = cursor.sort(sort_field, -1)
docs = list(cursor.limit(n))
if not docs:
columns = ["result"]
rows = [["(empty)"]]
else:
columns = sorted({k for d in docs for k in d if k != "_id"})
rows = [[_fmt(d.get(c)) for c in columns] for d in docs]
elif m := re.match(r"^DISTINCT\s+(\w+)\.(\w+)\s+(\w+)$", q, re.I):
db_name, coll, field = m.group(1), m.group(2), m.group(3)
columns = [field]
vals = client[db_name][coll].distinct(field)[:limit]
rows = [[_fmt(v)] for v in vals]
elif m := re.match(r"^AGGREGATE\s+(\w+)\.(\w+)\s+GROUP\s+(\w+)\s+TOP\s+(\d+)$", q, re.I):
db_name, coll, field, top_n = m.group(1), m.group(2), m.group(3), int(m.group(4))
pipe = [
{"$group": {"_id": f"${field}", "count": {"$sum": 1}}},
{"$sort": {"count": -1}},
{"$limit": min(top_n, limit)},
]
columns = [field, "count"]
rows = [[r.get("_id"), r.get("count")] for r in client[db_name][coll].aggregate(pipe, maxTimeMS=15000)]
elif m := re.match(r"^INDEXES\s+(\w+)\.(\w+)$", q, re.I):
db_name, coll = m.group(1), m.group(2)
columns = ["name", "keys"]
rows = [[i.get("name"), str(i.get("key"))] for i in client[db_name][coll].list_indexes()]
elif q.upper() == "SERVER STATUS":
status = client.admin.command("serverStatus")
columns = ["metric", "value"]
rows = [
["version", status.get("version")],
["uptime_secs", status.get("uptime")],
["connections_current", status.get("connections", {}).get("current")],
["mem_resident_mb", status.get("mem", {}).get("resident")],
]
else:
return {"ok": False, "error": f"Unknown MongoDB command: {q[:120]}"}
elapsed_ms = int((time.perf_counter() - t0) * 1000)
return _tabular(columns, rows[:limit], elapsed_ms, truncated=len(rows) > limit)
finally:
client.close()
def _cass_cluster() -> Cluster:
auth = PlainTextAuthProvider(CASS_USER, CASS_PASS) if CASS_USER else None
return Cluster([DB_HOST], port=CASS_PORT, auth_provider=auth, connect_timeout=8)
def _run_cassandra(cql: str, limit: int = 200) -> dict[str, Any]:
t0 = time.perf_counter()
cluster = _cass_cluster()
session = cluster.connect()
try:
rows_raw = list(session.execute(cql, timeout=30))
elapsed_ms = int((time.perf_counter() - t0) * 1000)
if not rows_raw:
return _tabular([], [], elapsed_ms, message="OK")
columns = list(rows_raw[0]._fields)
rows = [[_fmt(getattr(r, c)) for c in columns] for r in rows_raw[:limit]]
return _tabular(columns, rows, elapsed_ms, truncated=len(rows_raw) > limit)
finally:
cluster.shutdown()
def _neo4j_driver():
return GraphDatabase.driver(NEO4J_URI, auth=(NEO4J_USER, NEO4J_PASS))
def _run_neo4j(cypher: str, limit: int = 200) -> dict[str, Any]:
t0 = time.perf_counter()
driver = _neo4j_driver()
try:
with driver.session() as session:
result = session.run(cypher)
keys = result.keys()
columns = list(keys)
rows_raw = [record.values() for record in result]
elapsed_ms = int((time.perf_counter() - t0) * 1000)
rows = [[_fmt(v) for v in row] for row in rows_raw[:limit]]
return _tabular(columns, rows, elapsed_ms, truncated=len(rows_raw) > limit)
finally:
driver.close()
def _run_trino(sql: str, limit: int = 200) -> dict[str, Any]:
t0 = time.perf_counter()
headers = {"X-Trino-User": TRINO_USER, "Content-Type": "text/plain"}
with httpx.Client(timeout=120.0) as client:
resp = client.post(f"{TRINO_URL}/v1/statement", content=sql, headers=headers)
if resp.status_code >= 400:
return {"ok": False, "error": resp.text[:500]}
data = resp.json()
if data.get("error"):
return {"ok": False, "error": str(data["error"])[:500]}
columns: list[str] = []
rows: list[list[Any]] = []
while True:
if data.get("columns") and not columns:
columns = [c["name"] for c in data["columns"]]
if data.get("data"):
rows.extend(data["data"])
if len(rows) >= limit:
rows = rows[:limit]
break
if data.get("error"):
return {"ok": False, "error": str(data["error"])[:500]}
nxt = data.get("nextUri")
if not nxt:
break
data = client.get(nxt, headers=headers).json()
elapsed_ms = int((time.perf_counter() - t0) * 1000)
stats = data.get("stats") or {}
return _tabular(columns, rows, elapsed_ms, engine_stats_ms=stats.get("elapsedTimeMillis"), truncated=len(rows) >= limit)
class SqlRequest(BaseModel):
engine: str = Field(..., pattern="^(postgres|mysql|mongodb|trino|cassandra|neo4j)$")
sql: str = Field(..., min_length=1, max_length=8000)
def _connection_info(engine: str) -> dict[str, str]:
if engine == "postgres":
return {"host": DB_HOST, "port": str(PG_PORT), "database": PG_DB, "user": PG_USER}
if engine == "mysql":
return {"host": DB_HOST, "port": str(MYSQL_PORT), "database": MYSQL_DB, "user": MYSQL_USER}
if engine == "mongodb":
return {"host": MONGO_HOST, "port": str(MONGO_PORT), "database": MONGO_DB}
if engine == "cassandra":
return {"host": DB_HOST, "port": str(CASS_PORT), "keyspace": CASS_KS, "user": CASS_USER}
if engine == "neo4j":
return {"uri": NEO4J_URI, "user": NEO4J_USER}
return {"url": TRINO_URL, "user": TRINO_USER}
def _dispatch(engine: str, sql: str, limit: int = 200) -> dict[str, Any]:
if engine == "postgres":
return _run_postgres(sql, limit)
if engine == "mysql":
return _run_mysql(sql, limit)
if engine == "mongodb":
return _run_mongo(sql, limit)
if engine == "cassandra":
return _run_cassandra(sql, limit)
if engine == "neo4j":
return _run_neo4j(sql, limit)
return _run_trino(sql, limit)
def _catalog_postgres() -> dict[str, Any]:
conn = psycopg2.connect(host=DB_HOST, port=PG_PORT, user=PG_USER, password=PG_PASS, dbname=PG_DB, connect_timeout=8)
try:
conn.set_session(readonly=True, autocommit=True)
cur = conn.cursor()
cur.execute(
"SELECT schemaname, relname, n_live_tup FROM pg_stat_user_tables ORDER BY n_live_tup DESC NULLS LAST LIMIT 50"
)
objects = []
for schema, table, rows in cur.fetchall():
objects.append({"type": "table", "schema": schema, "name": table, "fqn": f"{schema}.{table}", "row_count": int(rows or 0)})
cur.execute("SELECT version()")
version = cur.fetchone()[0]
return {"engine": "postgres", "version": version, "objects": objects}
finally:
conn.close()
def _catalog_mysql() -> dict[str, Any]:
conn = pymysql.connect(host=DB_HOST, port=MYSQL_PORT, user=MYSQL_USER, password=MYSQL_PASS, database=MYSQL_DB, connect_timeout=8)
try:
cur = conn.cursor()
cur.execute(
"SELECT table_name, table_rows, ROUND((data_length+index_length)/1024/1024,1) AS mb "
"FROM information_schema.tables WHERE table_schema=%s ORDER BY table_rows DESC",
(MYSQL_DB,),
)
objects = [{"type": "table", "schema": MYSQL_DB, "name": r[0], "fqn": f"{MYSQL_DB}.{r[0]}", "row_count": int(r[1] or 0), "size_mb": float(r[2] or 0)} for r in cur.fetchall()]
cur.execute("SELECT VERSION()")
version = cur.fetchone()[0]
return {"engine": "mysql", "version": version, "objects": objects}
finally:
conn.close()
def _catalog_mongodb() -> dict[str, Any]:
client = _mongo_client()
try:
objects = []
for db_name in client.list_database_names():
if db_name in ("admin", "config", "local"):
continue
for coll in client[db_name].list_collection_names():
try:
cnt = client[db_name][coll].estimated_document_count()
except Exception:
cnt = None
objects.append({"type": "collection", "schema": db_name, "name": coll, "fqn": f"{db_name}.{coll}", "row_count": cnt})
build = client.admin.command("buildInfo")
return {"engine": "mongodb", "version": build.get("version"), "objects": objects}
finally:
client.close()
def _catalog_cassandra() -> dict[str, Any]:
cluster = _cass_cluster()
session = cluster.connect()
try:
objects = []
kss = [
r.keyspace_name for r in session.execute("SELECT keyspace_name FROM system_schema.keyspaces")
if not r.keyspace_name.startswith("system")
]
for ks in kss:
tables = session.execute(
"SELECT table_name FROM system_schema.tables WHERE keyspace_name=%s", (ks,)
)
for t in tables:
objects.append({"type": "table", "schema": ks, "name": t.table_name, "fqn": f"{ks}.{t.table_name}", "row_count": None})
ver = list(session.execute("SELECT release_version FROM system.local"))[0].release_version
return {"engine": "cassandra", "version": ver, "objects": objects}
finally:
cluster.shutdown()
def _catalog_neo4j() -> dict[str, Any]:
driver = _neo4j_driver()
try:
with driver.session() as session:
labels = [r["label"] for r in session.run("CALL db.labels() YIELD label RETURN label")]
rels = [r["relationshipType"] for r in session.run("CALL db.relationshipTypes() YIELD relationshipType RETURN relationshipType")]
objects = []
for label in labels:
cnt = session.run(f"MATCH (n:`{label}`) RETURN count(n) AS c").single()["c"]
objects.append({"type": "node_label", "schema": "graph", "name": label, "fqn": label, "row_count": cnt})
for rel in rels:
cnt = session.run(f"MATCH ()-[r:`{rel}`]->() RETURN count(r) AS c").single()["c"]
objects.append({"type": "relationship", "schema": "graph", "name": rel, "fqn": rel, "row_count": cnt})
total = session.run("MATCH (n) RETURN count(n) AS c").single()["c"]
return {"engine": "neo4j", "version": "4.4", "total_nodes": total, "objects": objects}
finally:
driver.close()
def _catalog(engine: str) -> dict[str, Any]:
if engine == "postgres":
return _catalog_postgres()
if engine == "mysql":
return _catalog_mysql()
if engine == "mongodb":
return _catalog_mongodb()
if engine == "cassandra":
return _catalog_cassandra()
if engine == "neo4j":
return _catalog_neo4j()
raise ValueError(f"Catalog not supported for {engine}")
def _sample_postgres(object_name: str, limit: int) -> dict[str, Any]:
if "." in object_name:
schema, table = object_name.split(".", 1)
sql = f'SELECT * FROM "{schema}"."{table}" LIMIT {limit}'
else:
sql = f'SELECT * FROM public."{object_name}" LIMIT {limit}'
return _run_postgres(sql, limit)
def _sample_mysql(object_name: str, limit: int) -> dict[str, Any]:
table = object_name.split(".")[-1]
return _run_mysql(f"SELECT * FROM `{table}` LIMIT {limit}", limit)
def _sample_mongodb(object_name: str, limit: int) -> dict[str, Any]:
if "." in object_name:
db_name, coll = object_name.split(".", 1)
else:
db_name, coll = MONGO_DB, object_name
return _run_mongo(f"FIND {db_name}.{coll} LIMIT {limit}", limit)
def _sample_cassandra(object_name: str, limit: int) -> dict[str, Any]:
if "." in object_name:
ks, table = object_name.split(".", 1)
else:
ks, table = CASS_KS, object_name
return _run_cassandra(f"SELECT * FROM {ks}.{table} LIMIT {limit}", limit)
def _graph_neo4j(edge_limit: int = 60, rel_type: str | None = None) -> dict[str, Any]:
"""Return a sampled subgraph (nodes + edges) for interactive graph visualization."""
t0 = time.perf_counter()
driver = _neo4j_driver()
try:
with driver.session() as session:
total = session.run("MATCH (n) RETURN count(n) AS c").single()["c"]
rel_types = [
r["relationshipType"]
for r in session.run(
"CALL db.relationshipTypes() YIELD relationshipType RETURN relationshipType"
)
]
if rel_type:
cypher = """
MATCH (a)-[r]->(b)
WHERE type(r) = $rel_type
WITH a, r, b LIMIT $limit
RETURN a, r, b, labels(a)[0] AS la, type(r) AS rt, labels(b)[0] AS lb
"""
params: dict[str, Any] = {"rel_type": rel_type, "limit": edge_limit}
else:
cypher = """
MATCH (a)-[r]->(b)
WITH a, r, b LIMIT $limit
RETURN a, r, b, labels(a)[0] AS la, type(r) AS rt, labels(b)[0] AS lb
"""
params = {"limit": edge_limit}
records = session.run(cypher, **params).data()
def _node_key(node: Any, label: str) -> str:
props = dict(node)
if label == "Product" and props.get("product_id") is not None:
return f"Product:{props['product_id']}"
if label == "Supplier" and props.get("supplier_id") is not None:
return f"Supplier:{props['supplier_id']}"
nid = getattr(node, "element_id", None) or getattr(node, "id", None)
return f"{label}:{nid}"
nodes_map: dict[str, dict[str, Any]] = {}
edges: list[dict[str, Any]] = []
for rec in records:
a, b, rt = rec["a"], rec["b"], rec["rt"]
la, lb = rec["la"], rec["lb"]
aid, bid = _node_key(a, la), _node_key(b, lb)
if aid not in nodes_map:
ap = dict(a)
nodes_map[aid] = {
"id": aid,
"label": la,
"name": ap.get("name") or str(ap.get("product_id") or ap.get("supplier_id") or aid),
"properties": {k: _fmt(v) for k, v in ap.items()},
}
if bid not in nodes_map:
bp = dict(b)
nodes_map[bid] = {
"id": bid,
"label": lb,
"name": bp.get("name") or str(bp.get("product_id") or bp.get("supplier_id") or bid),
"properties": {k: _fmt(v) for k, v in bp.items()},
}
edges.append({"id": f"{aid}|{rt}|{bid}", "source": aid, "target": bid, "type": rt})
elapsed_ms = int((time.perf_counter() - t0) * 1000)
return {
"ok": True,
"nodes": list(nodes_map.values()),
"edges": edges,
"rel_types": rel_types,
"stats": {
"total_nodes_db": total,
"returned_nodes": len(nodes_map),
"returned_edges": len(edges),
},
"elapsed_ms": elapsed_ms,
}
finally:
driver.close()
def _sample_neo4j(object_name: str, limit: int) -> dict[str, Any]:
label = object_name.split(".")[-1]
cypher = f"MATCH (n:`{label}`) RETURN n LIMIT {limit}"
t0 = time.perf_counter()
driver = _neo4j_driver()
try:
with driver.session() as session:
result = session.run(cypher)
docs = [dict(record["n"]) for record in result]
elapsed_ms = int((time.perf_counter() - t0) * 1000)
if not docs:
return _tabular(["result"], [["(empty)"]], elapsed_ms)
columns = sorted({k for d in docs for k in d})
rows = [[_fmt(d.get(c)) for c in columns] for d in docs]
return _tabular(columns, rows, elapsed_ms, object=object_name)
finally:
driver.close()
def _sample(engine: str, object_name: str, limit: int) -> dict[str, Any]:
if engine == "postgres":
return _sample_postgres(object_name, limit)
if engine == "mysql":
return _sample_mysql(object_name, limit)
if engine == "mongodb":
return _sample_mongodb(object_name, limit)
if engine == "cassandra":
return _sample_cassandra(object_name, limit)
if engine == "neo4j":
return _sample_neo4j(object_name, limit)
raise ValueError(f"Sample not supported for {engine}")
@router.get("/samples/{engine}")
async def get_samples(engine: str):
if engine not in SAMPLES:
return JSONResponse({"error": "unknown engine"}, status_code=404)
return {"engine": engine, "samples": SAMPLES[engine], "connection": _connection_info(engine)}
@router.get("/health")
async def sql_health():
out: dict[str, Any] = {}
for eng in SOURCE_ENGINES:
try:
if eng == "postgres":
_run_postgres("SELECT 1")
out[eng] = {"ok": True, "host": DB_HOST, "user": PG_USER, "error": None}
elif eng == "mysql":
_run_mysql("SELECT 1")
out[eng] = {"ok": True, "host": DB_HOST, "database": MYSQL_DB, "user": MYSQL_USER, "error": None}
elif eng == "mongodb":
_run_mongo("SHOW DATABASES")
out[eng] = {"ok": True, "host": MONGO_HOST, "database": MONGO_DB, "error": None}
elif eng == "cassandra":
_run_cassandra("SELECT release_version FROM system.local LIMIT 1")
out[eng] = {"ok": True, "host": DB_HOST, "keyspace": CASS_KS, "error": None}
else:
_run_neo4j("RETURN 1 AS n")
out[eng] = {"ok": True, "uri": NEO4J_URI, "user": NEO4J_USER, "error": None}
except Exception as exc:
out[eng] = {"ok": False, "error": str(exc)[:200]}
return out
@router.get("/catalog/{engine}")
async def get_catalog(engine: str):
if engine not in SOURCE_ENGINES:
return JSONResponse({"error": "unknown engine"}, status_code=404)
try:
return _catalog(engine)
except Exception as exc:
return JSONResponse({"error": str(exc)[:500]}, status_code=502)
@router.get("/sample/{engine}")
async def get_sample(engine: str, object: str = Query(..., min_length=1), limit: int = Query(50, ge=1, le=200)):
if engine not in SOURCE_ENGINES:
return JSONResponse({"error": "unknown engine"}, status_code=404)
try:
result = _sample(engine, object, limit)
if not result.get("ok"):
return JSONResponse(result, status_code=422)
return {**result, "engine": engine, "object": object}
except Exception as exc:
return JSONResponse({"ok": False, "error": str(exc)[:500]}, status_code=502)
@router.get("/graph/neo4j")
async def get_neo4j_graph(
limit: int = Query(60, ge=10, le=150),
rel_type: str | None = Query(None),
):
try:
result = _graph_neo4j(limit, rel_type or None)
if not result.get("ok"):
return JSONResponse(result, status_code=422)
return result
except Exception as exc:
return JSONResponse({"ok": False, "error": str(exc)[:500]}, status_code=502)
@router.post("/execute")
async def execute_sql(body: SqlRequest):
sql = body.sql.strip().rstrip(";")
if not sql:
return JSONResponse({"ok": False, "error": "Empty query"}, status_code=400)
try:
result = _dispatch(body.engine, sql)
if not result.get("ok"):
return JSONResponse(result, status_code=422)
return {**result, "engine": body.engine, "sql": sql}
except Exception as exc:
return JSONResponse({"ok": False, "error": str(exc)[:500]}, status_code=502)
@router.post("/benchmark")
async def benchmark():
"""Parallel analytics: PostgreSQL OLTP vs Trino distributed engine (5M sin() rows)."""
results: dict[str, Any] = {}
for engine, sql in BENCHMARK_SQL.items():
try:
results[engine] = _dispatch(engine, sql)
results[engine]["sql"] = sql
except Exception as exc:
results[engine] = {"ok": False, "error": str(exc)[:300], "sql": sql}
pg_ms = results.get("postgres", {}).get("elapsed_ms")
tr_ms = results.get("trino", {}).get("elapsed_ms")
faster = None
speedup = None
if pg_ms and tr_ms:
faster = "trino" if tr_ms < pg_ms else "postgres"
speedup = round(pg_ms / tr_ms, 2) if tr_ms < pg_ms else round(tr_ms / pg_ms, 2)
return {
"ok": True,
"postgres": results.get("postgres"),
"trino": results.get("trino"),
"comparison": {
"postgres_ms": pg_ms,
"trino_ms": tr_ms,
"faster": faster,
"speedup_factor": speedup,
"note": "5M row parallel compute — Trino distributes across workers; PostgreSQL single-node OLTP",
},
}