| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216 |
- """A `protoc` plugin that generates chattolib's Connect client stubs.
- Speaks the standard plugin protocol: reads a ``CodeGeneratorRequest`` from
- stdin, writes a ``CodeGeneratorResponse`` to stdout. For each proto that
- declares a ``service``, it emits one ``<name>_connect.py`` defining an async
- ``*ServiceClient(ConnectClient)`` and a sync ``*ServiceClientSync
- (ConnectClientSync)`` — both delegating to the shared runtime in
- ``chattolib._connect``.
- Installed as the console script ``protoc-gen-chattolib`` so ``protoc`` finds
- it on PATH and it is invoked with ``--chattolib_out=<dir>``.
- """
- from __future__ import annotations
- import importlib.util
- import os
- import sys
- from typing import Any
- _CODEGEN_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "_codegen")
- def _load_codegen_module(name: str) -> Any:
- """Load a vendored codegen descriptor module directly from its file path.
- The descriptors live under ``chattolib._codegen``, but importing them through
- the ``chattolib`` package would execute ``chattolib/__init__.py`` (which
- eagerly imports the client) — undesirable for a standalone ``protoc`` plugin
- process. Loading by file path avoids that.
- """
- path = os.path.join(_CODEGEN_DIR, f"{name}.py")
- spec = importlib.util.spec_from_file_location(f"_chattolib_codegen_{name}", path)
- assert spec is not None
- mod = importlib.util.module_from_spec(spec)
- sys.modules[spec.name] = mod
- assert spec.loader is not None
- spec.loader.exec_module(mod)
- return mod
- def _load_descriptors() -> tuple[Any, Any]:
- return (
- _load_codegen_module("protoc_gen_request_pb2"),
- _load_codegen_module("protoc_gen_response_pb2"),
- )
- _g, _r = _load_descriptors()
- g: Any = _g
- r: Any = _r
- _HEADER = "# Generated by chattolib's protoc plugin (protoc-gen-chattolib). DO NOT EDIT!\n"
- def _method_name(m: Any) -> str:
- """snake_case the RPC name, matching Python method naming."""
- out = []
- for i, ch in enumerate(m.name):
- if ch.isupper() and i:
- out.append("_")
- out.append(ch.lower())
- return "".join(out)
- def _pb2_module(fd: Any) -> str:
- """The import path of the pb2 module protoc emits for ``fd``."""
- name: str = fd.name
- return name[: -len(".proto")].replace("/", ".") + "_pb2"
- def _alias_for(module: str) -> str:
- """A valid, unique import alias for a pb2 module (dots -> underscores)."""
- return module.replace(".", "_")
- def _build_module_map(request: Any) -> dict[str, str]:
- """Map a fully-qualified message name (leading dot) to the pb2 module of
- the file that defines it.
- A service method's input/output may be a message declared in a *different*
- proto file than the service itself (e.g. ``MyAccountService`` in
- ``account.proto`` using ``ListExternalIdentitiesRequest`` from
- ``external_identities.proto``). The generated stub must import the module
- that actually defines the referenced message, not the service's own module.
- """
- mapping: dict[str, str] = {}
- for fd in request.proto_file:
- module = _pb2_module(fd)
- def walk(msgs: list[Any], prefix: list[str]) -> None:
- for mt in msgs:
- full = "." + ".".join([p for p in (fd.package, *prefix, mt.name) if p])
- mapping[full] = module
- walk(list(mt.nested_type), prefix + [mt.name])
- walk(list(fd.message_type), [])
- return mapping
- def _service_preamble(fd: Any, modules: list[str]) -> str:
- """The shared header + imports for a generated module (emitted once).
- ``modules`` is the (deduplicated) set of pb2 modules the file's services
- reference; each is imported under a unique alias so cross-file messages
- resolve to the module that defines them.
- """
- lines = [_HEADER, f"# source: {fd.name}", ""]
- for module in modules:
- lines.append(f"import {module} as {_alias_for(module)}")
- lines.append("from chattolib._connect import ConnectClient, ConnectClientSync, MethodInfo")
- lines += ["", ""]
- return "\n".join(lines)
- def _service_body(fd: Any, svc: Any, mod_by_fullname: dict[str, str]) -> str:
- """Render one service's async + sync client classes (no preamble)."""
- pkg = fd.package
- base = f"{pkg}.{svc.name}" if pkg else svc.name
- lines: list[str] = [f"class {svc.name}Client(ConnectClient):"]
- for m in svc.method:
- mn = _method_name(m)
- in_module = mod_by_fullname[m.input_type]
- out_module = mod_by_fullname[m.output_type]
- in_name = m.input_type.lstrip(".").split(".")[-1]
- out_name = m.output_type.lstrip(".").split(".")[-1]
- in_alias = _alias_for(in_module)
- out_alias = _alias_for(out_module)
- lines += [
- f" async def {mn}(self, request: {in_alias}.{in_name}, *,"
- f" headers: dict[str, str] | None = None) -> {out_alias}.{out_name}:",
- " return await self.execute_unary(",
- " request=request,",
- " method=MethodInfo(",
- f' name="{m.name}",',
- f' service_name="{base}",',
- f" input={in_alias}.{in_name},",
- f" output={out_alias}.{out_name},",
- " ),",
- " headers=headers,",
- " )",
- "",
- ]
- lines += [f"class {svc.name}ClientSync(ConnectClientSync):", ""]
- for m in svc.method:
- mn = _method_name(m)
- in_module = mod_by_fullname[m.input_type]
- out_module = mod_by_fullname[m.output_type]
- in_name = m.input_type.lstrip(".").split(".")[-1]
- out_name = m.output_type.lstrip(".").split(".")[-1]
- in_alias = _alias_for(in_module)
- out_alias = _alias_for(out_module)
- lines += [
- f" def {mn}(self, request: {in_alias}.{in_name}, *,"
- f" headers: dict[str, str] | None = None) -> {out_alias}.{out_name}:",
- " return self.execute_unary(",
- " request=request,",
- " method=MethodInfo(",
- f' name="{m.name}",',
- f' service_name="{base}",',
- f" input={in_alias}.{in_name},",
- f" output={out_alias}.{out_name},",
- " ),",
- " headers=headers,",
- " )",
- "",
- ]
- return "\n".join(lines)
- def generate(request: g.CodeGeneratorRequest) -> r.CodeGeneratorResponse:
- resp = r.CodeGeneratorResponse()
- # Tell protoc (>= the version that checks this) we can handle proto3
- # `optional` fields. Without this bit, recent protoc refuses to run the
- # plugin on any proto3 file that declares an optional field.
- resp.supported_features = r.CodeGeneratorResponse.Feature.FEATURE_PROTO3_OPTIONAL
- # file_to_generate is a repeated *string* (proto file names), not a list of
- # message objects — match by string membership.
- to_generate = set(request.file_to_generate)
- mod_by_fullname = _build_module_map(request)
- for fd in request.proto_file:
- if fd.name not in to_generate:
- continue
- services = [svc for svc in fd.service if svc.method]
- if not services:
- continue
- # One file per proto, containing every service it declares (e.g.
- # notifications.proto -> NotificationServiceClient and
- # NotificationPolicyServiceClient in the same module), matching the
- # historical connectrpc layout.
- required: list[str] = []
- seen: set[str] = set()
- for svc in services:
- for m in svc.method:
- for ref in (m.input_type, m.output_type):
- module = mod_by_fullname[ref]
- if module not in seen:
- seen.add(module)
- required.append(module)
- out = resp.file.add()
- out.name = fd.name.replace(".proto", "_connect.py")
- out.content = _service_preamble(fd, required) + "\n".join(
- _service_body(fd, svc, mod_by_fullname) for svc in services
- )
- return resp
- def main() -> None:
- data = sys.stdin.buffer.read()
- request = g.CodeGeneratorRequest.FromString(data)
- resp = generate(request)
- sys.stdout.buffer.write(resp.SerializeToString())
- if __name__ == "__main__":
- main()
|