From ed61a21a18a2d12a36b45a24cf8754d250c7cfce Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97SO?= <142557582+Linxiushen@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:12:53 +0800 Subject: [PATCH] Fix codegen imports for public tagging utilities --- fiddle/_src/codegen/import_manager.py | 15 +++++--- fiddle/_src/codegen/import_manager_test.py | 41 ++++++++++++++++++++++ fiddle/_src/codegen/legacy_codegen_test.py | 12 +++++++ fiddle/_src/codegen/new_codegen_test.py | 17 +++++++++ fiddle/_src/special_overrides.py | 6 ++++ 5 files changed, 87 insertions(+), 4 deletions(-) diff --git a/fiddle/_src/codegen/import_manager.py b/fiddle/_src/codegen/import_manager.py index d003393c..b18d65bc 100644 --- a/fiddle/_src/codegen/import_manager.py +++ b/fiddle/_src/codegen/import_manager.py @@ -19,7 +19,7 @@ import enum import functools import inspect -from typing import Any, Dict, Union +from typing import Any, Dict, Optional, Union from absl import logging from fiddle._src import special_overrides @@ -177,7 +177,9 @@ def _compatible_with_existing(self, stmt: AnyImport) -> bool: return get_full_module_name(node).split(".")[0] == base_module_name return None # pytype: disable=bad-return-type - def add_by_name(self, full_module_name: str) -> str: + def add_by_name( + self, full_module_name: str, *, symbol_name: Optional[str] = None + ) -> str: """Adds an import given a module name. This is a slightly lower-level API than `add`; you should only use it if @@ -185,6 +187,8 @@ def add_by_name(self, full_module_name: str) -> str: Args: full_module_name: String module name to try to import. + symbol_name: Qualified symbol name, if known, to select a symbol-specific + import alias instead of the module-wide alias. Returns: Name for the imported module. This is usually the last name, possibly @@ -196,7 +200,10 @@ def add_by_name(self, full_module_name: str) -> str: if result is None: result = _make_import(full_module_name) else: - result = parse_import(result.module_import_alias) + import_alias = result.module_import_alias + if symbol_name is not None: + import_alias = result.symbol_import_aliases.get(symbol_name, import_alias) + result = parse_import(import_alias) # Since multiple things could be aliased to the same import, rewrite # the module name to the alias' module name. @@ -242,7 +249,7 @@ def add(self, value: Any) -> str: ) return value_qualname - imported_name = self.add_by_name(module_name) + imported_name = self.add_by_name(module_name, symbol_name=value_qualname) return f"{imported_name}.{value_qualname}" def sorted_import_nodes(self): diff --git a/fiddle/_src/codegen/import_manager_test.py b/fiddle/_src/codegen/import_manager_test.py index 72f263bd..44c03c45 100644 --- a/fiddle/_src/codegen/import_manager_test.py +++ b/fiddle/_src/codegen/import_manager_test.py @@ -16,9 +16,13 @@ """Tests for import_manager.""" import textwrap +from unittest import mock from absl.testing import absltest from absl.testing import parameterized +import fiddle as fdl +from fiddle import tagging +from fiddle._src import special_overrides from fiddle._src.codegen import import_manager from fiddle._src.codegen import namespace import libcst as cst @@ -88,6 +92,43 @@ def test_import_manager_aliasing_builtin(self): module = cst.Module(body=manager.sorted_import_lines()) self.assertEqual(module.code.strip(), "import fiddle as fdl_2") + @parameterized.parameters((), ("fdl", "tagging")) + def test_tagging_imports_resolve_to_original_symbols(self, *reserved_names): + manager = import_manager.ImportManager( + namespace.Namespace(set(reserved_names)) + ) + values = [ + tagging.list_tags, + tagging.materialize_tags, + fdl.add_tag, + fdl.clear_tags, + fdl.get_tags, + fdl.remove_tag, + fdl.set_tagged, + fdl.set_tags, + fdl.Tag, + fdl.TaggedValue, + ] + names = [manager.add(value) for value in values] + module = cst.Module(body=manager.sorted_import_lines()) + imported_symbols = {} + exec(module.code, imported_symbols) # pylint: disable=exec-used + for name, value in zip(names, values): + with self.subTest(name=name): + self.assertIs(eval(name, imported_symbols), value) # pylint: disable=eval-used + self.assertLen(manager.sorted_import_nodes(), 2) + + def test_registered_tagging_alias_replaces_default_overrides(self): + with mock.patch.dict(special_overrides.SPECIAL_OVERRIDES_MAP): + import_manager.register_import_alias( + "fiddle._src.tagging", "from fiddle import tagging as tags" + ) + manager = import_manager.ImportManager(namespace.Namespace()) + self.assertEqual(manager.add(tagging.list_tags), "tags.list_tags") + self.assertEqual(manager.add(fdl.Tag), "tags.Tag") + module = cst.Module(body=manager.sorted_import_lines()) + self.assertEqual(module.code.strip(), "from fiddle import tagging as tags") + def test_dotted_special_import_okay(self): """Tests cases where dotted imports are used.""" import_manager.register_import_alias("lace.bar", "import lace.bar") diff --git a/fiddle/_src/codegen/legacy_codegen_test.py b/fiddle/_src/codegen/legacy_codegen_test.py index aeeae10c..292fea58 100644 --- a/fiddle/_src/codegen/legacy_codegen_test.py +++ b/fiddle/_src/codegen/legacy_codegen_test.py @@ -20,6 +20,7 @@ from absl.testing import absltest import fiddle as fdl +from fiddle import tagging from fiddle._src.codegen import legacy_codegen from fiddle._src.codegen import test_util from fiddle._src.codegen.test_submodule import test_util as submodule_test_util @@ -112,6 +113,17 @@ def partial_config() -> fdl.Partial[Baz]: class CodegenTest(absltest.TestCase): + def test_tagging_function_round_trip(self): + config = fdl.Partial(tagging.list_tags) + code = "\n".join(legacy_codegen.codegen_dot_syntax(config).lines()) + exec_globals = {} + exec(code, exec_globals) # pylint: disable=exec-used + restored_config = exec_globals["build_config"]() + self.assertEqual(restored_config, config) + tagged_config = fdl.Config(test_util.Foo, a=1) + fdl.add_tag(tagged_config, "a", fdl.Tag) + self.assertEqual(fdl.build(restored_config)(tagged_config), {fdl.Tag}) + def test_codegen_dot_syntax_shared(self): cfg = shared_config() result = legacy_codegen.codegen_dot_syntax(cfg) diff --git a/fiddle/_src/codegen/new_codegen_test.py b/fiddle/_src/codegen/new_codegen_test.py index 4ce1ed4b..47b2ce9a 100644 --- a/fiddle/_src/codegen/new_codegen_test.py +++ b/fiddle/_src/codegen/new_codegen_test.py @@ -24,6 +24,8 @@ from absl import flags from absl.testing import absltest +import fiddle as fdl +from fiddle import tagging from fiddle._src.codegen import new_codegen from fiddle._src.testing.example import fake_encoder_decoder @@ -58,6 +60,21 @@ def test_codegen(self): fixture = self._load_code_as_module(code).config_fixture self.assertEqual(fixture(), config) + def test_tagging_functions_round_trip(self): + config = fdl.Config( + dict, + list_tags=tagging.list_tags, + materialize_tags=tagging.materialize_tags, + add_tag=fdl.add_tag, + ) + code = new_codegen.new_codegen(config) + fixture = self._load_code_as_module(code).config_fixture + self.assertEqual(fixture(), config) + functions = fdl.build(fixture()) + self.assertIs(functions["list_tags"], tagging.list_tags) + self.assertIs(functions["materialize_tags"], tagging.materialize_tags) + self.assertIs(functions["add_tag"], fdl.add_tag) + def test_sub_fixtures(self): config = fake_encoder_decoder.fixture.as_buildable() code = new_codegen.new_codegen( diff --git a/fiddle/_src/special_overrides.py b/fiddle/_src/special_overrides.py index 5c5a942e..a03aec0e 100644 --- a/fiddle/_src/special_overrides.py +++ b/fiddle/_src/special_overrides.py @@ -32,6 +32,8 @@ class SpecialOverrides: migrated_symbol_destination_modules: Dict[str, str] = dataclasses.field( default_factory=dict ) + # Import aliases for symbols that are not exported by the module-wide alias. + symbol_import_aliases: Dict[str, str] = dataclasses.field(default_factory=dict) SPECIAL_OVERRIDES_MAP = { @@ -89,6 +91,10 @@ class SpecialOverrides: "fiddle._src.tagging": SpecialOverrides( module_name="fiddle._src.tagging", module_import_alias="import fiddle as fdl", + symbol_import_aliases={ + "list_tags": "from fiddle import tagging", + "materialize_tags": "from fiddle import tagging", + }, ), }