@@ -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 } " , []
0 commit comments