xref: /linux/tools/net/sunrpc/xdrgen/generators/program.py (revision d141ec2825b4d3ec52f27c43bdd864090159273a)
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