From 4f68a02c37be28b3564d25779515050f1d9c18fb Mon Sep 17 00:00:00 2001 From: zamax14 Date: Mon, 29 Jun 2026 02:16:10 -0600 Subject: [PATCH 1/2] feat: render inheritance relationships --- src/sqlalchemy_erd/html_renderer.py | 23 +++++++-- src/sqlalchemy_erd/introspect.py | 73 ++++++++++++++++++++++++++--- src/sqlalchemy_erd/renderer.py | 24 ++++++++-- src/sqlalchemy_erd/serialization.py | 2 + tests/conftest.py | 29 ++++++++++++ tests/test_html_renderer.py | 9 +++- tests/test_introspect.py | 20 ++++++++ tests/test_renderer.py | 17 +++++++ 8 files changed, 182 insertions(+), 15 deletions(-) diff --git a/src/sqlalchemy_erd/html_renderer.py b/src/sqlalchemy_erd/html_renderer.py index eb507a7..213d117 100644 --- a/src/sqlalchemy_erd/html_renderer.py +++ b/src/sqlalchemy_erd/html_renderer.py @@ -259,6 +259,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); @@ -293,18 +297,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); @@ -320,6 +326,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 @@ -390,11 +403,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); diff --git a/src/sqlalchemy_erd/introspect.py b/src/sqlalchemy_erd/introspect.py index a98fae0..9e6c8b0 100644 --- a/src/sqlalchemy_erd/introspect.py +++ b/src/sqlalchemy_erd/introspect.py @@ -7,7 +7,7 @@ Interval, JSON, LargeBinary, MetaData, Numeric, SmallInteger, String, Text, Time, Uuid, ) -from sqlalchemy.orm import DeclarativeBase +from sqlalchemy.orm import DeclarativeBase, Mapper from sqlalchemy.types import TypeDecorator @@ -35,6 +35,8 @@ class RelationshipInfo: from_card: str to_card: str fk_column: str + kind: str = "fk" + label: str | None = None @dataclass(frozen=True) @@ -112,17 +114,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( @@ -188,6 +191,53 @@ def _build_relationships( return relationships +def _inheritance_strategy(mapper: Mapper) -> str: + if mapper.concrete: + return "concrete" + if mapper.inherits is not None and mapper.local_table is mapper.inherits.local_table: + return "single" + return "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 + + fk_col = "" + for col in child_table.columns: + if any(fk.column.table.fullname == parent_name for fk in col.foreign_keys): + 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="inheritance", + label=strategy, + )) + return relationships + + def _collapse_association_tables( filtered_items: list[tuple[str, Any]], tables: list[TableInfo], @@ -236,7 +286,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) @@ -260,6 +310,15 @@ def introspect_models( filtered_items, tables, relationships, ) kept_names = {t.name for t in tables} + inheritance_relationships = _build_inheritance_relationships(mappers, kept_names) + inheritance_pairs = { + (rel.from_table, rel.to_table) for rel in inheritance_relationships + } + relationships = [ + rel for rel in relationships + if (rel.from_table, rel.to_table) not in inheritance_pairs + ] + relationships.extend(inheritance_relationships) relationships = [ rel for rel in relationships if rel.from_table in kept_names and rel.to_table in kept_names diff --git a/src/sqlalchemy_erd/renderer.py b/src/sqlalchemy_erd/renderer.py index cfb6a7f..edde993 100644 --- a/src/sqlalchemy_erd/renderer.py +++ b/src/sqlalchemy_erd/renderer.py @@ -126,6 +126,9 @@ def render_svg( + + + """) parts.append(f' ') @@ -164,9 +167,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 == "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"' @@ -176,11 +183,14 @@ def render_svg( fl = _label_pos(fpt, fs) tl = _label_pos(tpt, ts) - parts.append(f' ') + parts.append( + f' ' + ) parts.append(f' ') parts.append( f' ' + f'stroke="{theme.edge_color}" stroke-width="1.5"{dash} marker-end="url(#{marker})" />' ) parts.append( f' {rel.to_card}' ) + if rel.label: + mx = (fpt[0] + tpt[0]) / 2 + my = (fpt[1] + tpt[1]) / 2 - 8 + parts.append( + f' {escape(rel.label)}' + ) parts.append(" ") for table in tables: diff --git a/src/sqlalchemy_erd/serialization.py b/src/sqlalchemy_erd/serialization.py index 5d9d490..b02e1b7 100644 --- a/src/sqlalchemy_erd/serialization.py +++ b/src/sqlalchemy_erd/serialization.py @@ -60,6 +60,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 diff --git a/tests/conftest.py b/tests/conftest.py index a9500bb..a8dc493 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -273,3 +273,32 @@ class Task(MultiFkBase): @pytest.fixture def multi_fk_base(): return MultiFkBase + + +# -- Joined-table inheritance schema ----------------------------------------- + +class InheritanceBase(DeclarativeBase): + pass + + +class Employee(InheritanceBase): + __tablename__ = "employees" + id: Mapped[int] = mapped_column(primary_key=True) + kind: Mapped[str] = mapped_column(String(50)) + name: Mapped[str] = mapped_column(String(100)) + __mapper_args__ = { + "polymorphic_on": kind, + "polymorphic_identity": "employee", + } + + +class Manager(Employee): + __tablename__ = "managers" + id: Mapped[int] = mapped_column(ForeignKey("employees.id"), primary_key=True) + department: Mapped[str] = mapped_column(String(100)) + __mapper_args__ = {"polymorphic_identity": "manager"} + + +@pytest.fixture +def inheritance_base(): + return InheritanceBase diff --git a/tests/test_html_renderer.py b/tests/test_html_renderer.py index 1e4397d..41e8b69 100644 --- a/tests/test_html_renderer.py +++ b/tests/test_html_renderer.py @@ -109,7 +109,7 @@ def test_relation_has_expected_keys(self, blog_base): tables, rels = introspect_models(blog_base) positions = force_directed_layout(tables, rels) parsed = json.loads(_build_relations_json(rels, tables, positions)) - assert set(parsed[0]) == {"from", "to", "fromCard", "toCard", "fkCol"} + assert set(parsed[0]) == {"from", "to", "fromCard", "toCard", "fkCol", "kind", "label"} def test_via_relation_keeps_via_prefix(self, m2m_base): tables, rels = introspect_models(m2m_base) @@ -125,6 +125,13 @@ def test_relation_to_unknown_table_is_filtered_out(self, blog_base): parsed = json.loads(_build_relations_json(rels, kept, positions)) assert all(r["from"] != "comments" and r["to"] != "comments" for r in parsed) + def test_inheritance_relation_serializes_kind_and_label(self, inheritance_base): + tables, rels = introspect_models(inheritance_base) + positions = force_directed_layout(tables, rels) + parsed = json.loads(_build_relations_json(rels, tables, positions)) + rel = next(r for r in parsed if r["kind"] == "inheritance") + assert rel["label"] == "joined" + # ── render_html ────────────────────────────────────────────────────────────── diff --git a/tests/test_introspect.py b/tests/test_introspect.py index ff1c65d..9230e4b 100644 --- a/tests/test_introspect.py +++ b/tests/test_introspect.py @@ -354,3 +354,23 @@ def test_schema_attribute_set(self, multi_schema_metadata_fixture): tables, _ = introspect_models(multi_schema_metadata_fixture) schemas = {t.schema for t in tables} assert schemas == {"auth", "billing"} + + +# -- SQLAlchemy inheritance --------------------------------------------------- + +class TestIntrospectInheritance: + def test_joined_inheritance_edge_is_distinct(self, inheritance_base): + _, rels = introspect_models(inheritance_base) + inheritance = [r for r in rels if r.kind == "inheritance"] + assert len(inheritance) == 1 + rel = inheritance[0] + assert rel.from_table == "employees" + assert rel.to_table == "managers" + assert rel.from_card == "1" + assert rel.to_card == "1" + assert rel.label == "joined" + + def test_joined_inheritance_does_not_duplicate_fk_edge(self, inheritance_base): + _, rels = introspect_models(inheritance_base) + pairs = [(r.from_table, r.to_table) for r in rels] + assert pairs.count(("employees", "managers")) == 1 diff --git a/tests/test_renderer.py b/tests/test_renderer.py index 0e7d7cd..410284b 100644 --- a/tests/test_renderer.py +++ b/tests/test_renderer.py @@ -163,3 +163,20 @@ def test_custom_table_colors(self, blog_base): theme = get_theme("default", table_colors={"users": "#ff0000"}) svg = render_svg(tables, rels, positions, theme) assert "#ff0000" in svg + + +# -- Inheritance rendering ---------------------------------------------------- + +class TestInheritanceRendering: + def test_inheritance_edge_has_distinct_kind_and_marker(self, inheritance_base): + tables, rels = introspect_models(inheritance_base) + positions = force_directed_layout(tables, rels) + svg = render_svg(tables, rels, positions, get_theme("default")) + root = ET.fromstring(svg) + rel_groups = [g for g in root.findall("svg:g", NS) if g.get("data-kind") == "inheritance"] + assert len(rel_groups) == 1 + edge = rel_groups[0].find("svg:path[@class='erd-edge']", NS) + assert edge.get("marker-end") == "url(#inherit)" + assert edge.get("stroke-dasharray") == "3 4" + text = " ".join(t.text or "" for t in root.iter("{http://www.w3.org/2000/svg}text")) + assert "joined" in text From 58256918628dd1cd0d8e01f854431ed599315312 Mon Sep 17 00:00:00 2001 From: JoseVelazcoH Date: Mon, 17 Aug 2026 10:35:21 -0600 Subject: [PATCH 2/2] fix(introspect): keep extra foreign keys to an inherited parent Deduplicating FK edges by (from_table, to_table) removed every foreign key between a child and its inherited parent, not just the join column, so business relationships to the same parent silently vanished from the diagram. Match on the fk_column too, and resolve the join column as the child primary key that references the parent. Drop the unreachable "single" strategy branch: single-table inheritance is skipped earlier because parent and child share a table. Move the edge-kind and strategy literals into constants/relationships.py, and split the inheritance and cardinality schemas out of conftest.py. --- src/sqlalchemy_erd/constants/__init__.py | 1 + src/sqlalchemy_erd/constants/relationships.py | 7 ++ src/sqlalchemy_erd/introspect.py | 30 +++++--- src/sqlalchemy_erd/renderer.py | 3 +- tests/conftest.py | 72 +++---------------- tests/fixtures/__init__.py | 0 tests/fixtures/cardinality.py | 38 ++++++++++ tests/fixtures/inheritance.py | 66 +++++++++++++++++ tests/test_introspect.py | 7 ++ 9 files changed, 149 insertions(+), 75 deletions(-) create mode 100644 src/sqlalchemy_erd/constants/relationships.py create mode 100644 tests/fixtures/__init__.py create mode 100644 tests/fixtures/cardinality.py create mode 100644 tests/fixtures/inheritance.py diff --git a/src/sqlalchemy_erd/constants/__init__.py b/src/sqlalchemy_erd/constants/__init__.py index ef581c3..e1845b4 100644 --- a/src/sqlalchemy_erd/constants/__init__.py +++ b/src/sqlalchemy_erd/constants/__init__.py @@ -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``. """ diff --git a/src/sqlalchemy_erd/constants/relationships.py b/src/sqlalchemy_erd/constants/relationships.py new file mode 100644 index 0000000..b9bad5a --- /dev/null +++ b/src/sqlalchemy_erd/constants/relationships.py @@ -0,0 +1,7 @@ +"""Relationship edge kinds and inheritance strategy labels.""" + +KIND_FK = "fk" +KIND_INHERITANCE = "inheritance" + +STRATEGY_JOINED = "joined" +STRATEGY_CONCRETE = "concrete" diff --git a/src/sqlalchemy_erd/introspect.py b/src/sqlalchemy_erd/introspect.py index 9e1332c..cca86dc 100644 --- a/src/sqlalchemy_erd/introspect.py +++ b/src/sqlalchemy_erd/introspect.py @@ -10,6 +10,10 @@ 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: @@ -35,7 +39,7 @@ class RelationshipInfo: from_card: str to_card: str fk_column: str - kind: str = "fk" + kind: str = KIND_FK label: str | None = None @@ -210,11 +214,8 @@ def _build_relationships( def _inheritance_strategy(mapper: Mapper) -> str: - if mapper.concrete: - return "concrete" - if mapper.inherits is not None and mapper.local_table is mapper.inherits.local_table: - return "single" - return "joined" + """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( @@ -234,9 +235,14 @@ def _build_inheritance_relationships( 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: - if any(fk.column.table.fullname == parent_name for fk in col.foreign_keys): + 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: @@ -250,7 +256,7 @@ def _build_inheritance_relationships( from_card="1", to_card="1", fk_column=fk_col, - kind="inheritance", + kind=KIND_INHERITANCE, label=strategy, )) return relationships @@ -329,12 +335,14 @@ def introspect_models( ) kept_names = {t.name for t in tables} inheritance_relationships = _build_inheritance_relationships(mappers, kept_names) - inheritance_pairs = { - (rel.from_table, rel.to_table) for rel in inheritance_relationships + # 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) not in inheritance_pairs + if (rel.from_table, rel.to_table, rel.fk_column) not in inheritance_edges ] relationships.extend(inheritance_relationships) relationships = [ diff --git a/src/sqlalchemy_erd/renderer.py b/src/sqlalchemy_erd/renderer.py index edde993..5858cae 100644 --- a/src/sqlalchemy_erd/renderer.py +++ b/src/sqlalchemy_erd/renderer.py @@ -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, @@ -167,7 +168,7 @@ 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 == "inheritance" + 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 marker = "inherit" if is_inheritance else "arr" diff --git a/tests/conftest.py b/tests/conftest.py index 5c8470d..c62f808 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 ( @@ -10,6 +13,11 @@ from datetime import datetime, timedelta, time from decimal import Decimal +pytest_plugins = [ + "tests.fixtures.inheritance", + "tests.fixtures.cardinality", +] + # ── Single-table schema ────────────────────────────────────────────────────── @@ -274,65 +282,3 @@ class Task(MultiFkBase): def multi_fk_base(): return MultiFkBase - -# -- Joined-table inheritance schema ----------------------------------------- - -class InheritanceBase(DeclarativeBase): - pass - - -class Employee(InheritanceBase): - __tablename__ = "employees" - id: Mapped[int] = mapped_column(primary_key=True) - kind: Mapped[str] = mapped_column(String(50)) - name: Mapped[str] = mapped_column(String(100)) - __mapper_args__ = { - "polymorphic_on": kind, - "polymorphic_identity": "employee", - } - - -class Manager(Employee): - __tablename__ = "managers" - id: Mapped[int] = mapped_column(ForeignKey("employees.id"), primary_key=True) - department: Mapped[str] = mapped_column(String(100)) - __mapper_args__ = {"polymorphic_identity": "manager"} - - -@pytest.fixture -def inheritance_base(): - return InheritanceBase - - -# -- 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 diff --git a/tests/fixtures/__init__.py b/tests/fixtures/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/fixtures/cardinality.py b/tests/fixtures/cardinality.py new file mode 100644 index 0000000..570874e --- /dev/null +++ b/tests/fixtures/cardinality.py @@ -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 diff --git a/tests/fixtures/inheritance.py b/tests/fixtures/inheritance.py new file mode 100644 index 0000000..7fa912d --- /dev/null +++ b/tests/fixtures/inheritance.py @@ -0,0 +1,66 @@ +"""Fixtures for SQLAlchemy inheritance schemas.""" + +import pytest +from sqlalchemy import ForeignKey, String +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column + + +# -- Joined-table inheritance schema ----------------------------------------- + +class InheritanceBase(DeclarativeBase): + pass + + +class Employee(InheritanceBase): + __tablename__ = "employees" + id: Mapped[int] = mapped_column(primary_key=True) + kind: Mapped[str] = mapped_column(String(50)) + name: Mapped[str] = mapped_column(String(100)) + __mapper_args__ = { + "polymorphic_on": kind, + "polymorphic_identity": "employee", + } + + +class Manager(Employee): + __tablename__ = "managers" + id: Mapped[int] = mapped_column(ForeignKey("employees.id"), primary_key=True) + department: Mapped[str] = mapped_column(String(100)) + __mapper_args__ = {"polymorphic_identity": "manager"} + + +@pytest.fixture +def inheritance_base(): + return InheritanceBase + + +# -- Inheritance plus an extra FK to the same parent -------------------------- + +class InheritanceExtraFkBase(DeclarativeBase): + pass + + +class Staff(InheritanceExtraFkBase): + __tablename__ = "staff" + id: Mapped[int] = mapped_column(primary_key=True) + kind: Mapped[str] = mapped_column(String(50)) + __mapper_args__ = { + "polymorphic_on": kind, + "polymorphic_identity": "staff", + } + + +class Lead(Staff): + __tablename__ = "leads" + id: Mapped[int] = mapped_column(ForeignKey("staff.id"), primary_key=True) + mentor_id: Mapped[int] = mapped_column(ForeignKey("staff.id")) + __mapper_args__ = { + "polymorphic_identity": "lead", + "inherit_condition": id == Staff.id, + } + + +@pytest.fixture +def inheritance_extra_fk_base(): + return InheritanceExtraFkBase + diff --git a/tests/test_introspect.py b/tests/test_introspect.py index 633f857..fff7688 100644 --- a/tests/test_introspect.py +++ b/tests/test_introspect.py @@ -372,6 +372,13 @@ def test_joined_inheritance_does_not_duplicate_fk_edge(self, inheritance_base): pairs = [(r.from_table, r.to_table) for r in rels] assert pairs.count(("employees", "managers")) == 1 + def test_extra_fk_to_parent_survives_inheritance_edge( + self, inheritance_extra_fk_base, + ): + _, rels = introspect_models(inheritance_extra_fk_base) + edges = {(r.kind, r.fk_column) for r in rels if r.to_table == "leads"} + assert edges == {("inheritance", "id"), ("fk", "mentor_id")} + # -- Relationship cardinality -------------------------------------------------