"""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 ``_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=``. """ 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()