_protoc_plugin.py 8.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236
  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 _message_owner(request: Any) -> dict[str, str]:
  50. """Map each fully-qualified message name to the proto file that defines it.
  51. A service's RPC can reference a request/response message that is defined in
  52. a *different* ``.proto`` file than the service itself (e.g. ``MyAccountService
  53. .UpdatePresence`` in ``account.proto`` uses ``UpdatePresenceRequest`` from
  54. ``presence.proto``). To emit a correct import for each, we need to know which
  55. proto owns each message. We scan every ``FileDescriptorProto`` in the request
  56. (including nested messages) to build that map.
  57. """
  58. owner: dict[str, str] = {}
  59. def _walk(file_name: str, pkg: str, message_type: Any) -> None:
  60. for mt in message_type:
  61. fqn = f"{pkg}.{mt.name}" if pkg else mt.name
  62. owner.setdefault(fqn, file_name)
  63. _walk(file_name, f"{pkg}.{mt.name}" if pkg else mt.name, mt.nested_type)
  64. for fd in request.proto_file:
  65. _walk(fd.name, fd.package, fd.message_type)
  66. return owner
  67. def _module_alias(proto_name: str, current_proto: str) -> str:
  68. """A stable import alias for a proto's ``_pb2`` module.
  69. The proto being generated is always ``_pb2`` (so same-file references are
  70. unchanged); any other referenced proto gets ``_pb2_<last-path-segment>``.
  71. """
  72. if proto_name == current_proto:
  73. return "_pb2"
  74. stem = proto_name[: -len(".proto")].rsplit("/", 1)[-1]
  75. return f"_pb2_{stem}"
  76. def _pb2_module(proto_name: str) -> str:
  77. """The dotted import path for a proto's generated ``_pb2`` module."""
  78. return proto_name[: -len(".proto")].replace("/", ".") + "_pb2"
  79. def _service_preamble(fd: Any, extra_modules: list[str]) -> str:
  80. """The shared header + imports for a generated module (emitted once).
  81. ``extra_modules`` are proto file names (other than ``fd.name``) whose
  82. messages this file's services reference; each gets its own import so a
  83. cross-file RPC points at the module that actually defines its messages.
  84. """
  85. lines = [
  86. _HEADER,
  87. f"# source: {fd.name}",
  88. "",
  89. f"import {_pb2_module(fd.name)} as _pb2",
  90. ]
  91. for mod in extra_modules:
  92. lines.append(f"import {_pb2_module(mod)} as {_module_alias(mod, fd.name)}")
  93. lines += [
  94. "from chattolib._connect import ConnectClient, ConnectClientSync, MethodInfo",
  95. "",
  96. "",
  97. ]
  98. return "\n".join(lines)
  99. def _service_body(fd: Any, svc: Any, owner: dict[str, str]) -> str:
  100. """Render one service's async + sync client classes (no preamble)."""
  101. pkg = fd.package
  102. base = f"{pkg}.{svc.name}" if pkg else svc.name
  103. def _alias(fqn: str) -> str:
  104. """The import alias to use for a fully-qualified message name."""
  105. proto = owner.get(fqn, fd.name)
  106. return _module_alias(proto, fd.name)
  107. lines: list[str] = [f"class {svc.name}Client(ConnectClient):"]
  108. for m in svc.method:
  109. mn = _method_name(m)
  110. in_fqn = m.input_type.lstrip(".")
  111. out_fqn = m.output_type.lstrip(".")
  112. in_name = in_fqn.split(".")[-1]
  113. out_name = out_fqn.split(".")[-1]
  114. in_alias = _alias(in_fqn)
  115. out_alias = _alias(out_fqn)
  116. lines += [
  117. f" async def {mn}(self, request: {in_alias}.{in_name}, *,"
  118. f" headers: dict[str, str] | None = None) -> {out_alias}.{out_name}:",
  119. " return await self.execute_unary(",
  120. " request=request,",
  121. " method=MethodInfo(",
  122. f' name="{m.name}",',
  123. f' service_name="{base}",',
  124. f" input={in_alias}.{in_name},",
  125. f" output={out_alias}.{out_name},",
  126. " ),",
  127. " headers=headers,",
  128. " )",
  129. "",
  130. ]
  131. lines += [f"class {svc.name}ClientSync(ConnectClientSync):", ""]
  132. for m in svc.method:
  133. mn = _method_name(m)
  134. in_fqn = m.input_type.lstrip(".")
  135. out_fqn = m.output_type.lstrip(".")
  136. in_name = in_fqn.split(".")[-1]
  137. out_name = out_fqn.split(".")[-1]
  138. in_alias = _alias(in_fqn)
  139. out_alias = _alias(out_fqn)
  140. lines += [
  141. f" def {mn}(self, request: {in_alias}.{in_name}, *,"
  142. f" headers: dict[str, str] | None = None) -> {out_alias}.{out_name}:",
  143. " return self.execute_unary(",
  144. " request=request,",
  145. " method=MethodInfo(",
  146. f' name="{m.name}",',
  147. f' service_name="{base}",',
  148. f" input={in_alias}.{in_name},",
  149. f" output={out_alias}.{out_name},",
  150. " ),",
  151. " headers=headers,",
  152. " )",
  153. "",
  154. ]
  155. return "\n".join(lines)
  156. def generate(request: g.CodeGeneratorRequest) -> r.CodeGeneratorResponse:
  157. resp = r.CodeGeneratorResponse()
  158. # Tell protoc (>= the version that checks this) we can handle proto3
  159. # `optional` fields. Without this bit, recent protoc refuses to run the
  160. # plugin on any proto3 file that declares an optional field.
  161. resp.supported_features = r.CodeGeneratorResponse.Feature.FEATURE_PROTO3_OPTIONAL
  162. # file_to_generate is a repeated *string* (proto file names), not a list of
  163. # message objects — match by string membership.
  164. to_generate = set(request.file_to_generate)
  165. owner = _message_owner(request)
  166. for fd in request.proto_file:
  167. if fd.name not in to_generate:
  168. continue
  169. services = [svc for svc in fd.service if svc.method]
  170. if not services:
  171. continue
  172. # Collect the other proto files this file's services reference (for
  173. # cross-file RPC messages), in first-seen order, deduplicated.
  174. extra: list[str] = []
  175. for svc in services:
  176. for m in svc.method:
  177. for fqn in (m.input_type.lstrip("."), m.output_type.lstrip(".")):
  178. proto = owner.get(fqn, fd.name)
  179. if proto != fd.name and proto not in extra:
  180. extra.append(proto)
  181. # One file per proto, containing every service it declares (e.g.
  182. # notifications.proto -> NotificationServiceClient and
  183. # NotificationPolicyServiceClient in the same module), matching the
  184. # historical connectrpc layout.
  185. out = resp.file.add()
  186. out.name = fd.name.replace(".proto", "_connect.py")
  187. out.content = _service_preamble(fd, extra) + "\n".join(
  188. _service_body(fd, svc, owner) for svc in services
  189. )
  190. return resp
  191. def main() -> None:
  192. data = sys.stdin.buffer.read()
  193. request = g.CodeGeneratorRequest.FromString(data)
  194. resp = generate(request)
  195. sys.stdout.buffer.write(resp.SerializeToString())
  196. if __name__ == "__main__":
  197. main()