Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/sqlalchemy_erd/constants/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
- ``styling``: SVG card visual styling
- ``routing``: orthogonal edge routing
- ``layered``: the layered layout algorithm
- ``relationships``: edge kinds and inheritance strategy labels

Import from the relevant submodule, e.g. ``from sqlalchemy_erd.constants.geometry import NODE_W``.
"""
7 changes: 7 additions & 0 deletions src/sqlalchemy_erd/constants/relationships.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
"""Relationship edge kinds and inheritance strategy labels."""

KIND_FK = "fk"
KIND_INHERITANCE = "inheritance"

STRATEGY_JOINED = "joined"
STRATEGY_CONCRETE = "concrete"
23 changes: 19 additions & 4 deletions src/sqlalchemy_erd/html_renderer.py
Original file line number Diff line number Diff line change
Expand Up @@ -260,6 +260,10 @@ def render_html(
el('polygon', {{ points:'0 0, 10 3.5, 0 7', fill: THEME.edgeColor }}, mk);
const mkHi = el('marker', {{ id:'arr-hi', markerWidth:'10', markerHeight:'7', refX:'9', refY:'3.5', orient:'auto' }}, defs);
el('polygon', {{ points:'0 0, 10 3.5, 0 7', fill: THEME.hiColor }}, mkHi);
const mkInherit = el('marker', {{ id:'inherit', markerWidth:'12', markerHeight:'10', refX:'11', refY:'5', orient:'auto' }}, defs);
el('polygon', {{ points:'1 1, 11 5, 1 9', fill: THEME.bgColor, stroke: THEME.edgeColor, 'stroke-width':'1.4' }}, mkInherit);
const mkInheritHi = el('marker', {{ id:'inherit-hi', markerWidth:'12', markerHeight:'10', refX:'11', refY:'5', orient:'auto' }}, defs);
el('polygon', {{ points:'1 1, 11 5, 1 9', fill: THEME.bgColor, stroke: THEME.hiColor, 'stroke-width':'1.8' }}, mkInheritHi);

el('rect', {{ width:'100%', height:'100%', fill: THEME.bgColor }}, svg);
el('rect', {{ width:'100%', height:'100%', fill:'url(#erd-dots)' }}, svg);
Expand Down Expand Up @@ -294,18 +298,20 @@ def render_html(
}}

const d = makePath(fp, fs, tp, ts);
const isInheritance = rel.kind === 'inheritance';
const isNN = rel.fromCard === 'N' && rel.toCard === 'N';
const isCross = Object.keys(SCHEMA_COLORS).length > 0 &&
(fe.schema != null ? fe.schema : '_default') !== (te.schema != null ? te.schema : '_default');

const g = el('g', {{ 'data-rel-from': rel.from, 'data-rel-to': rel.to }}, svg);
const g = el('g', {{ 'data-rel-from': rel.from, 'data-rel-to': rel.to, 'data-rel-kind': rel.kind }}, svg);
el('path', {{ d, fill:'none', stroke:'transparent', 'stroke-width':'18' }}, g);
const edgeAttrs = {{
class: 'erd-edge', d, fill:'none',
stroke: THEME.edgeColor, 'stroke-width': '1.5',
'marker-end': 'url(#arr)'
'marker-end': isInheritance ? 'url(#inherit)' : 'url(#arr)'
}};
if (isNN) edgeAttrs['stroke-dasharray'] = '5 3';
if (isInheritance) edgeAttrs['stroke-dasharray'] = '3 4';
else if (isNN) edgeAttrs['stroke-dasharray'] = '5 3';
else if (isCross) edgeAttrs['stroke-dasharray'] = '8 4';
el('path', edgeAttrs, g);

Expand All @@ -321,6 +327,13 @@ def render_html(
fill: THEME.edgeColor, 'text-anchor':'middle', 'dominant-baseline':'middle',
textContent: rel.toCard
}}, g);
if (rel.label) {{
el('text', {{
class: 'erd-label', x: (fp[0] + tp[0]) / 2, y: (fp[1] + tp[1]) / 2 - 8,
'font-size':'10', 'font-family':'monospace', fill: THEME.edgeColor,
'text-anchor':'middle', 'dominant-baseline':'middle', textContent: rel.label
}}, g);
}}
}}

// Nodes — tagged with data attributes for updateHighlights
Expand Down Expand Up @@ -397,11 +410,13 @@ def render_html(
const from = g.getAttribute('data-rel-from');
const to = g.getAttribute('data-rel-to');
const hi = hoveredId && (from === hoveredId || to === hoveredId);
const kind = g.getAttribute('data-rel-kind');
const edge = g.querySelector('.erd-edge');
if (edge) {{
edge.setAttribute('stroke', hi ? THEME.hiColor : THEME.edgeColor);
edge.setAttribute('stroke-width', hi ? '2' : '1.5');
edge.setAttribute('marker-end', `url(#${{hi ? 'arr-hi' : 'arr'}})`);
const marker = kind === 'inheritance' ? (hi ? 'inherit-hi' : 'inherit') : (hi ? 'arr-hi' : 'arr');
edge.setAttribute('marker-end', `url(#${{marker}})`);
}}
g.querySelectorAll('.erd-label').forEach(lbl => {{
lbl.setAttribute('fill', hi ? THEME.hiColor : THEME.edgeColor);
Expand Down
81 changes: 74 additions & 7 deletions src/sqlalchemy_erd/introspect.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,13 @@
Interval, JSON, LargeBinary, MetaData, Numeric, SmallInteger, String,
Text, Time, UniqueConstraint, Uuid,
)
from sqlalchemy.orm import DeclarativeBase
from sqlalchemy.orm import DeclarativeBase, Mapper
from sqlalchemy.types import TypeDecorator

from sqlalchemy_erd.constants.relationships import (
KIND_FK, KIND_INHERITANCE, STRATEGY_CONCRETE, STRATEGY_JOINED,
)


@dataclass
class ColumnInfo:
Expand All @@ -36,6 +40,8 @@ class RelationshipInfo:
from_card: str
to_card: str
fk_column: str
kind: str = KIND_FK
label: str | None = None


@dataclass(frozen=True)
Expand Down Expand Up @@ -113,17 +119,18 @@ def _column_kind(col_info: Any, is_pk: bool, is_fk: bool) -> str:

def _resolve_metadata(
base_or_metadata: type[DeclarativeBase] | MetaData,
) -> tuple[MetaData, dict[str, str]]:
"""Return the metadata and a table-fullname → mapped-class-name lookup."""
) -> tuple[MetaData, dict[str, str], list[Mapper]]:
"""Return metadata plus mapper information when a DeclarativeBase is given."""
if isinstance(base_or_metadata, MetaData):
return base_or_metadata, {}
return base_or_metadata, {}, []

metadata = base_or_metadata.metadata
mappers = list(base_or_metadata.registry.mappers)
class_names = {
mapper.local_table.fullname: mapper.class_.__name__
for mapper in base_or_metadata.registry.mappers
for mapper in mappers
}
return metadata, class_names
return metadata, class_names, mappers


def _build_table(
Expand Down Expand Up @@ -208,6 +215,55 @@ def _build_relationships(
return relationships


def _inheritance_strategy(mapper: Mapper) -> str:
"""Name the strategy of a child mapper that owns its own table."""
return STRATEGY_CONCRETE if mapper.concrete else STRATEGY_JOINED


def _build_inheritance_relationships(
mappers: list[Mapper],
kept_names: set[str],
) -> list[RelationshipInfo]:
relationships: list[RelationshipInfo] = []
for mapper in sorted(mappers, key=lambda m: m.local_table.fullname):
if mapper.inherits is None:
continue
parent_table = mapper.inherits.local_table
child_table = mapper.local_table
parent_name = parent_table.fullname
child_name = child_table.fullname
if parent_name == child_name:
continue
if parent_name not in kept_names or child_name not in kept_names:
continue

# The join column is the child PK that also references the parent.
fk_col = ""
pk_names = {col.name for col in child_table.primary_key.columns}
for col in child_table.columns:
references_parent = any(
fk.column.table.fullname == parent_name for fk in col.foreign_keys
)
if references_parent and col.name in pk_names:
fk_col = col.name
break
if not fk_col:
pk_cols = list(child_table.primary_key.columns)
fk_col = pk_cols[0].name if pk_cols else ""

strategy = _inheritance_strategy(mapper)
relationships.append(RelationshipInfo(
from_table=parent_name,
to_table=child_name,
from_card="1",
to_card="1",
fk_column=fk_col,
kind=KIND_INHERITANCE,
label=strategy,
))
return relationships


def _collapse_association_tables(
filtered_items: list[tuple[str, Any]],
tables: list[TableInfo],
Expand Down Expand Up @@ -256,7 +312,7 @@ def introspect_models(
schemas: list[str] | None = None,
filters: Filters | None = None,
) -> tuple[list[TableInfo], list[RelationshipInfo]]:
metadata, class_names = _resolve_metadata(base_or_metadata)
metadata, class_names, mappers = _resolve_metadata(base_or_metadata)
filters = filters or Filters()
include = _compile(filters.include_tables)
exclude = _compile(filters.exclude_tables)
Expand All @@ -280,6 +336,17 @@ def introspect_models(
filtered_items, tables, relationships,
)
kept_names = {t.name for t in tables}
inheritance_relationships = _build_inheritance_relationships(mappers, kept_names)
# Drop only the FK edge the inheritance edge replaces.
inheritance_edges = {
(rel.from_table, rel.to_table, rel.fk_column)
for rel in inheritance_relationships
}
relationships = [
rel for rel in relationships
if (rel.from_table, rel.to_table, rel.fk_column) not in inheritance_edges
]
relationships.extend(inheritance_relationships)
relationships = [
rel for rel in relationships
if rel.from_table in kept_names and rel.to_table in kept_names
Expand Down
25 changes: 22 additions & 3 deletions src/sqlalchemy_erd/renderer.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from xml.sax.saxutils import escape

from sqlalchemy_erd.constants.geometry import FIELD_H, HEADER_H, NODE_W, PAD
from sqlalchemy_erd.constants.relationships import KIND_INHERITANCE
from sqlalchemy_erd.constants.styling import (
CARD_RADIUS, FIELD_FONT_SIZE, FIELD_INSET, HEADER_STRIP_H, KIND_FONT_SIZE,
KIND_INSET, TITLE_FONT_SIZE, TITLE_INSET,
Expand Down Expand Up @@ -126,6 +127,9 @@ def render_svg(
<marker id="arr-hi" markerWidth="10" markerHeight="7" refX="9" refY="3.5" orient="auto">
<polygon points="0 0, 10 3.5, 0 7" fill="{theme.highlight_color}" />
</marker>
<marker id="inherit" markerWidth="12" markerHeight="10" refX="11" refY="5" orient="auto">
<polygon points="1 1, 11 5, 1 9" fill="{theme.bg_color}" stroke="{theme.edge_color}" stroke-width="1.4" />
</marker>
</defs>""")

parts.append(f' <rect width="100%" height="100%" fill="{theme.bg_color}" />')
Expand Down Expand Up @@ -164,9 +168,13 @@ def render_svg(
tpt = _conn_pt(tp, tt, ts, to_idx, node_w)

path_d = orthogonal_path(fpt, fs, tpt, ts)
is_inheritance = rel.kind == KIND_INHERITANCE
is_nn = rel.from_card == "N" and rel.to_card == "N"
is_cross = multi_schema and ft.schema != tt.schema
if is_nn:
marker = "inherit" if is_inheritance else "arr"
if is_inheritance:
dash = ' stroke-dasharray="3 4"'
elif is_nn:
dash = ' stroke-dasharray="5 3"'
elif is_cross:
dash = ' stroke-dasharray="8 4"'
Expand All @@ -176,11 +184,14 @@ def render_svg(
fl = _label_pos(fpt, fs)
tl = _label_pos(tpt, ts)

parts.append(f' <g class="erd-rel" data-from="{rel.from_table}" data-to="{rel.to_table}">')
parts.append(
f' <g class="erd-rel" data-from="{rel.from_table}" data-to="{rel.to_table}" '
f'data-kind="{rel.kind}">'
)
parts.append(f' <path d="{path_d}" fill="none" stroke="transparent" stroke-width="18" />')
parts.append(
f' <path class="erd-edge" d="{path_d}" fill="none" '
f'stroke="{theme.edge_color}" stroke-width="1.5"{dash} marker-end="url(#arr)" />'
f'stroke="{theme.edge_color}" stroke-width="1.5"{dash} marker-end="url(#{marker})" />'
)
parts.append(
f' <text x="{fl[0]}" y="{fl[1]}" font-size="11" '
Expand All @@ -192,6 +203,14 @@ def render_svg(
f'font-family="monospace" fill="{theme.edge_color}" '
f'text-anchor="middle" dominant-baseline="middle">{rel.to_card}</text>'
)
if rel.label:
mx = (fpt[0] + tpt[0]) / 2
my = (fpt[1] + tpt[1]) / 2 - 8
parts.append(
f' <text x="{mx}" y="{my}" font-size="10" '
f'font-family="monospace" fill="{theme.edge_color}" '
f'text-anchor="middle" dominant-baseline="middle">{escape(rel.label)}</text>'
)
parts.append(" </g>")

for table in tables:
Expand Down
2 changes: 2 additions & 0 deletions src/sqlalchemy_erd/serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,8 @@ def build_relations_json(
"fromCard": r.from_card,
"toCard": r.to_card,
"fkCol": r.fk_column,
"kind": r.kind,
"label": r.label,
}
for r in relationships
if r.from_table in table_names and r.to_table in table_names
Expand Down
61 changes: 10 additions & 51 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
"""Shared fixtures for the ERDAlchemy test suite."""
"""Shared fixtures for the ERDAlchemy test suite.

Per-domain schemas live in ``tests/fixtures`` and are loaded as plugins.
"""

import pytest
from sqlalchemy import (
Expand All @@ -10,6 +13,12 @@
from datetime import datetime, timedelta, time
from decimal import Decimal

pytest_plugins = [
"tests.fixtures.inheritance",
"tests.fixtures.cardinality",
"tests.fixtures.comments",
]


# ── Single-table schema ──────────────────────────────────────────────────────

Expand Down Expand Up @@ -273,53 +282,3 @@ class Task(MultiFkBase):
@pytest.fixture
def multi_fk_base():
return MultiFkBase


# -- Cardinality schema -------------------------------------------------------

cardinality_metadata = MetaData()

Table(
"users", cardinality_metadata,
Column("id", Integer, primary_key=True),
Column("name", String(100), nullable=False),
)

Table(
"profiles", cardinality_metadata,
Column("user_id", Integer, ForeignKey("users.id"), primary_key=True),
Column("bio", Text),
)

Table(
"avatars", cardinality_metadata,
Column("id", Integer, primary_key=True),
Column("user_id", Integer, ForeignKey("users.id"), unique=True, nullable=False),
)

Table(
"tasks", cardinality_metadata,
Column("id", Integer, primary_key=True),
Column("assignee_id", Integer, ForeignKey("users.id"), nullable=True),
)


@pytest.fixture
def cardinality_metadata_fixture():
return cardinality_metadata


# -- Column comments schema ---------------------------------------------------

comments_metadata = MetaData()

Table(
"accounts", comments_metadata,
Column("id", Integer, primary_key=True),
Column("email", String(200), comment="Primary login email"),
)


@pytest.fixture
def comments_metadata_fixture():
return comments_metadata
Empty file added tests/fixtures/__init__.py
Empty file.
38 changes: 38 additions & 0 deletions tests/fixtures/cardinality.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
"""Fixtures for relationship cardinality schemas."""

import pytest
from sqlalchemy import Column, ForeignKey, Integer, MetaData, String, Table, Text


# -- Cardinality schema -------------------------------------------------------

cardinality_metadata = MetaData()

Table(
"users", cardinality_metadata,
Column("id", Integer, primary_key=True),
Column("name", String(100), nullable=False),
)

Table(
"profiles", cardinality_metadata,
Column("user_id", Integer, ForeignKey("users.id"), primary_key=True),
Column("bio", Text),
)

Table(
"avatars", cardinality_metadata,
Column("id", Integer, primary_key=True),
Column("user_id", Integer, ForeignKey("users.id"), unique=True, nullable=False),
)

Table(
"tasks", cardinality_metadata,
Column("id", Integer, primary_key=True),
Column("assignee_id", Integer, ForeignKey("users.id"), nullable=True),
)


@pytest.fixture
def cardinality_metadata_fixture():
return cardinality_metadata
Loading
Loading