Skip to content

ER/Schema Visualizer — Python source

Paste CREATE TABLE DDL and get an ER diagram as SVG: tables with typed columns, primary keys, and foreign-key arrows in a deterministic layered layout. Pan and zoom the live diagram; export the SVG.

This is the Python implementation — the same logic the interactive tool runs, in a shareable, citable form.

"""schema-visualizer — pure CREATE TABLE DDL → layered ER diagram as SVG.
Python port (canonical TS: src/lib/schema-visualizer.ts; Go twin:
cli/schema-visualizer). Tolerant common subset of Postgres/MySQL/SQLite:
unparseable statements degrade to notes, never raise. Integer geometry only,
so every port draws the identical diagram (byte-identical SVG).
"""
import re

LAYOUT = {"row_height": 24, "char_width": 7, "padding": 8, "layer_gap": 60, "column_gap": 40}
MODIFIERS = {"NOT", "NULL", "PRIMARY", "KEY", "UNIQUE", "DEFAULT", "REFERENCES",
             "AUTO_INCREMENT", "AUTOINCREMENT", "ON", "COMMENT", "CHECK", "CONSTRAINT"}

QUOTES = {"'": "'", '"': '"', "`": "`", "[": "]"}


def _r(v):
    """Round half up (JS Math.round parity — Python round() is banker's)."""
    return int(v + 0.5)


def _split_statements(ddl):
    """Split on `;` outside strings/quoted identifiers (depth-agnostic: an
    unterminated paren cannot swallow the statements after it)."""
    out, cur, i, n = [], "", 0, len(ddl)
    while i < n:
        ch = ddl[i]
        if ch in QUOTES:
            close = QUOTES[ch]
            cur += ch
            i += 1
            while i < n:
                cur += ddl[i]
                if ddl[i] == close:
                    if close == "'" and i + 1 < n and ddl[i + 1] == "'":
                        cur += ddl[i + 1]
                        i += 2
                        continue
                    break
                i += 1
            i += 1
            continue
        if ch == ";":
            out.append(cur)
            cur = ""
            i += 1
            continue
        cur += ch
        i += 1
    if cur.strip():
        out.append(cur)
    return out


def _tokenize(s):
    """Tokens: quoted identifiers/strings carry their text (quotes stripped);
    ( ) , . are punct; everything else is a word."""
    toks, i, n = [], 0, len(s)
    while i < n:
        ch = s[i]
        if ch.isspace():
            i += 1
            continue
        if ch in QUOTES:
            close = QUOTES[ch]
            text = ""
            i += 1
            while i < n:
                if s[i] == close:
                    if close == "'" and i + 1 < n and s[i + 1] == "'":
                        text += "'"
                        i += 2
                        continue
                    break
                text += s[i]
                i += 1
            i += 1
            toks.append(("string" if ch == "'" else "qident", text))
            continue
        if ch in "(),.":
            toks.append(("punct", ch))
            i += 1
            continue
        j = i
        while j < n and not s[j].isspace() and s[j] not in "',().`[]":
            j += 1
        toks.append(("word", s[i:j]))
        i = j
    return toks


def _is_p(tok, p):
    return tok is not None and tok[0] == "punct" and tok[1] == p


def _kw(tok, w):
    return tok is not None and tok[0] == "word" and tok[1].upper() == w


def _take_name(toks, i):
    if i >= len(toks) or toks[i][0] not in ("qident", "word"):
        return None
    name, j = toks[i][1], i + 1
    while (_is_p(_g(toks, j), ".") and _g(toks, j + 1) is not None
           and _g(toks, j + 1)[0] in ("qident", "word")):
        name += "." + toks[j + 1][1]
        j += 2
    return name, j


def _g(toks, i):
    return toks[i] if i < len(toks) else None


def _paren_list(toks, i):
    if not _is_p(_g(toks, i), "("):
        return None
    names, j = [], i + 1
    while True:
        name = _take_name(toks, j)
        if name is None:
            return None
        names.append(name[0])
        j = name[1]
        if _is_p(_g(toks, j), ","):
            j += 1
            continue
        if _is_p(_g(toks, j), ")"):
            return names, j + 1
        return None


def _join_type(toks):
    raw = " ".join(t[1] for t in toks)
    raw = re.sub(r"\s*\(\s*", "(", raw)
    raw = re.sub(r"\s*\)\s*", ")", raw)
    raw = re.sub(r"\s*,\s*", ",", raw)
    return raw.strip().upper()


def _parse_column(line, table_name, fks):
    name = _take_name(line, 0)
    if name is None:
        return None
    i = name[1]
    type_toks = []
    while i < len(line) and not (line[i][0] == "word" and line[i][1].upper() in MODIFIERS):
        type_toks.append(line[i])
        i += 1
    nullable, pk = True, False
    while i < len(line):
        t = line[i]
        if _kw(t, "NOT") and _kw(_g(line, i + 1), "NULL"):
            nullable = False
            i += 2
            continue
        if _kw(t, "NULL"):
            i += 1
            continue
        if _kw(t, "PRIMARY") and _kw(_g(line, i + 1), "KEY"):
            pk, nullable = True, False
            i += 2
            continue
        if _kw(t, "UNIQUE") or _kw(t, "AUTO_INCREMENT") or _kw(t, "AUTOINCREMENT"):
            i += 1
            continue
        if _kw(t, "DEFAULT"):
            i += 1
            if _is_p(_g(line, i), "("):
                depth = 0
                while i < len(line):
                    if _is_p(line[i], "("):
                        depth += 1
                    if _is_p(line[i], ")"):
                        depth -= 1
                    i += 1
                    if depth == 0:
                        break
            elif i < len(line):
                i += 1
            continue
        if _kw(t, "COMMENT"):
            i += 1
            if i < len(line) and line[i][0] == "string":
                i += 1
            continue
        if _kw(t, "ON"):
            i += 2
            if _kw(_g(line, i), "SET") or _kw(_g(line, i), "NO"):
                i += 2
            elif i < len(line):
                i += 1
            continue
        if _kw(t, "REFERENCES"):
            i += 1
            target = _take_name(line, i)
            if target is not None:
                i = target[1]
                to_col = None
                if _is_p(_g(line, i), "("):
                    lst = _paren_list(line, i)
                    if lst is not None:
                        to_col, i = lst[0][0], lst[1]
                fks.append({"fromTable": table_name, "fromColumn": name[0],
                            "toTable": target[0], "toColumn": to_col})
            continue
        i += 1  # unknown modifier tolerated
    return {"name": name[0], "type": _join_type(type_toks), "nullable": nullable, "isPrimaryKey": pk}


def parse_ddl(ddl):
    if not ddl.strip():
        return {"tables": [], "foreignKeys": [], "notes": ["No DDL input."]}
    tables, fks, notes = [], [], []
    for stmt in _split_statements(ddl):
        if not stmt.strip():
            continue
        toks = _tokenize(stmt)
        try:
            i = 0
            if not _kw(_g(toks, i), "CREATE"):
                raise ValueError()
            i += 1
            while _kw(_g(toks, i), "TEMP") or _kw(_g(toks, i), "TEMPORARY") or _kw(_g(toks, i), "UNLOGGED"):
                i += 1
            if not _kw(_g(toks, i), "TABLE"):
                notes.append("Skipped non-table statement.")
                continue
            i += 1
            if _kw(_g(toks, i), "IF") and _kw(_g(toks, i + 1), "NOT") and _kw(_g(toks, i + 2), "EXISTS"):
                i += 3
            name = _take_name(toks, i)
            if name is None or not _is_p(_g(toks, name[1]), "("):
                raise ValueError()
            i = name[1] + 1
            body, depth = [], 0
            while i < len(toks):
                if _is_p(toks[i], "("):
                    depth += 1
                if _is_p(toks[i], ")"):
                    if depth == 0:
                        break
                    depth -= 1
                body.append(toks[i])
                i += 1
            if i >= len(toks):
                raise ValueError()
            lines, line, depth = [], [], 0
            for t in body:
                if _is_p(t, "("):
                    depth += 1
                if _is_p(t, ")"):
                    depth -= 1
                if _is_p(t, ",") and depth == 0:
                    lines.append(line)
                    line = []
                    continue
                line.append(t)
            if line:
                lines.append(line)
            table = {"name": name[0], "columns": []}
            tables.append(table)
            for toks2 in lines:
                if not toks2:
                    continue
                first = toks2[0]
                u = first[1].upper() if first[0] == "word" else ""
                if u == "PRIMARY" and _kw(_g(toks2, 1), "KEY"):
                    lst = _paren_list(toks2, 2)
                    if lst:
                        for cn in lst[0]:
                            for col in table["columns"]:
                                if col["name"] == cn:
                                    col["isPrimaryKey"] = True
                                    col["nullable"] = False
                    continue
                if u == "FOREIGN" and _kw(_g(toks2, 1), "KEY"):
                    frm = _paren_list(toks2, 2)
                    if frm and _kw(_g(toks2, frm[1]), "REFERENCES"):
                        target = _take_name(toks2, frm[1] + 1)
                        if target is not None:
                            to_cols = None
                            if _is_p(_g(toks2, target[1]), "("):
                                to = _paren_list(toks2, target[1])
                                if to:
                                    to_cols = to[0]
                            for idx, fc in enumerate(frm[0]):
                                to_col = None
                                if to_cols:
                                    to_col = to_cols[idx] if idx < len(to_cols) else to_cols[-1]
                                fks.append({"fromTable": table["name"], "fromColumn": fc,
                                            "toTable": target[0], "toColumn": to_col})
                    continue
                if u in ("UNIQUE", "KEY", "INDEX", "CHECK", "EXCLUDE", "CONSTRAINT"):
                    continue
                col = _parse_column(toks2, table["name"], fks)
                if col is not None:
                    table["columns"].append(col)
        except ValueError:
            notes.append("Skipped unparseable statement.")
    resolved = []
    for fk in fks:
        if fk["toColumn"]:
            resolved.append(fk)
            continue
        target = next((t for t in tables if t["name"] == fk["toTable"]), None)
        pk = next((c for c in (target or {"columns": []})["columns"] if c["isPrimaryKey"]), None)
        fk = dict(fk, toColumn=pk["name"] if pk else "id")
        resolved.append(fk)
    return {"tables": tables, "foreignKeys": resolved, "notes": notes}


def layout_schema(schema):
    o = LAYOUT
    if not schema["tables"]:
        return {"width": 0, "height": 0, "tables": [], "edges": []}
    index = {}
    for i, t in enumerate(schema["tables"]):
        index.setdefault(t["name"], i)
    boxes = []
    for t in schema["tables"]:
        text_len = max([t["name"]] + [f"{c['name']} {c['type']}" for c in t["columns"]] + [""], key=len)
        text_len = len(text_len) or 1
        boxes.append({"x": 0, "y": 0,
                      "w": _r(text_len * o["char_width"] + 2 * o["padding"]),
                      "h": _r(o["row_height"] * (1 + len(t["columns"])) + o["padding"])})
    layer_of = [0] * len(schema["tables"])
    for _ in range(len(schema["tables"])):
        changed = False
        for fk in schema["foreignKeys"]:
            ti, tj = index.get(fk["fromTable"]), index.get(fk["toTable"])
            if ti is None or tj is None or ti == tj:
                continue
            if layer_of[ti] < layer_of[tj] + 1:
                layer_of[ti] = layer_of[tj] + 1
                changed = True
        if not changed:
            break
    layers = {}
    for i, l in enumerate(layer_of):
        layers.setdefault(l, []).append(i)
    y = width = height = 0
    for li in sorted(layers):
        x = layer_h = 0
        for i in layers[li]:
            boxes[i]["x"], boxes[i]["y"] = x, y
            x += boxes[i]["w"] + o["column_gap"]
            layer_h = max(layer_h, boxes[i]["h"])
        width = max(width, x - o["column_gap"])
        height = max(height, y + layer_h)
        y += layer_h + o["layer_gap"]
    edges = []
    for fk in schema["foreignKeys"]:
        frm, to = index.get(fk["fromTable"]), index.get(fk["toTable"])
        if frm is None or to is None:
            continue
        x1 = boxes[to]["x"] + _r(boxes[to]["w"] / 2)
        y1 = boxes[to]["y"] + boxes[to]["h"]
        x2 = boxes[frm]["x"] + _r(boxes[frm]["w"] / 2)
        y2 = boxes[frm]["y"]
        mid_y = _r((y1 + y2) / 2)
        edges.append({"fk": fk, "path": f"M {x1} {y1} V {mid_y} H {x2} V {y2}",
                      "label": f"{fk['fromColumn']} → {fk['toColumn']}"})
    laid = []
    for i, t in enumerate(schema["tables"]):
        rows = [{"x": boxes[i]["x"], "y": boxes[i]["y"] + o["row_height"] * (1 + ci),
                 "w": boxes[i]["w"], "h": o["row_height"]} for ci in range(len(t["columns"]))]
        laid.append({"table": t, "box": boxes[i],
                     "titleBar": {"x": boxes[i]["x"], "y": boxes[i]["y"], "w": boxes[i]["w"], "h": o["row_height"]},
                     "columnRows": rows})
    return {"width": width, "height": height, "tables": laid, "edges": edges}


def _esc(s):
    return re.sub(r"[&<>\"]", lambda m: {"&": "&amp;", "<": "&lt;", ">": "&gt;", '"': "&quot;"}[m.group()], s)


def render_svg(geo):
    out = [f'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 {geo["width"]} {geo["height"]}" '
           f'class="sv-root" role="img"><title>Schema diagram</title>']
    box_of = {t["table"]["name"]: t["box"] for t in geo["tables"]}
    for e in geo["edges"]:
        frm = box_of.get(e["fk"]["fromTable"])
        if frm is None:
            continue
        ax = frm["x"] + _r(frm["w"] / 2)
        out.append(f'<path class="sv-edge" d="{e["path"]}"/>'
                   f'<polygon class="sv-arrow" points="{ax - 5},{frm["y"] - 8} {ax + 5},{frm["y"] - 8} {ax},{frm["y"]}"/>')
    for t in geo["tables"]:
        b, tb = t["box"], t["titleBar"]
        out.append(f'<g class="sv-table"><rect class="sv-box" x="{b["x"]}" y="{b["y"]}" width="{b["w"]}" height="{b["h"]}" rx="6"/>'
                   f'<rect class="sv-titlebar" x="{tb["x"]}" y="{tb["y"]}" width="{tb["w"]}" height="{tb["h"]}" rx="6"/>'
                   f'<text class="sv-title" x="{b["x"] + 8}" y="{tb["y"] + 17}">{_esc(t["table"]["name"])}</text>')
        for c, row in zip(t["table"]["columns"], t["columnRows"]):
            cls = "sv-pk" if c["isPrimaryKey"] else "sv-col"
            out.append(f'<text class="{cls}" x="{row["x"] + 8}" y="{row["y"] + 17}">{_esc(c["name"])} {_esc(c["type"])}</text>')
        out.append("</g>")
    out.append("</svg>")
    return "".join(out)


def ddl_to_svg(ddl):
    schema = parse_ddl(ddl)
    return {"svg": render_svg(layout_schema(schema)), "schema": schema}


# Example:
#   ddl_to_svg("CREATE TABLE users (id INT PRIMARY KEY);"
#              "CREATE TABLE posts (id INT PRIMARY KEY, user_id INT REFERENCES users(id), title TEXT);")
# → users box on layer 0, posts below, one FK edge — byte-identical to the
#   TS/Go/… ports (integer geometry, same defaults).

Also available in 13 other languages

Every CosmoDev tool ships its pure logic in TypeScript (web) and Go (CLI), with authored implementations in a dozen-plus languages — the same contract, ported. Compare all languages side by side →