Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 11 additions & 4 deletions fiddle/_src/codegen/import_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -177,14 +177,18 @@ 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
you don't have access to a function or class to pass to `add`.

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
Expand All @@ -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.
Expand Down Expand Up @@ -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):
Expand Down
41 changes: 41 additions & 0 deletions fiddle/_src/codegen/import_manager_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
12 changes: 12 additions & 0 deletions fiddle/_src/codegen/legacy_codegen_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
17 changes: 17 additions & 0 deletions fiddle/_src/codegen/new_codegen_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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(
Expand Down
6 changes: 6 additions & 0 deletions fiddle/_src/special_overrides.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand Down Expand Up @@ -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",
},
),
}

Expand Down
Loading