_protoc_plugin.py 8.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216
  1. """A `protoc` plugin that generates chattolib's Connect client stubs.
  2. Speaks the standard plugin protocol: reads a ``CodeGeneratorRequest`` from
  3. stdin, writes a ``CodeGeneratorResponse`` to stdout. For each proto that
  4. declares a ``service``, it emits one ``<name>_connect.py`` defining an async
  5. ``*ServiceClient(ConnectClient)`` and a sync ``*ServiceClientSync
  6. (ConnectClientSync)`` — both delegating to the shared runtime in
  7. ``chattolib._connect``.
  8. Installed as the console script ``protoc-gen-chattolib`` so ``protoc`` finds
  9. it on PATH and it is invoked with ``--chattolib_out=<dir>``.
  10. """
  11. from __future__ import annotations
  12. import importlib.util
  13. import os
  14. import sys
  15. from typing import Any
  16. _CODEGEN_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "_codegen")
  17. def _load_codegen_module(name: str) -> Any:
  18. """Load a vendored codegen descriptor module directly from its file path.
  19. The descriptors live under ``chattolib._codegen``, but importing them through
  20. the ``chattolib`` package would execute ``chattolib/__init__.py`` (which
  21. eagerly imports the client) — undesirable for a standalone ``protoc`` plugin
  22. process. Loading by file path avoids that.
  23. """
  24. path = os.path.join(_CODEGEN_DIR, f"{name}.py")
  25. spec = importlib.util.spec_from_file_location(f"_chattolib_codegen_{name}", path)
  26. assert spec is not None
  27. mod = importlib.util.module_from_spec(spec)
  28. sys.modules[spec.name] = mod
  29. assert spec.loader is not None
  30. spec.loader.exec_module(mod)
  31. return mod
  32. def _load_descriptors() -> tuple[Any, Any]:
  33. return (
  34. _load_codegen_module("protoc_gen_request_pb2"),
  35. _load_codegen_module("protoc_gen_response_pb2"),
  36. )
  37. _g, _r = _load_descriptors()
  38. g: Any = _g
  39. r: Any = _r
  40. _HEADER = "# Generated by chattolib's protoc plugin (protoc-gen-chattolib). DO NOT EDIT!\n"
  41. def _method_name(m: Any) -> str:
  42. """snake_case the RPC name, matching Python method naming."""
  43. out = []
  44. for i, ch in enumerate(m.name):
  45. if ch.isupper() and i:
  46. out.append("_")
  47. out.append(ch.lower())
  48. return "".join(out)
  49. def _pb2_module(fd: Any) -> str:
  50. """The import path of the pb2 module protoc emits for ``fd``."""
  51. name: str = fd.name
  52. return name[: -len(".proto")].replace("/", ".") + "_pb2"
  53. def _alias_for(module: str) -> str:
  54. """A valid, unique import alias for a pb2 module (dots -> underscores)."""
  55. return module.replace(".", "_")
  56. def _build_module_map(request: Any) -> dict[str, str]:
  57. """Map a fully-qualified message name (leading dot) to the pb2 module of
  58. the file that defines it.
  59. A service method's input/output may be a message declared in a *different*
  60. proto file than the service itself (e.g. ``MyAccountService`` in
  61. ``account.proto`` using ``ListExternalIdentitiesRequest`` from
  62. ``external_identities.proto``). The generated stub must import the module
  63. that actually defines the referenced message, not the service's own module.
  64. """
  65. mapping: dict[str, str] = {}
  66. for fd in request.proto_file:
  67. module = _pb2_module(fd)
  68. def walk(msgs: list[Any], prefix: list[str]) -> None:
  69. for mt in msgs:
  70. full = "." + ".".join([p for p in (fd.package, *prefix, mt.name) if p])
  71. mapping[full] = module
  72. walk(list(mt.nested_type), prefix + [mt.name])
  73. walk(list(fd.message_type), [])
  74. return mapping
  75. def _service_preamble(fd: Any, modules: list[str]) -> str:
  76. """The shared header + imports for a generated module (emitted once).
  77. ``modules`` is the (deduplicated) set of pb2 modules the file's services
  78. reference; each is imported under a unique alias so cross-file messages
  79. resolve to the module that defines them.
  80. """
  81. lines = [_HEADER, f"# source: {fd.name}", ""]
  82. for module in modules:
  83. lines.append(f"import {module} as {_alias_for(module)}")
  84. lines.append("from chattolib._connect import ConnectClient, ConnectClientSync, MethodInfo")
  85. lines += ["", ""]
  86. return "\n".join(lines)
  87. def _service_body(fd: Any, svc: Any, mod_by_fullname: dict[str, str]) -> str:
  88. """Render one service's async + sync client classes (no preamble)."""
  89. pkg = fd.package
  90. base = f"{pkg}.{svc.name}" if pkg else svc.name
  91. lines: list[str] = [f"class {svc.name}Client(ConnectClient):"]
  92. for m in svc.method:
  93. mn = _method_name(m)
  94. in_module = mod_by_fullname[m.input_type]
  95. out_module = mod_by_fullname[m.output_type]
  96. in_name = m.input_type.lstrip(".").split(".")[-1]
  97. out_name = m.output_type.lstrip(".").split(".")[-1]
  98. in_alias = _alias_for(in_module)
  99. out_alias = _alias_for(out_module)
  100. lines += [
  101. f" async def {mn}(self, request: {in_alias}.{in_name}, *,"
  102. f" headers: dict[str, str] | None = None) -> {out_alias}.{out_name}:",
  103. " return await self.execute_unary(",
  104. " request=request,",
  105. " method=MethodInfo(",
  106. f' name="{m.name}",',
  107. f' service_name="{base}",',
  108. f" input={in_alias}.{in_name},",
  109. f" output={out_alias}.{out_name},",
  110. " ),",
  111. " headers=headers,",
  112. " )",
  113. "",
  114. ]
  115. lines += [f"class {svc.name}ClientSync(ConnectClientSync):", ""]
  116. for m in svc.method:
  117. mn = _method_name(m)
  118. in_module = mod_by_fullname[m.input_type]
  119. out_module = mod_by_fullname[m.output_type]
  120. in_name = m.input_type.lstrip(".").split(".")[-1]
  121. out_name = m.output_type.lstrip(".").split(".")[-1]
  122. in_alias = _alias_for(in_module)
  123. out_alias = _alias_for(out_module)
  124. lines += [
  125. f" def {mn}(self, request: {in_alias}.{in_name}, *,"
  126. f" headers: dict[str, str] | None = None) -> {out_alias}.{out_name}:",
  127. " return self.execute_unary(",
  128. " request=request,",
  129. " method=MethodInfo(",
  130. f' name="{m.name}",',
  131. f' service_name="{base}",',
  132. f" input={in_alias}.{in_name},",
  133. f" output={out_alias}.{out_name},",
  134. " ),",
  135. " headers=headers,",
  136. " )",
  137. "",
  138. ]
  139. return "\n".join(lines)
  140. def generate(request: g.CodeGeneratorRequest) -> r.CodeGeneratorResponse:
  141. resp = r.CodeGeneratorResponse()
  142. # Tell protoc (>= the version that checks this) we can handle proto3
  143. # `optional` fields. Without this bit, recent protoc refuses to run the
  144. # plugin on any proto3 file that declares an optional field.
  145. resp.supported_features = r.CodeGeneratorResponse.Feature.FEATURE_PROTO3_OPTIONAL
  146. # file_to_generate is a repeated *string* (proto file names), not a list of
  147. # message objects — match by string membership.
  148. to_generate = set(request.file_to_generate)
  149. mod_by_fullname = _build_module_map(request)
  150. for fd in request.proto_file:
  151. if fd.name not in to_generate:
  152. continue
  153. services = [svc for svc in fd.service if svc.method]
  154. if not services:
  155. continue
  156. # One file per proto, containing every service it declares (e.g.
  157. # notifications.proto -> NotificationServiceClient and
  158. # NotificationPolicyServiceClient in the same module), matching the
  159. # historical connectrpc layout.
  160. required: list[str] = []
  161. seen: set[str] = set()
  162. for svc in services:
  163. for m in svc.method:
  164. for ref in (m.input_type, m.output_type):
  165. module = mod_by_fullname[ref]
  166. if module not in seen:
  167. seen.add(module)
  168. required.append(module)
  169. out = resp.file.add()
  170. out.name = fd.name.replace(".proto", "_connect.py")
  171. out.content = _service_preamble(fd, required) + "\n".join(
  172. _service_body(fd, svc, mod_by_fullname) for svc in services
  173. )
  174. return resp
  175. def main() -> None:
  176. data = sys.stdin.buffer.read()
  177. request = g.CodeGeneratorRequest.FromString(data)
  178. resp = generate(request)
  179. sys.stdout.buffer.write(resp.SerializeToString())
  180. if __name__ == "__main__":
  181. main()