1#!/usr/bin/env python3 2# ex: set filetype=python: 3 4"""Generate code for an RPC program's procedures""" 5 6from jinja2 import Environment 7 8from generators import SourceGenerator, create_jinja2_environment, get_jinja2_template 9from xdr_ast import _RpcProgram, _RpcVersion, excluded_apis 10from xdr_ast import max_widths, get_header_name 11 12 13def emit_version_definitions( 14 environment: Environment, program: str, version: _RpcVersion 15) -> None: 16 """Emit procedure numbers for each RPC version's procedures""" 17 template = environment.get_template("definition/open.j2") 18 print(template.render(program=program.upper())) 19 20 template = environment.get_template("definition/procedure.j2") 21 for procedure in version.procedures: 22 if procedure.name not in excluded_apis: 23 print( 24 template.render( 25 name=procedure.name, 26 value=procedure.number, 27 ) 28 ) 29 30 template = environment.get_template("definition/close.j2") 31 print(template.render()) 32 33 34def emit_version_declarations( 35 environment: Environment, program: str, version: _RpcVersion 36) -> None: 37 """Emit declarations for each RPC version's procedures""" 38 arguments = dict.fromkeys([]) 39 for procedure in version.procedures: 40 if procedure.name not in excluded_apis: 41 if procedure.argument.type_name == "void": 42 continue 43 arguments[procedure.argument.type_name] = None 44 if len(arguments) > 0: 45 print("") 46 template = environment.get_template("declaration/argument.j2") 47 for argument in arguments: 48 print(template.render(program=program, argument=argument)) 49 50 results = dict.fromkeys([]) 51 for procedure in version.procedures: 52 if procedure.name not in excluded_apis: 53 if procedure.result.type_name == "void": 54 continue 55 results[procedure.result.type_name] = None 56 if len(results) > 0: 57 print("") 58 template = environment.get_template("declaration/result.j2") 59 for result in results: 60 print(template.render(program=program, result=result)) 61 62 63def emit_version_argument_decoders( 64 environment: Environment, program: str, version: _RpcVersion 65) -> None: 66 """Emit server argument decoders for each RPC version's procedures""" 67 arguments = dict.fromkeys([]) 68 for procedure in version.procedures: 69 if procedure.name not in excluded_apis: 70 if procedure.argument.type_name == "void": 71 continue 72 arguments[procedure.argument.type_name] = None 73 74 template = environment.get_template("decoder/argument.j2") 75 for argument in arguments: 76 print(template.render(program=program, argument=argument)) 77 78 79def emit_version_result_decoders( 80 environment: Environment, program: str, version: _RpcVersion 81) -> None: 82 """Emit client result decoders for each RPC version's procedures""" 83 results = dict.fromkeys([]) 84 for procedure in version.procedures: 85 if procedure.name not in excluded_apis: 86 results[procedure.result.type_name] = None 87 88 template = environment.get_template("decoder/result.j2") 89 for result in results: 90 print(template.render(program=program, result=result)) 91 92 93def emit_version_argument_encoders( 94 environment: Environment, program: str, version: _RpcVersion 95) -> None: 96 """Emit client argument encoders for each RPC version's procedures""" 97 arguments = dict.fromkeys([]) 98 for procedure in version.procedures: 99 if procedure.name not in excluded_apis: 100 arguments[procedure.argument.type_name] = None 101 102 template = environment.get_template("encoder/argument.j2") 103 for argument in arguments: 104 print(template.render(program=program, argument=argument)) 105 106 107def emit_version_result_encoders( 108 environment: Environment, program: str, version: _RpcVersion 109) -> None: 110 """Emit server result encoders for each RPC version's procedures""" 111 results = dict.fromkeys([]) 112 for procedure in version.procedures: 113 if procedure.name not in excluded_apis: 114 if procedure.result.type_name == "void": 115 continue 116 results[procedure.result.type_name] = None 117 118 template = environment.get_template("encoder/result.j2") 119 for result in results: 120 print(template.render(program=program, result=result)) 121 122 123class XdrProgramGenerator(SourceGenerator): 124 """Generate source code for an RPC program's procedures""" 125 126 def __init__(self, language: str, peer: str): 127 """Initialize an instance of this class""" 128 self.environment = create_jinja2_environment(language, "program") 129 self.peer = peer 130 131 def emit_definition(self, node: _RpcProgram) -> None: 132 """Emit procedure numbers for each of an RPC programs's procedures""" 133 raw_name = node.name 134 program = raw_name.lower().removesuffix("_program").removesuffix("_prog") 135 136 for version in node.versions: 137 emit_version_definitions(self.environment, program, version) 138 139 template = self.environment.get_template("definition/program.j2") 140 print(template.render(name=raw_name, value=node.number)) 141 142 def emit_declaration(self, node: _RpcProgram) -> None: 143 """Emit a declaration pair for each of an RPC programs's procedures""" 144 raw_name = node.name 145 program = raw_name.lower().removesuffix("_program").removesuffix("_prog") 146 147 for version in node.versions: 148 emit_version_declarations(self.environment, program, version) 149 150 def emit_decoder(self, node: _RpcProgram) -> None: 151 """Emit all decoder functions for an RPC program's procedures""" 152 raw_name = node.name 153 program = raw_name.lower().removesuffix("_program").removesuffix("_prog") 154 match self.peer: 155 case "server": 156 for version in node.versions: 157 emit_version_argument_decoders( 158 self.environment, program, version, 159 ) 160 case "client": 161 for version in node.versions: 162 emit_version_result_decoders( 163 self.environment, program, version, 164 ) 165 166 def emit_encoder(self, node: _RpcProgram) -> None: 167 """Emit all encoder functions for an RPC program's procedures""" 168 raw_name = node.name 169 program = raw_name.lower().removesuffix("_program").removesuffix("_prog") 170 match self.peer: 171 case "server": 172 for version in node.versions: 173 emit_version_result_encoders( 174 self.environment, program, version, 175 ) 176 case "client": 177 for version in node.versions: 178 emit_version_argument_encoders( 179 self.environment, program, version, 180 ) 181 182 def emit_maxsize(self, node: _RpcProgram) -> None: 183 """Emit maxsize macro for maximum RPC argument size""" 184 header = get_header_name().upper() 185 186 # Find the largest argument across all versions 187 max_arg_width = 0 188 max_arg_name = None 189 for version in node.versions: 190 for procedure in version.procedures: 191 if procedure.name in excluded_apis: 192 continue 193 arg_name = procedure.argument.type_name 194 if arg_name == "void": 195 continue 196 if arg_name not in max_widths: 197 continue 198 if max_widths[arg_name] > max_arg_width: 199 max_arg_width = max_widths[arg_name] 200 max_arg_name = arg_name 201 202 if max_arg_name is None: 203 return 204 205 macro_name = header + "_MAX_ARGS_SZ" 206 template = get_jinja2_template(self.environment, "maxsize", "max_args") 207 print( 208 template.render( 209 macro=macro_name, 210 width=header + "_" + max_arg_name + "_sz", 211 ) 212 ) 213