Skip to content

Commit 0b3ef0d

Browse files
committed
Add polymorphic field support
1 parent 00413ee commit 0b3ef0d

16 files changed

Lines changed: 605 additions & 290 deletions

CHANGELOG.md

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,9 @@
1+
# v0.4.5
2+
3+
- Added support for polymorphic foreign keys
4+
- Removed Python 3.8, 3.9 support and added 3.13, 3.14 support
5+
- Updated dependencies
6+
17
# v0.4.4
28

39
- Improved query performance following foreign key relationships

subsetter/_version.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
__version__ = "0.4.4"
1+
__version__ = "0.4.5"

subsetter/common.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -241,7 +241,7 @@ def _push(key: Any, value: Any):
241241
data = stack.pop()
242242
if isinstance(data, BaseModel):
243243
yield data
244-
for field, _ in data.model_fields.items():
244+
for field, _ in data.__class__.model_fields.items():
245245
_push(field, getattr(data, field))
246246

247247
if isinstance(data, list):

subsetter/config_model.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,32 @@ def check_columns_match(self):
4343
raise ValueError("each column in src_columns must be unique")
4444
return self
4545

46+
class PolymorphicFKConfig(ForbidBaseModel):
47+
class KeyDestination(ForbidBaseModel):
48+
table: str
49+
columns: List[str]
50+
51+
table: str
52+
columns: List[str]
53+
discriminator_column: str
54+
destinations: dict[str, KeyDestination]
55+
56+
@model_validator(mode="after")
57+
def check_columns_match(self):
58+
col_count = len(self.columns)
59+
if not col_count:
60+
raise ValueError("columns cannot be empty")
61+
if len(set(self.columns)) != col_count:
62+
raise ValueError("each column in columns must be unique")
63+
for key_dest in self.destinations.values():
64+
if len(key_dest.columns) != col_count:
65+
raise ValueError(
66+
"src_columns and dst_columns must be the same length"
67+
)
68+
if len(set(key_dest.columns)) != col_count:
69+
raise ValueError("each column in src_columns must be unique")
70+
return self
71+
4672
class ColumnConstraint(ForbidBaseModel):
4773
column: str
4874
operator: SQLKnownOperator
@@ -55,6 +81,7 @@ class ColumnConstraint(ForbidBaseModel):
5581
passthrough: List[str] = []
5682
ignore_fks: List[IgnoreFKConfig] = []
5783
extra_fks: List[ExtraFKConfig] = []
84+
polymorphic_fks: List[PolymorphicFKConfig] = []
5885
infer_foreign_keys: Literal["none", "schema", "all"] = "none"
5986
include_dependencies: bool = True
6087

subsetter/metadata.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,8 @@ class ForeignKey:
2222
dst_schema: str
2323
dst_table: str
2424
dst_columns: Tuple[str, ...]
25+
src_discriminator: Optional[Tuple[str, str]] = None
26+
dst_discriminator: Optional[Tuple[str, str]] = None
2527

2628
@classmethod
2729
def from_schema(cls, fk: sa.ForeignKeyConstraint) -> "ForeignKey":
@@ -188,6 +190,8 @@ def compute_reverse_keys(self) -> None:
188190
dst_schema=table.schema,
189191
dst_table=table.name,
190192
dst_columns=fk.columns,
193+
src_discriminator=fk.dst_discriminator,
194+
dst_discriminator=fk.src_discriminator,
191195
)
192196
)
193197

subsetter/plan_model.py

Lines changed: 37 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -237,6 +237,8 @@ class SQLLeftJoin(BaseModel):
237237
left_columns: List[str]
238238
right_columns: List[str]
239239
half_unique: bool = True
240+
left_discriminator: List[str] = []
241+
right_discriminator: List[str] = []
240242

241243

242244
class SQLStatementSelect(BaseModel):
@@ -245,8 +247,9 @@ class SQLStatementSelect(BaseModel):
245247
from_: SQLTableIdentifier = Field(..., alias="from")
246248
where: Optional[SQLWhereClause] = None
247249
limit: Optional[int] = None
248-
joins: Optional[List[SQLLeftJoin]] = None
249-
joins_outer: bool = False
250+
251+
# Joins are combined in CNF format - inner lists of joins must have one matching joined row
252+
joins: List[List[SQLLeftJoin]] = []
250253

251254
model_config = ConfigDict(populate_by_name=True)
252255

@@ -259,50 +262,47 @@ def build(self, context: SQLBuildContext):
259262
else:
260263
stmt = sa.select(table_obj)
261264

262-
if self.joins:
263-
joined_cols: List[sa.ColumnElement] = []
264-
joined: sa.FromClause = table_obj
265-
exists_constraints: List[sa.ColumnExpressionArgument] = []
266-
for join in self.joins: # pylint: disable=not-an-iterable
265+
joined: sa.FromClause = table_obj
266+
join_and_conditions = []
267+
for join_list in self.joins:
268+
join_or_conditions: List[sa.ColumnExpressionArgument] = []
269+
for join in join_list:
267270
right = join.right.build(context).alias()
268271

272+
join_on = [
273+
table_obj.c[lft_col] == right.c[rht_col]
274+
for lft_col, rht_col in zip(join.left_columns, join.right_columns)
275+
]
276+
if join.left_discriminator:
277+
disc_col, disc_val = join.left_discriminator
278+
join_on.append(table_obj.c[disc_col] == disc_val)
279+
if join.right_discriminator:
280+
disc_col, disc_val = join.right_discriminator
281+
join_on.append(right.c[disc_col] == disc_val)
282+
269283
if join.half_unique and table_obj.primary_key:
270284
joined = joined.join(
271285
right,
272-
onclause=sa.and_(
273-
*(
274-
table_obj.c[lft_col] == right.c[rht_col]
275-
for lft_col, rht_col in zip(
276-
join.left_columns, join.right_columns
277-
)
278-
)
279-
),
280-
isouter=self.joins_outer,
281-
)
282-
joined_cols.extend(
283-
right.c[rht_col] for rht_col in join.right_columns
286+
onclause=sa.and_(*join_on),
287+
isouter=len(join_list) > 1,
284288
)
285-
else:
286-
exists_constraints.append(
287-
sa.exists().where(
288-
*(
289-
table_obj.c[lft_col] == right.c[rht_col]
290-
for lft_col, rht_col in zip(
291-
join.left_columns, join.right_columns
292-
)
293-
)
289+
if len(join_list) > 1:
290+
join_or_conditions.extend(
291+
right.c[rht_col].is_not(None)
292+
for rht_col in join.right_columns
294293
)
295-
)
294+
else:
295+
join_or_conditions.append(sa.exists().where(*join_on))
296+
297+
if join_or_conditions:
298+
join_and_conditions.append(sa.or_(*join_or_conditions))
296299

297-
stmt = stmt.select_from(joined)
298-
if joined is not table_obj:
299-
stmt = stmt.group_by(*table_obj.primary_key.columns)
300+
stmt = stmt.select_from(joined)
301+
if joined is not table_obj:
302+
stmt = stmt.group_by(*table_obj.primary_key.columns)
300303

301-
if self.joins_outer:
302-
exists_constraints.extend(col.is_not(None) for col in joined_cols)
303-
stmt = stmt.where(sa.or_(*exists_constraints))
304-
elif exists_constraints:
305-
stmt = stmt.where(sa.and_(*exists_constraints))
304+
if join_and_conditions:
305+
stmt = stmt.where(sa.and_(*join_and_conditions))
306306

307307
if self.where:
308308
stmt = stmt.where(self.where.build(context, table_obj))
@@ -329,7 +329,6 @@ def simplify(self) -> "SQLStatementSelect":
329329
kwargs["limit"] = self.limit
330330
if self.joins:
331331
kwargs["joins"] = self.joins
332-
kwargs["joins_outer"] = self.joins_outer
333332

334333
return SQLStatementSelect(**kwargs) # type: ignore
335334

subsetter/planner.py

Lines changed: 123 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -69,6 +69,7 @@ def _plan_internal(self) -> SubsetPlan:
6969
)
7070
self._remove_ignore_fks()
7171
self._add_extra_fks()
72+
self._add_polymorphic_fks()
7273
if self.config.include_dependencies:
7374
self._check_ignore_tables()
7475
self._check_passthrough_tables()
@@ -202,6 +203,85 @@ def _add_extra_fks(self) -> None:
202203
),
203204
)
204205

206+
def _add_polymorphic_fks(self) -> None:
207+
"""Add in configured polymorphic foreign keys requested."""
208+
for index, poly_fk in enumerate(self.config.polymorphic_fks):
209+
src_schema, src_table_name = parse_table_name(poly_fk.table)
210+
table = self.meta.tables.get((src_schema, src_table_name))
211+
if table is None:
212+
LOGGER.warning(
213+
"Found no source table %s.%s referenced in polymorphic_fks[%d]",
214+
src_schema,
215+
src_table_name,
216+
index,
217+
)
218+
continue
219+
220+
src_missing_cols = {
221+
col for col in poly_fk.columns if col not in table.table_obj.columns
222+
}
223+
if src_missing_cols:
224+
LOGGER.warning(
225+
"Columns %s do not exist in %s.%s referenced in poly_fks[%d]",
226+
src_missing_cols,
227+
src_schema,
228+
src_table_name,
229+
index,
230+
)
231+
continue
232+
233+
if poly_fk.discriminator_column not in table.table_obj.columns:
234+
LOGGER.warning(
235+
"Column %s does not exist in %s.%s referenced in poly_fks[%d].discriminator_column",
236+
poly_fk.discriminator_column,
237+
src_schema,
238+
src_table_name,
239+
index,
240+
)
241+
continue
242+
243+
for discriminator_value, key_dest in poly_fk.destinations.items():
244+
dst_schema, dst_table_name = parse_table_name(key_dest.table)
245+
dst_table = self.meta.tables.get((dst_schema, dst_table_name))
246+
if dst_table is None:
247+
LOGGER.warning(
248+
"Found no destination table %s.%s referenced in poly_fks[%d].destinations[%s]",
249+
dst_schema,
250+
dst_table_name,
251+
index,
252+
discriminator_value,
253+
)
254+
continue
255+
256+
dst_missing_cols = {
257+
col
258+
for col in key_dest.columns
259+
if col not in dst_table.table_obj.columns
260+
}
261+
if dst_missing_cols:
262+
LOGGER.warning(
263+
"Columns %s do not exist in %s.%s referenced in poly_fks[%d].destinations[%s]",
264+
dst_missing_cols,
265+
dst_schema,
266+
dst_table_name,
267+
index,
268+
discriminator_value,
269+
)
270+
continue
271+
272+
table.foreign_keys.append(
273+
ForeignKey(
274+
columns=tuple(poly_fk.columns),
275+
dst_schema=dst_schema,
276+
dst_table=dst_table_name,
277+
dst_columns=tuple(key_dest.columns),
278+
src_discriminator=(
279+
poly_fk.discriminator_column,
280+
discriminator_value,
281+
),
282+
),
283+
)
284+
205285
def _remove_ignore_fks(self) -> None:
206286
"""Remove requested foreign keys"""
207287
for ignore_fk in self.config.ignore_fks:
@@ -322,24 +402,51 @@ def _is_distinct(table_obj: sa.Table, cols: Iterable[str]) -> bool:
322402
return True
323403
return False
324404

405+
# Create joins in conjunctive normal form.
406+
fks_to_join = []
407+
408+
# reverse foreign keys just get OR'ed together
409+
if rev_foreign_keys:
410+
fks_to_join.append(rev_foreign_keys)
411+
412+
# forward foreign keys get AND'ed except for polymorphic fks which OR when using the same
413+
# discriminator column.
414+
fk_disc_index: dict[str, int] = {}
415+
for fk in foreign_keys:
416+
if fk.src_discriminator:
417+
disc_col = fk.src_discriminator[0]
418+
if disc_col in fk_disc_index:
419+
fks_to_join[fk_disc_index[disc_col]].append(fk)
420+
else:
421+
fk_disc_index[disc_col] = len(fks_to_join)
422+
fks_to_join.append([fk])
423+
else:
424+
fks_to_join.append([fk])
425+
325426
fk_joins = []
326-
for fk in foreign_keys or rev_foreign_keys:
327-
dst_table = self.meta.tables[(fk.dst_schema, fk.dst_table)]
328-
half_unique = _is_distinct(table.table_obj, fk.columns) or _is_distinct(
329-
dst_table.table_obj, fk.dst_columns
330-
)
331-
fk_joins.append(
332-
SQLLeftJoin(
333-
right=SQLTableIdentifier(
334-
table_schema=fk.dst_schema,
335-
table_name=fk.dst_table,
336-
sampled=True,
337-
),
338-
left_columns=list(fk.columns),
339-
right_columns=list(fk.dst_columns),
340-
half_unique=half_unique,
427+
for fk_join_list in fks_to_join:
428+
or_joins = []
429+
for fk in fk_join_list:
430+
dst_table = self.meta.tables[(fk.dst_schema, fk.dst_table)]
431+
half_unique = _is_distinct(table.table_obj, fk.columns) or _is_distinct(
432+
dst_table.table_obj, fk.dst_columns
341433
)
342-
)
434+
435+
or_joins.append(
436+
SQLLeftJoin(
437+
right=SQLTableIdentifier(
438+
table_schema=fk.dst_schema,
439+
table_name=fk.dst_table,
440+
sampled=True,
441+
),
442+
left_columns=list(fk.columns),
443+
right_columns=list(fk.dst_columns),
444+
half_unique=half_unique,
445+
left_discriminator=list(fk.src_discriminator or ()),
446+
right_discriminator=list(fk.dst_discriminator or ()),
447+
)
448+
)
449+
fk_joins.append(or_joins)
343450

344451
conf_constraints = self.config.table_constraints.get(
345452
f"{table.schema}.{table.name}", []

tests/data/big_join.yaml

Lines changed: 9 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -36,16 +36,15 @@ expected_plan:
3636
schema: test
3737
table: users
3838
joins:
39-
- half_unique: false
40-
left_columns:
41-
- state
42-
right:
43-
sampled: true
44-
schema: test
45-
table: homes
46-
right_columns:
47-
- state
48-
joins_outer: true
39+
- - half_unique: false
40+
left_columns:
41+
- state
42+
right:
43+
sampled: true
44+
schema: test
45+
table: homes
46+
right_columns:
47+
- state
4948
type: select
5049

5150
expected_sample:

0 commit comments

Comments
 (0)