@@ -430,6 +430,7 @@ def _find_links(node_a, all_nodes):
430430
431431
432432def map_it (sources , extension , no_trimming , exclude_namespaces , exclude_functions ,
433+ include_only_namespaces , include_only_functions ,
433434 skip_parse_errors , lang_params ):
434435 '''
435436 Given a language implementation and a list of filenames, do these things:
@@ -447,6 +448,8 @@ def map_it(sources, extension, no_trimming, exclude_namespaces, exclude_function
447448 :param bool no_trimming:
448449 :param list exclude_namespaces:
449450 :param list exclude_functions:
451+ :param list include_only_namespaces:
452+ :param list include_only_functions:
450453 :param bool skip_parse_errors:
451454 :param LanguageParams lang_params:
452455
@@ -475,11 +478,11 @@ def map_it(sources, extension, no_trimming, exclude_namespaces, exclude_function
475478 file_group = make_file_group (file_ast_tree , source , extension )
476479 file_groups .append (file_group )
477480
478- # 3. Trim namespaces / functions that we don't want
479- if exclude_namespaces :
480- file_groups = _exclude_namespaces (file_groups , exclude_namespaces )
481- if exclude_functions :
482- file_groups = _exclude_functions (file_groups , exclude_functions )
481+ # 3. Trim namespaces / functions to exactly what we want
482+ if exclude_namespaces or include_only_namespaces :
483+ file_groups = _limit_namespaces (file_groups , exclude_namespaces , include_only_namespaces )
484+ if exclude_functions or include_only_functions :
485+ file_groups = _limit_functions (file_groups , exclude_functions , include_only_functions )
483486
484487 # 4. Consolidate structures
485488 all_subgroups = flatten (g .all_groups () for g in file_groups )
@@ -557,48 +560,71 @@ def map_it(sources, extension, no_trimming, exclude_namespaces, exclude_function
557560 return file_groups , all_nodes , edges
558561
559562
560- def _exclude_namespaces (file_groups , exclude_namespaces ):
563+ def _limit_namespaces (file_groups , exclude_namespaces , include_only_namespaces ):
561564 """
562565 Exclude namespaces (classes/modules) which match any of the exclude_namespaces
563566
564567 :param list[Group] file_groups:
565568 :param list exclude_namespaces:
569+ :param list include_only_namespaces:
566570 :rtype: list[Group]
567571 """
572+
573+ removed_namespaces = set ()
574+
575+ for group in list (file_groups ):
576+ if group .token in exclude_namespaces :
577+ for node in group .all_nodes ():
578+ node .remove_from_parent ()
579+ removed_namespaces .add (group .token )
580+ if include_only_namespaces and group .token not in include_only_namespaces :
581+ for node in group .nodes :
582+ node .remove_from_parent ()
583+ removed_namespaces .add (group .token )
584+
585+ for subgroup in group .all_groups ():
586+ print (subgroup , subgroup .all_parents ())
587+ if subgroup .token in exclude_namespaces :
588+ for node in subgroup .all_nodes ():
589+ node .remove_from_parent ()
590+ removed_namespaces .add (subgroup .token )
591+ if include_only_namespaces and \
592+ subgroup .token not in include_only_namespaces and \
593+ all (p .token not in include_only_namespaces for p in subgroup .all_parents ()):
594+ for node in subgroup .nodes :
595+ node .remove_from_parent ()
596+ removed_namespaces .add (group .token )
597+
568598 for namespace in exclude_namespaces :
569- found = False
570- for group in list (file_groups ):
571- if group .token == namespace :
572- file_groups .remove (group )
573- found = True
574- for subgroup in group .all_groups ():
575- if subgroup .token == namespace :
576- subgroup .remove_from_parent ()
577- found = True
578- if not found :
599+ if namespace not in removed_namespaces :
579600 logging .warning (f"Could not exclude namespace '{ namespace } ' "
580- "because it was not found." )
601+ "because it was not found." )
581602 return file_groups
582603
583604
584- def _exclude_functions (file_groups , exclude_functions ):
605+ def _limit_functions (file_groups , exclude_functions , include_only_functions ):
585606 """
586607 Exclude nodes (functions) which match any of the exclude_functions
587608
588609 :param list[Group] file_groups:
589610 :param list exclude_functions:
611+ :param list include_only_functions:
590612 :rtype: list[Group]
591613 """
614+
615+ removed_functions = set ()
616+
617+ for group in list (file_groups ):
618+ for node in group .all_nodes ():
619+ if node .token in exclude_functions or \
620+ (include_only_functions and node .token not in include_only_functions ):
621+ node .remove_from_parent ()
622+ removed_functions .add (node .token )
623+
592624 for function_name in exclude_functions :
593- found = False
594- for group in list (file_groups ):
595- for node in group .all_nodes ():
596- if node .token == function_name :
597- node .remove_from_parent ()
598- found = True
599- if not found :
625+ if function_name not in removed_functions :
600626 logging .warning (f"Could not exclude function '{ function_name } ' "
601- "because it was not found." )
627+ "because it was not found." )
602628 return file_groups
603629
604630
@@ -636,6 +662,7 @@ def _generate_final_img(output_file, extension, final_img_filename, num_edges):
636662
637663def code2flow (raw_source_paths , output_file , language = None , hide_legend = True ,
638664 exclude_namespaces = None , exclude_functions = None ,
665+ include_only_namespaces = None , include_only_functions = None ,
639666 no_grouping = False , no_trimming = False , skip_parse_errors = False ,
640667 lang_params = None , subset_params = None , level = logging .INFO ):
641668 """
@@ -648,6 +675,8 @@ def code2flow(raw_source_paths, output_file, language=None, hide_legend=True,
648675 :param bool hide_legend: Omit the legend from the output
649676 :param list exclude_namespaces: List of namespaces to exclude
650677 :param list exclude_functions: List of functions to exclude
678+ :param list include_only_namespaces: List of namespaces to include
679+ :param list include_only_functions: List of functions to include
651680 :param bool no_grouping: Don't group functions into namespaces in the final output
652681 :param bool no_trimming: Don't trim orphaned functions / namespaces
653682 :param bool skip_parse_errors: If a language parser fails to parse a file, skip it
@@ -660,11 +689,16 @@ def code2flow(raw_source_paths, output_file, language=None, hide_legend=True,
660689
661690 if not isinstance (raw_source_paths , list ):
662691 raw_source_paths = [raw_source_paths ]
692+ lang_params = lang_params or LanguageParams ()
693+
663694 exclude_namespaces = exclude_namespaces or []
664695 assert isinstance (exclude_namespaces , list )
665696 exclude_functions = exclude_functions or []
666697 assert isinstance (exclude_functions , list )
667- lang_params = lang_params or LanguageParams ()
698+ include_only_namespaces = include_only_namespaces or []
699+ assert isinstance (include_only_namespaces , list )
700+ include_only_functions = include_only_functions or []
701+ assert isinstance (include_only_functions , list )
668702
669703 logging .basicConfig (format = "Code2Flow: %(message)s" , level = level )
670704
@@ -691,6 +725,7 @@ def code2flow(raw_source_paths, output_file, language=None, hide_legend=True,
691725
692726 file_groups , all_nodes , edges = map_it (sources , language , no_trimming ,
693727 exclude_namespaces , exclude_functions ,
728+ include_only_namespaces , include_only_functions ,
694729 skip_parse_errors , lang_params )
695730
696731 if subset_params :
@@ -761,6 +796,12 @@ def main(sys_argv=None):
761796 parser .add_argument (
762797 '--exclude-namespaces' ,
763798 help = 'exclude namespaces (Classes, modules, etc) from the output. Comma delimited.' )
799+ parser .add_argument (
800+ '--include-only-functions' ,
801+ help = 'include only functions in the output. Comma delimited.' )
802+ parser .add_argument (
803+ '--include-only-namespaces' ,
804+ help = 'include only namespaces (Classes, modules, etc) in the output. Comma delimited.' )
764805 parser .add_argument (
765806 '--no-grouping' , action = 'store_true' ,
766807 help = 'instead of grouping functions into namespaces, let functions float.' )
@@ -801,6 +842,9 @@ def main(sys_argv=None):
801842
802843 exclude_namespaces = list (filter (None , (args .exclude_namespaces or "" ).split (',' )))
803844 exclude_functions = list (filter (None , (args .exclude_functions or "" ).split (',' )))
845+ include_only_namespaces = list (filter (None , (args .include_only_namespaces or "" ).split (',' )))
846+ include_only_functions = list (filter (None , (args .include_only_functions or "" ).split (',' )))
847+
804848 lang_params = LanguageParams (args .source_type , args .ruby_version )
805849 subset_params = SubsetParams .generate (args .target_function , args .upstream_depth ,
806850 args .downstream_depth )
@@ -812,6 +856,8 @@ def main(sys_argv=None):
812856 hide_legend = args .hide_legend ,
813857 exclude_namespaces = exclude_namespaces ,
814858 exclude_functions = exclude_functions ,
859+ include_only_namespaces = include_only_namespaces ,
860+ include_only_functions = include_only_functions ,
815861 no_grouping = args .no_grouping ,
816862 no_trimming = args .no_trimming ,
817863 skip_parse_errors = args .skip_parse_errors ,
0 commit comments