SQLAlchemy mapping many-to-many join table to list of keys

Viewed 25

I have two tables with a many to many relationship:

purchase_order = Table(
    "purchase_order",
    metadata,
    Column("name", String, nullable=False, primary_key=True),
    ...
)

shipment = Table(
    "shipment",
    metadata,
    Column("identifier", String, nullable=False, primary_key=True),
)

shipment_po = Table(
    "shipment_po",
    metadata,
    Column(
        "shipment_identifier",
        String,
        ForeignKey("shipment.identifier"),
        nullable=False,
        primary_key=True,
    ),
    Column(
        "purchase_order_name",
        String,
        ForeignKey("purchase_order.name"),
        nullable=False,
        primary_key=True,
    ),
)

I'd like to map them to a class like this:

class Shipment:
    def __init__(
        self,
        identifier: str,
        purchase_order_names: list[str],
    ):
    ...

(Note that purchase_order_names is just a list of str, not a list of some PurchaseOrder object.)

So that I can create them like this (it being my responsibility to ensure the purchase_order rows here exist):

session.add(Shipment("my_shipment", purchase_order_names=["my_po_1", "my_po_2"]))

fetch like this:

session.query(Shipment).filter_by(identifier="my_shipment").one()

and update like this:

shipment.purchase_order_names.append("my_po_3")

Is this supported?

1 Answers

You can do this by adding a validator to Shipment that creates a PurchaseOrder from the provided name:

from sqlalchemy import orm

class Shipment:
    ...
    @orm.validates('purchase_orders')
    def validate_purchase_orders(self, key, purchase_order):
        return PurchaseOrder(name=purchase_order)

Here is a complete example, based on the code in the question:

import sqlalchemy as sa
from sqlalchemy import orm

mapper_registry = orm.registry()
metadata = mapper_registry.metadata

purchase_order = sa.Table(
    'purchase_order',
    metadata,
    sa.Column('name', sa.String, nullable=False, primary_key=True),
)

shipment = sa.Table(
    'shipment',
    metadata,
    sa.Column('identifier', sa.String, nullable=False, primary_key=True),
)

shipment_po = sa.Table(
    'shipment_po',
    metadata,
    sa.Column(
        'shipment_identifier',
        sa.String,
        sa.ForeignKey('shipment.identifier'),
        nullable=False,
        primary_key=True,
    ),
    sa.Column(
        'purchase_order_name',
        sa.String,
        sa.ForeignKey('purchase_order.name'),
        nullable=False,
        primary_key=True,
    ),
)


class Shipment:
    def __init__(
        self,
        identifier: str,
        purchase_order_names: list[str],
    ):
        self.identifier = identifier
        self.purchase_orders = purchase_order_names

    @orm.validates('purchase_orders')
    def validate_purchase_orders(self, key, purchase_order):
        return PurchaseOrder(name=purchase_order)


class PurchaseOrder:
    pass


mapper_registry.map_imperatively(Shipment, shipment)
mapper_registry.map_imperatively(
    PurchaseOrder,
    purchase_order,
    properties={
        'shipments': orm.relationship(
            'Shipment', secondary=shipment_po, backref='purchase_orders'
        )
    },
)

engine = sa.create_engine('sqlite://', echo=True, future=True)
mapper_registry.metadata.create_all(engine)
Session = orm.sessionmaker(engine, future=True)

with Session.begin() as s:
    s.add(Shipment("my_shipment", purchase_order_names=["my_po_1", "my_po_2"]))

with Session() as s:
    shipment = s.scalars(sa.select(Shipment).filter_by(identifier="my_shipment")).one()
    assert len(shipment.purchase_orders) == 2
    assert all(isinstance(po, PurchaseOrder) for po in shipment.purchase_orders)
    assert sorted(po.name for po in shipment.purchase_orders) == ["my_po_1", "my_po_2"]

    shipment.purchase_orders.append("my_po_3")
    s.commit()

with Session() as s:
    shipment = s.scalars(sa.select(Shipment).filter_by(identifier="my_shipment")).one()
    assert len(shipment.purchase_orders) == 3
    assert all(isinstance(po, PurchaseOrder) for po in shipment.purchase_orders)
    assert sorted(po.name for po in shipment.purchase_orders) == ["my_po_1", "my_po_2", "my_po_3"]

If you require more control you could create a custom collection class.

Related