diff --git a/docs/writing-plugins/options.md b/docs/writing-plugins/options.md index db72899..0b9000a 100644 --- a/docs/writing-plugins/options.md +++ b/docs/writing-plugins/options.md @@ -38,7 +38,7 @@ Supported field types: `str`, `bool`, `int`, `float`, `Literal[...]`, `StrEnum`, ## Framework-reserved options -The framework reserves certain option names for all plugins, currently: `no_fmt_off`, `escape_module_with_hash`, and `rewrite_imports`. +The framework reserves certain option names for all plugins, currently: `no_fmt_off`, `escape_module_with_hash`, and `map_imports`. If your `Options` dataclass defines a field with a reserved name, `run()` raises a `ValueError`. ### no_fmt_off @@ -66,45 +66,48 @@ plugins: opt: escape_module_with_hash ``` -### rewrite_imports +### map_imports -Rewrites imports of generated modules that match a glob pattern to an absolute package. -This makes it possible to reference generated code published from a separate package. +By default, generated code imports dependencies from the local output with a relative import. +For example, the module generated for `foo/bar.proto` imports a message from `buf/validate/validate.proto` as `from ..buf.validate.validate_pb import Rule`. -The option takes the form `rewrite_imports=:` and can be given multiple times; the first matching pattern wins. -The pattern is a very reduced subset of glob: - -- `*` matches zero or more characters except `/`. -- `**/` matches zero or more path elements, where an element is one or more characters with a trailing `/`. - -The pattern is matched against the import path of the module before it is made relative to the file importing it. -A generated module such as `.google.type.foo_pb` is matched as the file path `./google/type/foo_pb.py`, relative to the generation root. -On a match, the target package is prepended: +If a dependency is provided by a package instead, use `map_imports` to import it from there. +The option takes the form `map_imports=:` and can be given multiple times. +The pattern is matched against the path of the Protobuf file, and the first matching pattern wins. +The target is a Python package that is prepended to the module path derived from the Protobuf file: ```yaml title="buf.gen.yaml" plugins: - local: protoc-gen-hello out: src/gen - opt: rewrite_imports=./google/type/**/*_pb.py:mypkg.gen + opt: map_imports=google/type/:mypkg.gen ``` -With this option, `from .google.type.foo_pb import Foo` is generated as `from mypkg.gen.google.type.foo_pb import Foo` instead. -References to symbols defined in the file being generated are not imports and are never rewritten. +With this option, a message from `google/type/date.proto` is imported as `from mypkg.gen.google.type.date_pb import Date`. + +Patterns support a subset of glob: + +- `*` matches zero or more characters except `/`. +- `**` matches zero or more characters, including `/`. +- `**/` matches zero or more directories. +- A trailing `/` matches every file in the directory and its subdirectories. -An empty target rewrites matching imports to the canonical import path, the module path derived from the proto file name, relative to the root of `sys.path`. +An empty target imports from the canonical module path: the module path derived from the Protobuf file, relative to the root of `sys.path`. +This is the layout of generated SDKs installed as separate packages, such as those from the Buf Python registry. For example, when `buf/validate/validate.proto` is provided by an installed package: ```yaml title="buf.gen.yaml" plugins: - local: protoc-gen-hello out: src/gen - opt: rewrite_imports=./buf/validate/**/*_pb.py: + opt: "map_imports=buf/validate/:" ``` -This generates `from buf.validate import validate_pb` instead of a relative import that points at a location that does not exist in the output directory. -A rewritten module import that lands at the top level (for example a proto file at the root of the module) is written as a plain `import foo_pb` statement. +This generates `from buf.validate import validate_pb` and `from buf.validate.validate_pb import Rule`. +A mapped module at the top level, such as one generated for a Protobuf file at the root, is written as a plain `import foo_pb` statement. -Absolute imports (such as `protobuf.wkt`) are matched without a leading `./` and `.py` extension (for example `protobuf/wkt`), and are replaced by the target entirely. +Mapping applies to imports derived from descriptors. +Identifiers constructed directly from a `Module` are not mapped, and well-known types are always imported from `protobuf.wkt`. ## Example: Sensitive fields plugin diff --git a/src/protobuf/plugin/_file.py b/src/protobuf/plugin/_file.py index 70c7341..c2b4893 100644 --- a/src/protobuf/plugin/_file.py +++ b/src/protobuf/plugin/_file.py @@ -22,12 +22,12 @@ from protobuf import DescEnum, DescExtension, DescFile, DescMessage, ScalarType from protobuf.plugin._ident import Ident, Module -from protobuf.plugin._rewrite_imports import rewrite_module_path +from protobuf.plugin._map_imports import map_import_target if TYPE_CHECKING: from collections.abc import Generator, Iterable, Iterator - from protobuf.plugin._rewrite_imports import RewriteImports + from protobuf.plugin._map_imports import MapImports _INDENT = " " * 4 @@ -223,7 +223,7 @@ def __init__( parameter: str, *, escape_module_with_hash: bool = False, - rewrite_imports: RewriteImports = (), + map_imports: MapImports = (), ) -> None: self.path = path self.module = module @@ -232,7 +232,7 @@ def __init__( self._plugin_version = plugin_version self._parameter = parameter self._escape_module_with_hash = escape_module_with_hash - self._rewrite_imports = rewrite_imports + self._map_imports = map_imports self._indent = 0 self._type_checking = False self._in_doc = False @@ -334,39 +334,45 @@ def _to_el(self, v: object) -> str | Ident: ) case _: return repr(v) - ident = self._relativize(self._rewrite_import(ident)) + ident = self._relativize(self._map_import(ident)) if ident.type_only: self._type_imports[ident.module].add(ident) else: self._runtime_imports[ident.module].add(ident) return ident - def _rewrite_import(self, ident: Ident) -> Ident: - if not self._rewrite_imports: + def _map_import(self, ident: Ident) -> Ident: + # Only imports derived from a descriptor know their Protobuf file. + if not self._map_imports or ident._desc is None: + return ident + if not _is_relative(ident.module): return ident - # A DescFile ident imports the module itself, so its import - # path includes the ident name. is_module_import = isinstance(ident._desc, DescFile) + if not is_module_import and _module_segments(ident.module) == _module_segments( + self.module + ): + # References to symbols in this file are not imports. + return ident + file = ident._desc if isinstance(ident._desc, DescFile) else ident._desc.file + target = map_import_target(file.name, self._map_imports) + if target is None: + return ident module_path = ident.module.path if is_module_import: + # A DescFile ident imports the module itself, so the module + # path includes the ident name. sep = "" if module_path.endswith(".") else "." module_path = f"{module_path}{sep}{ident.name}" - elif _is_relative(ident.module) and _module_segments( - ident.module - ) == _module_segments(self.module): - # References to symbols in this file are not imports. - return ident - rewritten = rewrite_module_path(module_path, self._rewrite_imports) - if rewritten is None: - return ident + dotted = module_path.removeprefix(".") + mapped = f"{target}.{dotted}" if target else dotted if is_module_import: - parent, _, name = rewritten.rpartition(".") + parent, _, name = mapped.rpartition(".") # A module import that lands at the top level is written as # a plain `import X` statement, keyed by the module itself. module = Module(parent) if parent else Module(name) return Ident(name, module, type_only=ident.type_only, _desc=ident._desc) return Ident( - ident.name, Module(rewritten), type_only=ident.type_only, _desc=ident._desc + ident.name, Module(mapped), type_only=ident.type_only, _desc=ident._desc ) def _relativize(self, ident: Ident) -> Ident: diff --git a/src/protobuf/plugin/_map_imports.py b/src/protobuf/plugin/_map_imports.py new file mode 100644 index 0000000..b4c0b00 --- /dev/null +++ b/src/protobuf/plugin/_map_imports.py @@ -0,0 +1,68 @@ +# Copyright (c) 2025-2026 Buf Technologies, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from __future__ import annotations + +import re + +MapImports = tuple[tuple[re.Pattern[str], str], ...] + +_OPTION_NAME = "map_imports" + +_TARGET = re.compile(r"[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)*") + + +def compile_map_imports(mappings: dict[str, str]) -> MapImports: + compiled: list[tuple[re.Pattern[str], str]] = [] + for pattern, target in mappings.items(): + if not pattern: + msg = f"option '{_OPTION_NAME}': pattern must not be empty" + raise ValueError(msg) + # An empty target (or ".") maps to the canonical module path. + normalized = target.removesuffix(".") + if normalized and not _TARGET.fullmatch(normalized): + msg = f"option '{_OPTION_NAME}': target '{target}' must be a Python package path" + raise ValueError(msg) + compiled.append((_glob_to_regex(pattern), normalized)) + return tuple(compiled) + + +def map_import_target(proto_name: str, map_imports: MapImports) -> str | None: + for pattern, target in map_imports: + if pattern.fullmatch(proto_name): + return target + return None + + +def _glob_to_regex(pattern: str) -> re.Pattern[str]: + parts: list[str] = [] + i = 0 + while i < len(pattern): + char = pattern[i] + if char == "*": + if pattern[i + 1 : i + 2] == "*": + if pattern[i + 2 : i + 3] == "/": + parts.append(r"([^/]+/)*") + i += 3 + continue + parts.append(".*") + i += 2 + continue + parts.append(r"[^/]*") + elif char == "/" and i == len(pattern) - 1: + # A trailing slash matches everything in the directory. + parts.append("/.*") + else: + parts.append(re.escape(char)) + i += 1 + return re.compile("".join(parts)) diff --git a/src/protobuf/plugin/_rewrite_imports.py b/src/protobuf/plugin/_rewrite_imports.py deleted file mode 100644 index 8fbb806..0000000 --- a/src/protobuf/plugin/_rewrite_imports.py +++ /dev/null @@ -1,64 +0,0 @@ -# Copyright (c) 2025-2026 Buf Technologies, Inc. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from __future__ import annotations - -import re - -RewriteImports = tuple[tuple[re.Pattern[str], str], ...] - -_OPTION_NAME = "rewrite_imports" - - -def compile_rewrite_imports(rewrites: dict[str, str]) -> RewriteImports: - return tuple( - (_glob_to_regex(pattern), target.rstrip(".")) - for pattern, target in rewrites.items() - ) - - -def rewrite_module_path(module_path: str, rewrites: RewriteImports) -> str | None: - relative = module_path.startswith(".") - segments = module_path.removeprefix(".").split(".") - if not all(segments): - return None - import_path = f"./{'/'.join(segments)}.py" if relative else "/".join(segments) - for pattern, target in rewrites: - if pattern.fullmatch(import_path): - if relative: - dotted = ".".join(segments) - return f"{target}.{dotted}" if target else dotted - if not target: - msg = f"option '{_OPTION_NAME}': cannot rewrite absolute import '{module_path}' to an empty target" - raise ValueError(msg) - return target - return None - - -def _glob_to_regex(pattern: str) -> re.Pattern[str]: - parts: list[str] = [] - i = 0 - while i < len(pattern): - char = pattern[i] - if char == "*": - if pattern[i + 1 : i + 3] == "*/": - parts.append(r"([^/]+/)*") - i += 3 - continue - parts.append(r"[^/]*") - else: - parts.append(re.escape(char)) - i += 1 - return re.compile("".join(parts)) diff --git a/src/protobuf/plugin/_run.py b/src/protobuf/plugin/_run.py index 5d90a31..4eb065c 100644 --- a/src/protobuf/plugin/_run.py +++ b/src/protobuf/plugin/_run.py @@ -20,8 +20,8 @@ from protobuf import maximum_supported_edition, minimum_supported_edition from protobuf.plugin._file import write +from protobuf.plugin._map_imports import compile_map_imports from protobuf.plugin._options import parse_options -from protobuf.plugin._rewrite_imports import compile_rewrite_imports from protobuf.plugin._schema import _Schema from protobuf.wkt import CodeGeneratorRequest, CodeGeneratorResponse @@ -168,7 +168,7 @@ def run( name=name, version=version, escape_module_with_hash=fw_opts.escape_module_with_hash, - rewrite_imports=compile_rewrite_imports(fw_opts.rewrite_imports), + map_imports=compile_map_imports(fw_opts.map_imports), ) generate(schema) @@ -232,4 +232,4 @@ def _parse_plugin_options( class _FrameworkOptions: no_fmt_off: bool = False escape_module_with_hash: bool = False - rewrite_imports: dict[str, str] = dataclasses.field(default_factory=dict) + map_imports: dict[str, str] = dataclasses.field(default_factory=dict) diff --git a/src/protobuf/plugin/_schema.py b/src/protobuf/plugin/_schema.py index 6f3197a..3102b48 100644 --- a/src/protobuf/plugin/_schema.py +++ b/src/protobuf/plugin/_schema.py @@ -27,7 +27,7 @@ from protobuf.wkt import CodeGeneratorRequest - from ._rewrite_imports import RewriteImports + from ._map_imports import MapImports T_co = TypeVar("T_co", covariant=True) @@ -98,7 +98,7 @@ def __init__( name: str, version: str, escape_module_with_hash: bool, - rewrite_imports: RewriteImports = (), + map_imports: MapImports = (), ) -> None: self._options = options self._name = name @@ -106,7 +106,7 @@ def __init__( self._parameter = req.parameter self._generated_files: dict[str, _File] = {} self._escape_module_with_hash = escape_module_with_hash - self._rewrite_imports = rewrite_imports + self._map_imports = map_imports file_to_generate = frozenset(req.file_to_generate) source_by_name = {s.name: s for s in req.source_file_descriptors} @@ -170,7 +170,7 @@ def generate_file( self._version, self._parameter, escape_module_with_hash=self._escape_module_with_hash, - rewrite_imports=self._rewrite_imports, + map_imports=self._map_imports, ) self._generated_files[path] = f return f diff --git a/tests/buf.gen.yaml b/tests/buf.gen.yaml index 3485bc1..22f1317 100644 --- a/tests/buf.gen.yaml +++ b/tests/buf.gen.yaml @@ -3,4 +3,4 @@ clean: true plugins: - local: protoc-gen-py out: gen_buf - opt: "rewrite_imports=./google/**/*_pb.py:" + opt: "map_imports=google/:" diff --git a/tests/gen_buf/escaping_pb.py b/tests/gen_buf/escaping_pb.py index cc3b32f..4e925a5 100644 --- a/tests/gen_buf/escaping_pb.py +++ b/tests/gen_buf/escaping_pb.py @@ -13,7 +13,7 @@ # limitations under the License. # Generated from escaping.proto. DO NOT EDIT. -# Generated by protoc-gen-py v0.4.0 with parameter "rewrite_imports=./google/**/*_pb.py:". +# Generated by protoc-gen-py v0.4.0 with parameter "map_imports=google/:". # ruff: noqa: PGH004 # ruff: noqa # fmt: off diff --git a/tests/gen_buf/escaping_proto2_pb.py b/tests/gen_buf/escaping_proto2_pb.py index f6c00be..26a70ae 100644 --- a/tests/gen_buf/escaping_proto2_pb.py +++ b/tests/gen_buf/escaping_proto2_pb.py @@ -13,7 +13,7 @@ # limitations under the License. # Generated from escaping_proto2.proto. DO NOT EDIT. -# Generated by protoc-gen-py v0.4.0 with parameter "rewrite_imports=./google/**/*_pb.py:". +# Generated by protoc-gen-py v0.4.0 with parameter "map_imports=google/:". # ruff: noqa: PGH004 # ruff: noqa # fmt: off diff --git a/tests/gen_buf/local_dep/dep_pb.py b/tests/gen_buf/local_dep/dep_pb.py index 62a0141..e6c1357 100644 --- a/tests/gen_buf/local_dep/dep_pb.py +++ b/tests/gen_buf/local_dep/dep_pb.py @@ -13,7 +13,7 @@ # limitations under the License. # Generated from local_dep/dep.proto. DO NOT EDIT. -# Generated by protoc-gen-py v0.4.0 with parameter "rewrite_imports=./google/**/*_pb.py:". +# Generated by protoc-gen-py v0.4.0 with parameter "map_imports=google/:". # ruff: noqa: PGH004 # ruff: noqa # fmt: off diff --git a/tests/gen_buf/local_import/importer_pb.py b/tests/gen_buf/local_import/importer_pb.py index b84e40e..024014b 100644 --- a/tests/gen_buf/local_import/importer_pb.py +++ b/tests/gen_buf/local_import/importer_pb.py @@ -13,7 +13,7 @@ # limitations under the License. # Generated from local_import/importer.proto. DO NOT EDIT. -# Generated by protoc-gen-py v0.4.0 with parameter "rewrite_imports=./google/**/*_pb.py:". +# Generated by protoc-gen-py v0.4.0 with parameter "map_imports=google/:". # ruff: noqa: PGH004 # ruff: noqa # fmt: off diff --git a/tests/plugin/test_rewrite_imports.py b/tests/plugin/test_map_imports.py similarity index 60% rename from tests/plugin/test_rewrite_imports.py rename to tests/plugin/test_map_imports.py index 63dd73d..8a42674 100644 --- a/tests/plugin/test_rewrite_imports.py +++ b/tests/plugin/test_map_imports.py @@ -20,85 +20,71 @@ from protobuf.plugin import Ident, Module from protobuf.plugin._file import _File, write as gen_write -from protobuf.plugin._rewrite_imports import ( - compile_rewrite_imports, - rewrite_module_path, -) +from protobuf.plugin._map_imports import compile_map_imports, map_import_target if TYPE_CHECKING: from protobuf import DescFile from tests.conftest import Protoc -class TestRewriteModulePath: +class TestMapImportTarget: @pytest.mark.parametrize( - ("pattern", "module_path", "expected"), + ("pattern", "matches"), [ - pytest.param("./foo/*_pb.py", ".foo.bar_pb", "pkg.foo.bar_pb", id="star"), - pytest.param( - "./foo/*_pb.py", - ".foo.baz.bar_pb", - None, - id="star_does_not_cross_separator", - ), - pytest.param( - "./foo/**/*_pb.py", - ".foo.baz.qux.bar_pb", - "pkg.foo.baz.qux.bar_pb", - id="globstar", - ), - pytest.param( - "./foo/**/*_pb.py", - ".foo.bar_pb", - "pkg.foo.bar_pb", - id="globstar_matches_zero_elements", - ), - pytest.param("./**/*_pb.py", ".bar_pb", "pkg.bar_pb", id="root_globstar"), - pytest.param("./foo/*_pb.py", ".foo.bar_px", None, id="suffix_mismatch"), - pytest.param("./bar/*_pb.py", ".foo.bar_pb", None, id="prefix_mismatch"), - pytest.param("foo/*_pb.py", ".foo.bar_pb", None, id="missing_leading_dot"), - pytest.param("./b.r_pb.py", ".bar_pb", None, id="dot_is_literal"), - pytest.param( - "protobuf/wkt", "protobuf.wkt", "pkg", id="absolute_replaced_entirely" - ), - pytest.param("protobuf/*", "protobuf.wkt", "pkg", id="absolute_star"), - pytest.param( - "./protobuf/wkt.py", "protobuf.wkt", None, id="absolute_not_relative" - ), + pytest.param("google/rpc/status.proto", True, id="exact"), + pytest.param("google/rpc/", True, id="trailing_slash"), + pytest.param("google/", True, id="trailing_slash_parent"), + pytest.param("google/rpc/*", True, id="star"), + pytest.param("google/rpc/*.proto", True, id="star_suffix"), + pytest.param("google/rpc/**", True, id="trailing_globstar"), + pytest.param("google/**", True, id="trailing_globstar_parent"), + pytest.param("**", True, id="globstar_only"), + pytest.param("**/status.proto", True, id="leading_globstar"), + pytest.param("google/**/status.proto", True, id="globstar_zero_elements"), + pytest.param("**/*.proto", True, id="globstar_star"), + pytest.param("google/rpc", False, id="directory_without_slash"), + pytest.param("google/*", False, id="star_does_not_cross_separator"), + pytest.param("google/*.proto", False, id="star_suffix_wrong_depth"), + pytest.param("google/rpc/status", False, id="missing_extension"), + pytest.param("google/rpc/s.atus.proto", False, id="dot_is_literal"), + pytest.param("rpc/", False, id="not_anchored"), + pytest.param("google/rpc/status.proto/", False, id="file_with_slash"), ], ) - def test_single_pattern( - self, pattern: str, module_path: str, expected: str | None - ) -> None: - rewrites = compile_rewrite_imports({pattern: "pkg"}) - assert rewrite_module_path(module_path, rewrites) == expected + def test_pattern(self, pattern: str, *, matches: bool) -> None: + mappings = compile_map_imports({pattern: "pkg"}) + expected = "pkg" if matches else None + assert map_import_target("google/rpc/status.proto", mappings) == expected def test_first_match_wins(self) -> None: - rewrites = compile_rewrite_imports( - {"./foo/*_pb.py": "first", "./**/*_pb.py": "second"} - ) - assert rewrite_module_path(".foo.bar_pb", rewrites) == "first.foo.bar_pb" - assert rewrite_module_path(".other.bar_pb", rewrites) == "second.other.bar_pb" + mappings = compile_map_imports({"google/rpc/": "first", "google/": "second"}) + assert map_import_target("google/rpc/status.proto", mappings) == "first" + assert map_import_target("google/type/date.proto", mappings) == "second" def test_target_trailing_dot_stripped(self) -> None: - rewrites = compile_rewrite_imports({"./*_pb.py": "pkg."}) - assert rewrite_module_path(".bar_pb", rewrites) == "pkg.bar_pb" + mappings = compile_map_imports({"**": "pkg."}) + assert map_import_target("foo.proto", mappings) == "pkg" @pytest.mark.parametrize("target", ["", "."]) def test_empty_target_is_canonical(self, target: str) -> None: - rewrites = compile_rewrite_imports({"./**/*_pb.py": target}) - assert rewrite_module_path(".foo.bar_pb", rewrites) == "foo.bar_pb" - assert rewrite_module_path(".bar_pb", rewrites) == "bar_pb" + mappings = compile_map_imports({"**": target}) + assert map_import_target("foo.proto", mappings) == "" - def test_empty_target_absolute_raises(self) -> None: - rewrites = compile_rewrite_imports({"protobuf/*": ""}) - with pytest.raises(ValueError, match="rewrite_imports"): - rewrite_module_path("protobuf.wkt", rewrites) + @pytest.mark.parametrize( + "target", ["my-pkg", ".pkg", "..", "a..b", "pkg/sub", "1pkg", "pkg:x"] + ) + def test_invalid_target_raises(self, target: str) -> None: + with pytest.raises(ValueError, match="map_imports"): + compile_map_imports({"**": target}) + def test_empty_pattern_raises(self) -> None: + with pytest.raises(ValueError, match="map_imports"): + compile_map_imports({"": "pkg"}) -class TestFileRewrites: + +class TestFileMaps: def test_symbol_import(self, desc: DescFile) -> None: - f = _file(desc, {"./dep_pb.py": "mypkg.gen"}) + f = _file(desc, {"dep.proto": "mypkg.gen"}) f.print("x: ", desc.dependencies[0].messages[0]) assert gen_write(f, f.path) == dedent( """\ @@ -112,7 +98,7 @@ def test_symbol_import(self, desc: DescFile) -> None: ) def test_module_import(self, desc: DescFile) -> None: - f = _file(desc, {"./pkg/**/*_pb.py": "mypkg.gen"}) + f = _file(desc, {"pkg/": "mypkg.gen"}) f.print("d = ", desc.dependencies[1], ".desc()") assert gen_write(f, f.path) == dedent( """\ @@ -125,8 +111,8 @@ def test_module_import(self, desc: DescFile) -> None: """ ) - def test_own_symbols_not_rewritten(self, desc: DescFile) -> None: - f = _file(desc, {"./**/*_pb.py": "mypkg.gen"}) + def test_own_symbols_not_mapped(self, desc: DescFile) -> None: + f = _file(desc, {"**": "mypkg.gen"}) f.print("x: ", desc.messages[0]) f.print("y: ", desc.dependencies[0].messages[0]) assert gen_write(f, f.path) == dedent( @@ -142,7 +128,7 @@ def test_own_symbols_not_rewritten(self, desc: DescFile) -> None: ) def test_unmatched_import_stays_relative(self, desc: DescFile) -> None: - f = _file(desc, {"./pkg/**/*_pb.py": "mypkg.gen"}) + f = _file(desc, {"pkg/": "mypkg.gen"}) f.print("x: ", desc.dependencies[0].messages[0]) assert gen_write(f, f.path) == dedent( """\ @@ -156,7 +142,7 @@ def test_unmatched_import_stays_relative(self, desc: DescFile) -> None: ) def test_type_only_import(self, desc: DescFile) -> None: - f = _file(desc, {"./dep_pb.py": "mypkg.gen"}) + f = _file(desc, {"dep.proto": "mypkg.gen"}) f.print("x: ", Ident.for_desc(desc.dependencies[0].messages[0], type_only=True)) assert gen_write(f, f.path) == dedent( """\ @@ -172,14 +158,35 @@ def test_type_only_import(self, desc: DescFile) -> None: """ ) - def test_absolute_import_replaced(self, desc: DescFile) -> None: - f = _file(desc, {"protobuf/wkt": "vendored.wkt"}) - f.print("x: ", Module("protobuf.wkt").ident("Timestamp")) + def test_plain_ident_not_mapped(self, desc: DescFile) -> None: + f = _file(desc, {"**": "mypkg.gen"}) + f.print("x: ", Module(".dep_pb").ident("Dep")) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from .dep_pb import Dep + + + x: Dep + """ + ) + + def test_wkt_not_mapped(self, protoc: Protoc) -> None: + desc = protoc.compile_file( + """ + syntax = "proto3"; + import "google/protobuf/timestamp.proto"; + message Foo { google.protobuf.Timestamp ts = 1; } + """ + ) + f = _file(desc, {"google/": "mypkg.gen"}) + f.print("x: ", desc.dependencies[0].messages[0]) assert gen_write(f, f.path) == dedent( """\ from __future__ import annotations - from vendored.wkt import Timestamp + from protobuf.wkt import Timestamp x: Timestamp @@ -206,7 +213,7 @@ def test_canonical_external_dependency(self, protoc: Protoc) -> None: "include_imports", ) dep = files["buf/validate/validate.proto"] - f = _file(files["app/main.proto"], {"./buf/validate/**/*_pb.py": ""}) + f = _file(files["app/main.proto"], {"buf/validate/": ""}) f.print(dep, ".desc()") f.print("rule: ", dep.messages[0]) assert gen_write(f, f.path) == dedent( @@ -241,7 +248,7 @@ def test_canonical_root_external_dependency(self, protoc: Protoc) -> None: "include_imports", ) dep = files["dep.proto"] - f = _file(files["main.proto"], {"./dep_pb.py": ""}) + f = _file(files["main.proto"], {"dep.proto": ""}) f.print(dep, ".desc()") f.print("dep: ", dep.messages[0]) assert gen_write(f, f.path) == dedent( @@ -277,7 +284,7 @@ def desc(self, protoc: Protoc) -> DescFile: )["input.proto"] -def _file(desc: DescFile, rewrites: dict[str, str]) -> _File: +def _file(desc: DescFile, mappings: dict[str, str]) -> _File: module = Module.for_desc(desc, "_pb") return _File( path=f"{module.path.removeprefix('.').replace('.', '/')}.py", @@ -286,5 +293,5 @@ def _file(desc: DescFile, rewrites: dict[str, str]) -> _File: plugin_name="test", plugin_version="0.0.0", parameter="", - rewrite_imports=compile_rewrite_imports(rewrites), + map_imports=compile_map_imports(mappings), ) diff --git a/tests/plugin/test_protoc_gen_py.py b/tests/plugin/test_protoc_gen_py.py index b30ca9c..f97d7bc 100644 --- a/tests/plugin/test_protoc_gen_py.py +++ b/tests/plugin/test_protoc_gen_py.py @@ -148,7 +148,7 @@ def test_external_dependency(self, protoc: Protoc) -> None: """, }, files_to_generate=["app/main.proto"], - parameter="rewrite_imports=./buf/validate/**/*_pb.py:", + parameter="map_imports=buf/validate/:", ) assert resp.error == "" diff --git a/tests/plugin/test_run.py b/tests/plugin/test_run.py index cff82af..04675f9 100644 --- a/tests/plugin/test_run.py +++ b/tests/plugin/test_run.py @@ -220,7 +220,7 @@ def generate(schema: Schema[None]) -> None: assert len(resp.file) == 1 assert resp.file[0].name == "foo/bar_module/baz_service_pb.py" - def test_rewrite_imports(self, protoc: Protoc) -> None: + def test_map_imports(self, protoc: Protoc) -> None: def generate(schema: Schema[None]) -> None: for desc in schema.files_to_generate: f = schema.generate_file(desc, "_pb.py") @@ -234,14 +234,14 @@ def generate(schema: Schema[None]) -> None: "dep/dep.proto": 'syntax = "proto3"; package dep; message Dep {}', }, files_to_generate=["main.proto"], - parameter="rewrite_imports=./dep/**/*_pb.py:mypkg.gen", + parameter="map_imports=dep/:mypkg.gen", ) assert resp.error == "" content = resp.file[0].content assert "from mypkg.gen.dep import dep_pb" in content assert "from mypkg.gen.dep.dep_pb import Dep" in content - def test_rewrite_imports_in_preamble(self, protoc: Protoc) -> None: + def test_map_imports_in_preamble(self, protoc: Protoc) -> None: def generate(schema: Schema[None]) -> None: for desc in schema.files_to_generate: f = schema.generate_file(desc, "_pb.py") @@ -250,22 +250,20 @@ def generate(schema: Schema[None]) -> None: resp = protoc.run_plugin( Plugin(generate), {"test.proto": 'syntax = "proto3";'}, - parameter="no_fmt_off,rewrite_imports=./**/*_pb.py:mypkg", + parameter="no_fmt_off,map_imports=**:mypkg", ) assert resp.error == "" content = resp.file[0].content - assert ( - 'with parameter "no_fmt_off,rewrite_imports=./**/*_pb.py:mypkg"' in content - ) + assert 'with parameter "no_fmt_off,map_imports=**:mypkg"' in content - def test_rewrite_imports_without_colon_is_error(self, protoc: Protoc) -> None: + def test_map_imports_without_colon_is_error(self, protoc: Protoc) -> None: resp = protoc.run_plugin( Plugin(lambda _: None), {"test.proto": 'syntax = "proto3";'}, - parameter="rewrite_imports=./foo_pb.py", + parameter="map_imports=foo.proto", ) assert resp.error != "" - assert "rewrite_imports" in resp.error + assert "map_imports" in resp.error def test_invalid_framework_option_value_is_error(self, protoc: Protoc) -> None: resp = protoc.run_plugin( diff --git a/tests/test_gen_buf.py b/tests/test_gen_buf.py index c889f38..541a681 100644 --- a/tests/test_gen_buf.py +++ b/tests/test_gen_buf.py @@ -17,7 +17,7 @@ The proto_buf module depends on buf.build/googleapis/googleapis, which is not generated into gen_buf. Instead, its generated code is provided by the googleapis-googleapis-bufbuild-py package, and gen_buf is generated with -rewrite_imports mapping google/** to the canonical import path it provides. +map_imports mapping google/ to the canonical import path it provides. """ from __future__ import annotations