diff --git a/docs/writing-plugins/options.md b/docs/writing-plugins/options.md index 51e2105..db72899 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` and `escape_module_with_hash`. +The framework reserves certain option names for all plugins, currently: `no_fmt_off`, `escape_module_with_hash`, and `rewrite_imports`. If your `Options` dataclass defines a field with a reserved name, `run()` raises a `ValueError`. ### no_fmt_off @@ -66,6 +66,46 @@ plugins: opt: escape_module_with_hash ``` +### rewrite_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. + +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: + +```yaml title="buf.gen.yaml" +plugins: + - local: protoc-gen-hello + out: src/gen + opt: rewrite_imports=./google/type/**/*_pb.py: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. + +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`. +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: +``` + +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. + +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. + ## Example: Sensitive fields plugin Here is a plugin that generates a `_sensitive.py` file for each proto file, listing all fields marked with the `sensitive` custom option from [extensions](../extensions.md#extensions-in-custom-options): diff --git a/pyproject.toml b/pyproject.toml index ab9201e..24fe86e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,7 @@ dev-required = [ "upstream-protobuf", "buf-bin==1.72.0", + "googleapis-googleapis-bufbuild-py==0.3.0.2.20260414192239+c17df5b2beca", "hypothesis==6.167.1", "license-header==0.0.1", "poethepoet==0.48.0", @@ -476,6 +477,11 @@ exclude = [ reinstall-package = ["protobuf-py-ext"] no-build-isolation-package = ["protobuf-py-ext"] +[[tool.uv.index]] +name = "buf" +url = "https://buf.build/gen/python" +explicit = true + [tool.uv.build-backend] module-name = ["protobuf"] @@ -488,6 +494,8 @@ example-plugin = { workspace = true } protobuf-py-ext = { workspace = true } protobuf-py-bench = { workspace = true } +googleapis-googleapis-bufbuild-py = { index = "buf" } + [tool.uv.workspace] members = [ "packages/protobuf-py-ext", diff --git a/src/protobuf/plugin/_file.py b/src/protobuf/plugin/_file.py index 92a5667..70c7341 100644 --- a/src/protobuf/plugin/_file.py +++ b/src/protobuf/plugin/_file.py @@ -22,10 +22,13 @@ from protobuf import DescEnum, DescExtension, DescFile, DescMessage, ScalarType from protobuf.plugin._ident import Ident, Module +from protobuf.plugin._rewrite_imports import rewrite_module_path if TYPE_CHECKING: from collections.abc import Generator, Iterable, Iterator + from protobuf.plugin._rewrite_imports import RewriteImports + _INDENT = " " * 4 _TYPING = Module("typing") @@ -220,6 +223,7 @@ def __init__( parameter: str, *, escape_module_with_hash: bool = False, + rewrite_imports: RewriteImports = (), ) -> None: self.path = path self.module = module @@ -228,6 +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._indent = 0 self._type_checking = False self._in_doc = False @@ -329,19 +334,48 @@ def _to_el(self, v: object) -> str | Ident: ) case _: return repr(v) - ident = self._relativize(ident) + ident = self._relativize(self._rewrite_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: + return ident + # A DescFile ident imports the module itself, so its import + # path includes the ident name. + is_module_import = isinstance(ident._desc, DescFile) + module_path = ident.module.path + if is_module_import: + 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 + if is_module_import: + parent, _, name = rewritten.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 + ) + def _relativize(self, ident: Ident) -> Ident: if not _is_relative(ident.module): return Ident( ident.name, ident.module, type_only=ident.type_only or self._type_checking, + _desc=ident._desc, ) self_segments = _module_segments(self.module) @@ -486,8 +520,24 @@ def _write_imports( # Ruff splits the imports into groups of std, global, and relative. We also do the same: for group in _group_and_sort_imports(imports): for module, idents in group.items(): - deduped = sorted({aliases.resolve_import(ident) for ident in idents}) - lines.append(f"{indent}from {module.path} import {', '.join(deduped)}") + module_imports = [ + ident for ident in idents if _is_module_import(module, ident) + ] + lines.extend( + f"{indent}import {import_}" + for import_ in sorted( + {aliases.resolve_import(ident) for ident in module_imports} + ) + ) + + from_imports = [ + ident for ident in idents if not _is_module_import(module, ident) + ] + if from_imports: + deduped = sorted( + {aliases.resolve_import(ident) for ident in from_imports} + ) + lines.append(f"{indent}from {module.path} import {', '.join(deduped)}") lines.append("") @@ -600,6 +650,11 @@ def _is_relative(module: Module) -> bool: return module.path.startswith(".") +def _is_module_import(module: Module, ident: Ident) -> bool: + """Return True if the identifier imports a top-level module itself.""" + return isinstance(ident._desc, DescFile) and module.path == ident.name + + def _module_segments(module: Module) -> list[str]: path = module.path.removeprefix(".") return path.split(".") if path else [] diff --git a/src/protobuf/plugin/_rewrite_imports.py b/src/protobuf/plugin/_rewrite_imports.py new file mode 100644 index 0000000..8fbb806 --- /dev/null +++ b/src/protobuf/plugin/_rewrite_imports.py @@ -0,0 +1,64 @@ +# 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 2a59f1f..5d90a31 100644 --- a/src/protobuf/plugin/_run.py +++ b/src/protobuf/plugin/_run.py @@ -21,6 +21,7 @@ from protobuf import maximum_supported_edition, minimum_supported_edition from protobuf.plugin._file import write 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 @@ -167,6 +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), ) generate(schema) @@ -230,3 +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) diff --git a/src/protobuf/plugin/_schema.py b/src/protobuf/plugin/_schema.py index 657a3c9..6f3197a 100644 --- a/src/protobuf/plugin/_schema.py +++ b/src/protobuf/plugin/_schema.py @@ -27,6 +27,8 @@ from protobuf.wkt import CodeGeneratorRequest + from ._rewrite_imports import RewriteImports + T_co = TypeVar("T_co", covariant=True) @@ -96,6 +98,7 @@ def __init__( name: str, version: str, escape_module_with_hash: bool, + rewrite_imports: RewriteImports = (), ) -> None: self._options = options self._name = name @@ -103,6 +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 file_to_generate = frozenset(req.file_to_generate) source_by_name = {s.name: s for s in req.source_file_descriptors} @@ -166,6 +170,7 @@ def generate_file( self._version, self._parameter, escape_module_with_hash=self._escape_module_with_hash, + rewrite_imports=self._rewrite_imports, ) self._generated_files[path] = f return f diff --git a/tests/buf.gen.yaml b/tests/buf.gen.yaml index f3d3780..3485bc1 100644 --- a/tests/buf.gen.yaml +++ b/tests/buf.gen.yaml @@ -3,3 +3,4 @@ clean: true plugins: - local: protoc-gen-py out: gen_buf + opt: "rewrite_imports=./google/**/*_pb.py:" diff --git a/tests/buf.lock b/tests/buf.lock new file mode 100644 index 0000000..8447589 --- /dev/null +++ b/tests/buf.lock @@ -0,0 +1,6 @@ +# Generated by buf. DO NOT EDIT. +version: v2 +deps: + - name: buf.build/googleapis/googleapis + commit: c17df5b2beca46928cc87d5656bd5343 + digest: b5:648a01e0170d4512dea7d564016165decd1ed6e34bef79fe54753e51ad7e27545709ad9157d7551270147d551155c595a2fb0bf5bb33b1c83040ddbce915c604 diff --git a/tests/buf.yaml b/tests/buf.yaml index 9acfb79..4354aa8 100644 --- a/tests/buf.yaml +++ b/tests/buf.yaml @@ -1,3 +1,5 @@ version: v2 modules: - path: proto_buf +deps: + - buf.build/googleapis/googleapis diff --git a/tests/gen_buf/__init__.py b/tests/gen_buf/__init__.py new file mode 100644 index 0000000..e7c353f --- /dev/null +++ b/tests/gen_buf/__init__.py @@ -0,0 +1,14 @@ +# 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 diff --git a/tests/gen_buf/escaping_pb.py b/tests/gen_buf/escaping_pb.py index 88f2495..9cecc7e 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.3.0 with parameter "". +# Generated by protoc-gen-py v0.3.0 with parameter "rewrite_imports=./google/**/*_pb.py:". # 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 84c85e6..69ab501 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.3.0 with parameter "". +# Generated by protoc-gen-py v0.3.0 with parameter "rewrite_imports=./google/**/*_pb.py:". # ruff: noqa: PGH004 # ruff: noqa # fmt: off diff --git a/tests/gen_buf/local_dep/__init__.py b/tests/gen_buf/local_dep/__init__.py new file mode 100644 index 0000000..e7c353f --- /dev/null +++ b/tests/gen_buf/local_dep/__init__.py @@ -0,0 +1,14 @@ +# 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 diff --git a/tests/gen_buf/local_dep/dep_pb.py b/tests/gen_buf/local_dep/dep_pb.py new file mode 100644 index 0000000..49aaae0 --- /dev/null +++ b/tests/gen_buf/local_dep/dep_pb.py @@ -0,0 +1,82 @@ +# 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. + +# Generated from local_dep/dep.proto. DO NOT EDIT. +# Generated by protoc-gen-py v0.3.0 with parameter "rewrite_imports=./google/**/*_pb.py:". +# ruff: noqa: PGH004 +# ruff: noqa +# fmt: off + +from __future__ import annotations + +from typing import Literal, TYPE_CHECKING, TypeAlias + +from google.rpc import status_pb +from protobuf import Message +from protobuf._codegen import file_desc + +if TYPE_CHECKING: + from google.rpc.status_pb import Status + from protobuf import DescFile + + +_DepFields: TypeAlias = Literal["value", "status"] + +class Dep(Message[_DepFields]): + """ + ```proto + message local_dep.Dep + ``` + + Attributes: + value: + ```proto + string value = 1; + ``` + status: + ```proto + optional google.rpc.Status status = 2; + ``` + """ + + __slots__ = ("value", "status") + + if TYPE_CHECKING: + + def __init__( + self, + *, + value: str = "", + status: Status | None = None, + ) -> None: + pass + + value: str + status: Status | None + + +_DESC = file_desc( + b'\n\x13local_dep/dep.proto\x12\tlocal_dep\x1a\x17google/rpc/status.proto"G\n\x03Dep\x12\x14\n\x05value\x18\x01 \x01(\tR\x05value\x12*\n\x06status\x18\x02 \x01(\x0b2\x12.google.rpc.StatusR\x06statusb\x06proto3', + [ + status_pb.desc(), + ], + { + "Dep": Dep, + }, +) + + +def desc() -> DescFile: + """Returns the descriptor for the file `local_dep/dep.proto`.""" + return _DESC diff --git a/tests/gen_buf/local_import/__init__.py b/tests/gen_buf/local_import/__init__.py new file mode 100644 index 0000000..e7c353f --- /dev/null +++ b/tests/gen_buf/local_import/__init__.py @@ -0,0 +1,14 @@ +# 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 diff --git a/tests/gen_buf/local_import/importer_pb.py b/tests/gen_buf/local_import/importer_pb.py new file mode 100644 index 0000000..11393aa --- /dev/null +++ b/tests/gen_buf/local_import/importer_pb.py @@ -0,0 +1,78 @@ +# 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. + +# Generated from local_import/importer.proto. DO NOT EDIT. +# Generated by protoc-gen-py v0.3.0 with parameter "rewrite_imports=./google/**/*_pb.py:". +# ruff: noqa: PGH004 +# ruff: noqa +# fmt: off + +from __future__ import annotations + +from typing import Literal, TYPE_CHECKING, TypeAlias + +from protobuf import Message +from protobuf._codegen import file_desc + +from ..local_dep import dep_pb + +if TYPE_CHECKING: + from protobuf import DescFile + + from ..local_dep.dep_pb import Dep + + +_ImporterFields: TypeAlias = Literal["dep"] + +class Importer(Message[_ImporterFields]): + """ + ```proto + message local_import.Importer + ``` + + Attributes: + dep: + ```proto + optional local_dep.Dep dep = 1; + ``` + """ + + __slots__ = ("dep",) + + if TYPE_CHECKING: + + def __init__( + self, + *, + dep: Dep | None = None, + ) -> None: + pass + + dep: Dep | None + + +_DESC = file_desc( + b'\n\x1blocal_import/importer.proto\x12\x0clocal_import\x1a\x13local_dep/dep.proto",\n\x08Importer\x12 \n\x03dep\x18\x01 \x01(\x0b2\x0e.local_dep.DepR\x03depb\x06proto3', + [ + dep_pb.desc(), + ], + { + "Importer": Importer, + }, +) + + +def desc() -> DescFile: + """Returns the descriptor for the file `local_import/importer.proto`.""" + return _DESC diff --git a/tests/plugin/test_protoc_gen_py.py b/tests/plugin/test_protoc_gen_py.py index 18a8f8e..b30ca9c 100644 --- a/tests/plugin/test_protoc_gen_py.py +++ b/tests/plugin/test_protoc_gen_py.py @@ -128,3 +128,33 @@ def desc() -> DescFile: ''' ).lstrip() ) + + def test_external_dependency(self, protoc: Protoc) -> None: + resp = protoc.run_plugin( + Plugin(_generate, options=_Options), + { + "app/main.proto": """ + syntax = "proto3"; + package app; + import "buf/validate/validate.proto"; + message Main { + buf.validate.Rule rule = 1; + } + """, + "buf/validate/validate.proto": """ + syntax = "proto3"; + package buf.validate; + message Rule {} + """, + }, + files_to_generate=["app/main.proto"], + parameter="rewrite_imports=./buf/validate/**/*_pb.py:", + ) + + assert resp.error == "" + files = {f.name: f.content for f in resp.file} + assert set(files) == {"__init__.py", "app/__init__.py", "app/main_pb.py"} + main = files["app/main_pb.py"] + assert "from buf.validate import validate_pb" in main + assert "from buf.validate.validate_pb import Rule" in main + assert "from ..buf.validate" not in main diff --git a/tests/plugin/test_rewrite_imports.py b/tests/plugin/test_rewrite_imports.py new file mode 100644 index 0000000..63dd73d --- /dev/null +++ b/tests/plugin/test_rewrite_imports.py @@ -0,0 +1,290 @@ +# 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 + +from textwrap import dedent +from typing import TYPE_CHECKING + +import pytest + +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, +) + +if TYPE_CHECKING: + from protobuf import DescFile + from tests.conftest import Protoc + + +class TestRewriteModulePath: + @pytest.mark.parametrize( + ("pattern", "module_path", "expected"), + [ + 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" + ), + ], + ) + 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_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" + + 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" + + @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" + + 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) + + +class TestFileRewrites: + def test_symbol_import(self, desc: DescFile) -> None: + f = _file(desc, {"./dep_pb.py": "mypkg.gen"}) + f.print("x: ", desc.dependencies[0].messages[0]) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from mypkg.gen.dep_pb import Dep + + + x: Dep + """ + ) + + def test_module_import(self, desc: DescFile) -> None: + f = _file(desc, {"./pkg/**/*_pb.py": "mypkg.gen"}) + f.print("d = ", desc.dependencies[1], ".desc()") + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from mypkg.gen.pkg import nested_pb + + + d = nested_pb.desc() + """ + ) + + def test_own_symbols_not_rewritten(self, desc: DescFile) -> None: + f = _file(desc, {"./**/*_pb.py": "mypkg.gen"}) + f.print("x: ", desc.messages[0]) + f.print("y: ", desc.dependencies[0].messages[0]) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from mypkg.gen.dep_pb import Dep + + + x: Foo + y: Dep + """ + ) + + def test_unmatched_import_stays_relative(self, desc: DescFile) -> None: + f = _file(desc, {"./pkg/**/*_pb.py": "mypkg.gen"}) + f.print("x: ", desc.dependencies[0].messages[0]) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from .dep_pb import Dep + + + x: Dep + """ + ) + + def test_type_only_import(self, desc: DescFile) -> None: + f = _file(desc, {"./dep_pb.py": "mypkg.gen"}) + f.print("x: ", Ident.for_desc(desc.dependencies[0].messages[0], type_only=True)) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from typing import TYPE_CHECKING + + if TYPE_CHECKING: + from mypkg.gen.dep_pb import Dep + + + x: Dep + """ + ) + + def test_absolute_import_replaced(self, desc: DescFile) -> None: + f = _file(desc, {"protobuf/wkt": "vendored.wkt"}) + f.print("x: ", Module("protobuf.wkt").ident("Timestamp")) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from vendored.wkt import Timestamp + + + x: Timestamp + """ + ) + + def test_canonical_external_dependency(self, protoc: Protoc) -> None: + files = protoc.compile( + { + "app/main.proto": """ + syntax = "proto3"; + package app; + import "buf/validate/validate.proto"; + message Main { + buf.validate.Rule rule = 1; + } + """, + "buf/validate/validate.proto": """ + syntax = "proto3"; + package buf.validate; + message Rule {} + """, + }, + "include_imports", + ) + dep = files["buf/validate/validate.proto"] + f = _file(files["app/main.proto"], {"./buf/validate/**/*_pb.py": ""}) + f.print(dep, ".desc()") + f.print("rule: ", dep.messages[0]) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + from buf.validate import validate_pb + from buf.validate.validate_pb import Rule + + + validate_pb.desc() + rule: Rule + """ + ) + + def test_canonical_root_external_dependency(self, protoc: Protoc) -> None: + """A root-level external dependency becomes a plain `import X`.""" + files = protoc.compile( + { + "main.proto": """ + syntax = "proto3"; + import "dep.proto"; + message Main { + Dep dep = 1; + } + """, + "dep.proto": """ + syntax = "proto3"; + message Dep {} + """, + }, + "include_imports", + ) + dep = files["dep.proto"] + f = _file(files["main.proto"], {"./dep_pb.py": ""}) + f.print(dep, ".desc()") + f.print("dep: ", dep.messages[0]) + assert gen_write(f, f.path) == dedent( + """\ + from __future__ import annotations + + import dep_pb + from dep_pb import Dep + + + dep_pb.desc() + dep: Dep + """ + ) + + @pytest.fixture + def desc(self, protoc: Protoc) -> DescFile: + return protoc.compile( + { + "input.proto": """ + syntax = "proto3"; + import "dep.proto"; + import "pkg/nested.proto"; + message Foo { + Dep dep = 1; + pkg.Nested nested = 2; + } + """, + "dep.proto": 'syntax = "proto3"; message Dep {}', + "pkg/nested.proto": 'syntax = "proto3"; package pkg; message Nested {}', + }, + "include_imports", + )["input.proto"] + + +def _file(desc: DescFile, rewrites: dict[str, str]) -> _File: + module = Module.for_desc(desc, "_pb") + return _File( + path=f"{module.path.removeprefix('.').replace('.', '/')}.py", + module=module, + file_to_generate=frozenset(), + plugin_name="test", + plugin_version="0.0.0", + parameter="", + rewrite_imports=compile_rewrite_imports(rewrites), + ) diff --git a/tests/plugin/test_run.py b/tests/plugin/test_run.py index 1686393..cff82af 100644 --- a/tests/plugin/test_run.py +++ b/tests/plugin/test_run.py @@ -220,6 +220,53 @@ 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 generate(schema: Schema[None]) -> None: + for desc in schema.files_to_generate: + f = schema.generate_file(desc, "_pb.py") + f.print("d = ", desc.dependencies[0], ".desc()") + f.print("x: ", desc.dependencies[0].messages[0]) + + resp = protoc.run_plugin( + Plugin(generate), + { + "main.proto": 'syntax = "proto3"; import "dep/dep.proto"; message Main { dep.Dep d = 1; }', + "dep/dep.proto": 'syntax = "proto3"; package dep; message Dep {}', + }, + files_to_generate=["main.proto"], + parameter="rewrite_imports=./dep/**/*_pb.py: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 generate(schema: Schema[None]) -> None: + for desc in schema.files_to_generate: + f = schema.generate_file(desc, "_pb.py") + f.preamble(desc) + + resp = protoc.run_plugin( + Plugin(generate), + {"test.proto": 'syntax = "proto3";'}, + parameter="no_fmt_off,rewrite_imports=./**/*_pb.py:mypkg", + ) + assert resp.error == "" + content = resp.file[0].content + assert ( + 'with parameter "no_fmt_off,rewrite_imports=./**/*_pb.py:mypkg"' in content + ) + + def test_rewrite_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", + ) + assert resp.error != "" + assert "rewrite_imports" in resp.error + def test_invalid_framework_option_value_is_error(self, protoc: Protoc) -> None: resp = protoc.run_plugin( Plugin(lambda _: None), diff --git a/tests/proto_buf/local_dep/dep.proto b/tests/proto_buf/local_dep/dep.proto new file mode 100644 index 0000000..f001aac --- /dev/null +++ b/tests/proto_buf/local_dep/dep.proto @@ -0,0 +1,24 @@ +// 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. +syntax = "proto3"; + +package local_dep; + +import "google/rpc/status.proto"; + +message Dep { + string value = 1; + + google.rpc.Status status = 2; +} diff --git a/tests/proto_buf/local_import/importer.proto b/tests/proto_buf/local_import/importer.proto new file mode 100644 index 0000000..ede2225 --- /dev/null +++ b/tests/proto_buf/local_import/importer.proto @@ -0,0 +1,22 @@ +// 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. +syntax = "proto3"; + +package local_import; + +import "local_dep/dep.proto"; + +message Importer { + local_dep.Dep dep = 1; +} diff --git a/tests/test_gen_buf.py b/tests/test_gen_buf.py new file mode 100644 index 0000000..c889f38 --- /dev/null +++ b/tests/test_gen_buf.py @@ -0,0 +1,43 @@ +# 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. + +"""Tests for buf-generated code in gen_buf. + +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. +""" + +from __future__ import annotations + +from google.rpc.status_pb import Status + +from tests.gen_buf.local_dep.dep_pb import Dep +from tests.gen_buf.local_import.importer_pb import Importer + + +def test_external_dependency_roundtrip() -> None: + msg = Importer(dep=Dep(value="v", status=Status(code=3, message="oops"))) + decoded = Importer.from_binary(msg.to_binary()) + assert decoded.dep is not None + assert decoded.dep.value == "v" + assert decoded.dep.status is not None + assert decoded.dep.status.code == 3 + assert decoded.dep.status.message == "oops" + + +def test_external_dependency_descriptor() -> None: + deps = Dep.desc().file.dependencies + assert [d.name for d in deps] == ["google/rpc/status.proto"] diff --git a/uv.lock b/uv.lock index d1e2da1..6729286 100644 --- a/uv.lock +++ b/uv.lock @@ -294,6 +294,17 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f7/ec/67fbef5d497f86283db54c22eec6f6140243aae73265799baaaa19cd17fb/ghp_import-2.1.0-py3-none-any.whl", hash = "sha256:8337dd7b50877f163d4c0289bc1f1c7f127550241988d568c1db512c4324a619", size = 11034, upload-time = "2022-05-02T15:47:14.552Z" }, ] +[[package]] +name = "googleapis-googleapis-bufbuild-py" +version = "0.3.0.2.20260414192239+c17df5b2beca" +source = { registry = "https://buf.build/gen/python" } +dependencies = [ + { name = "protobuf-py" }, +] +wheels = [ + { url = "https://buf.build/gen/python/googleapis-googleapis-bufbuild-py/googleapis_googleapis_bufbuild_py-0.3.0.2.20260414192239+c17df5b2beca-py3-none-any.whl", upload-time = "2026-08-10T18:49:41Z" }, +] + [[package]] name = "griffelib" version = "2.2.0" @@ -736,6 +747,7 @@ dev = [ { name = "example" }, { name = "example-plugin" }, { name = "fix-protobuf-imports" }, + { name = "googleapis-googleapis-bufbuild-py" }, { name = "hypothesis" }, { name = "license-header" }, { name = "maturin" }, @@ -763,6 +775,7 @@ dev-required = [ { name = "buf-bin" }, { name = "example" }, { name = "example-plugin" }, + { name = "googleapis-googleapis-bufbuild-py" }, { name = "hypothesis" }, { name = "license-header" }, { name = "poethepoet" }, @@ -799,6 +812,7 @@ dev = [ { name = "example", editable = "examples/protobuf" }, { name = "example-plugin", editable = "examples/plugin" }, { name = "fix-protobuf-imports", specifier = "==0.1.7" }, + { name = "googleapis-googleapis-bufbuild-py", specifier = "==0.3.0.2.20260414192239+c17df5b2beca", index = "https://buf.build/gen/python" }, { name = "hypothesis", specifier = "==6.167.1" }, { name = "license-header", specifier = "==0.0.1" }, { name = "maturin", specifier = "==1.15.0" }, @@ -824,6 +838,7 @@ dev-required = [ { name = "buf-bin", specifier = "==1.72.0" }, { name = "example", editable = "examples/protobuf" }, { name = "example-plugin", editable = "examples/plugin" }, + { name = "googleapis-googleapis-bufbuild-py", specifier = "==0.3.0.2.20260414192239+c17df5b2beca", index = "https://buf.build/gen/python" }, { name = "hypothesis", specifier = "==6.167.1" }, { name = "license-header", specifier = "==0.0.1" }, { name = "poethepoet", specifier = "==0.48.0" },