Skip to main content

capnpc/
codegen_types.rs

1// Copyright (c) 2013-2015 Sandstorm Development Group, Inc. and contributors
2// Licensed under the MIT License:
3//
4// Permission is hereby granted, free of charge, to any person obtaining a copy
5// of this software and associated documentation files (the "Software"), to deal
6// in the Software without restriction, including without limitation the rights
7// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
8// copies of the Software, and to permit persons to whom the Software is
9// furnished to do so, subject to the following conditions:
10//
11// The above copyright notice and this permission notice shall be included in
12// all copies or substantial portions of the Software.
13//
14// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
15// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
16// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
17// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
18// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
19// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
20// THE SOFTWARE.
21
22use crate::codegen::{fmt, GeneratorContext};
23use capnp::schema_capnp::{brand, node, type_};
24use capnp::Error;
25use std::collections::hash_map::HashMap;
26
27#[derive(Copy, Clone, PartialEq)]
28pub enum Leaf {
29    Reader(&'static str),
30    Builder(&'static str),
31    Owned,
32    Client,
33    Server,
34    ServerDispatch,
35    Pipeline,
36    GetType,
37}
38
39impl ::std::fmt::Display for Leaf {
40    fn fmt(&self, fmt: &mut ::std::fmt::Formatter) -> Result<(), ::std::fmt::Error> {
41        let display_string = match *self {
42            Self::Reader(lt) => format!("Reader<{lt}>"),
43            Self::Builder(lt) => format!("Builder<{lt}>"),
44            Self::Owned => "Owned".to_string(),
45            Self::Client => "Client".to_string(),
46            Self::Server => "Server".to_string(),
47            Self::ServerDispatch => "ServerDispatch".to_string(),
48            Self::Pipeline => "Pipeline".to_string(),
49            Self::GetType => "get_type".to_string(),
50        };
51        ::std::fmt::Display::fmt(&display_string, fmt)
52    }
53}
54
55impl Leaf {
56    fn bare_name(&self) -> &'static str {
57        match *self {
58            Self::Reader(_) => "Reader",
59            Self::Builder(_) => "Builder",
60            Self::Owned => "Owned",
61            Self::Client => "Client",
62            Self::Server => "Server",
63            Self::ServerDispatch => "ServerDispatch",
64            Self::GetType => "get_type",
65            Self::Pipeline => "Pipeline",
66        }
67    }
68
69    fn _have_lifetime(&self) -> bool {
70        match self {
71            &Self::Reader(_) | &Self::Builder(_) => true,
72            &Self::Owned
73            | &Self::Client
74            | &Self::Server
75            | &Self::ServerDispatch
76            | &Self::GetType
77            | &Self::Pipeline => false,
78        }
79    }
80}
81
82pub struct TypeParameterTexts {
83    pub expanded_list: Vec<String>,
84    pub params: String,
85    pub where_clause: String,
86    pub pipeline_where_clause: String,
87    pub phantom_data_value: String,
88    pub phantom_data_type: String,
89}
90
91// this is a collection of helpers acting on a "Node" (most of them are Type definitions)
92pub trait RustNodeInfo {
93    fn parameters_texts(&self, ctx: &GeneratorContext) -> TypeParameterTexts;
94}
95
96// this is a collection of helpers acting on a "Type" (someplace where a Type is used, not defined)
97pub trait RustTypeInfo {
98    fn is_prim(&self) -> Result<bool, Error>;
99    fn is_pointer(&self) -> Result<bool, Error>;
100    fn is_parameter(&self) -> Result<bool, Error>;
101    fn is_branded(&self) -> Result<bool, Error>;
102    fn type_string(&self, ctx: &GeneratorContext, module: Leaf) -> Result<String, Error>;
103}
104
105impl RustNodeInfo for node::Reader<'_> {
106    fn parameters_texts(&self, ctx: &GeneratorContext) -> TypeParameterTexts {
107        if self.get_is_generic() {
108            let params = get_type_parameters(ctx, self.get_id());
109            let type_parameters = params
110                .iter()
111                .map(|param| param.to_string())
112                .collect::<Vec<String>>()
113                .join(",");
114            let where_clause = "where ".to_string()
115                + &*(params
116                    .iter()
117                    .map(|param| fmt!(ctx, "{param}: {capnp}::traits::Owned + 'static"))
118                    .collect::<Vec<String>>()
119                    .join(", ")
120                    + " ");
121            let pipeline_where_clause = "where ".to_string() + &*(params.iter().map(|param| {
122                fmt!(ctx, "{param}: {capnp}::traits::Pipelined, <{param} as {capnp}::traits::Pipelined>::Pipeline: {capnp}::capability::FromTypelessPipeline")
123            }).collect::<Vec<String>>().join(", ") + " ");
124            let phantom_data_type = if params.len() == 1 {
125                // omit parens to avoid linter error
126                format!("_phantom: ::core::marker::PhantomData<{type_parameters}>")
127            } else {
128                format!("_phantom: ::core::marker::PhantomData<({type_parameters})>")
129            };
130            let phantom_data_value = "_phantom: ::core::marker::PhantomData,".to_string();
131
132            TypeParameterTexts {
133                expanded_list: params,
134                params: type_parameters,
135                where_clause,
136                pipeline_where_clause,
137                phantom_data_type,
138                phantom_data_value,
139            }
140        } else {
141            TypeParameterTexts {
142                expanded_list: vec![],
143                params: "".to_string(),
144                where_clause: "".to_string(),
145                pipeline_where_clause: "".to_string(),
146                phantom_data_type: "".to_string(),
147                phantom_data_value: "".to_string(),
148            }
149        }
150    }
151}
152
153impl RustTypeInfo for type_::Reader<'_> {
154    fn type_string(&self, ctx: &GeneratorContext, module: Leaf) -> Result<String, Error> {
155        // A scalar `type` newtype: render the alias's name in place of the underlying type. The
156        // alias module (emitted at the type node) provides Reader/Builder/Owned; a value newtype
157        // has no lifetime, so only its pointer counterpart keeps `<'a>`.
158        if self.get_type_id() != 0 {
159            let m = ctx.get_qualified_module(self.get_type_id());
160            let bare = module.bare_name();
161            return Ok(match module {
162                Leaf::Reader(lt) | Leaf::Builder(lt) if self.is_pointer()? => {
163                    format!("{m}::{bare}<{lt}>")
164                }
165                _ => format!("{m}::{bare}"),
166            });
167        }
168
169        let local_lifetime = match module {
170            Leaf::Reader(lt) => lt,
171            Leaf::Builder(lt) => lt,
172            _ => "",
173        };
174
175        let lifetime_comma = if local_lifetime.is_empty() {
176            "".to_string()
177        } else {
178            format!("{local_lifetime},")
179        };
180
181        match self.which()? {
182            type_::Void(()) => Ok("()".to_string()),
183            type_::Bool(()) => Ok("bool".to_string()),
184            type_::Int8(()) => Ok("i8".to_string()),
185            type_::Int16(()) => Ok("i16".to_string()),
186            type_::Int32(()) => Ok("i32".to_string()),
187            type_::Int64(()) => Ok("i64".to_string()),
188            type_::Uint8(()) => Ok("u8".to_string()),
189            type_::Uint16(()) => Ok("u16".to_string()),
190            type_::Uint32(()) => Ok("u32".to_string()),
191            type_::Uint64(()) => Ok("u64".to_string()),
192            type_::Float32(()) => Ok("f32".to_string()),
193            type_::Float64(()) => Ok("f64".to_string()),
194            type_::Text(()) => Ok(fmt!(ctx, "{capnp}::text::{module}")),
195            type_::Data(()) => Ok(fmt!(ctx, "{capnp}::data::{module}")),
196            type_::Struct(st) => do_branding(
197                ctx,
198                st.get_type_id(),
199                st.get_brand()?,
200                module,
201                &ctx.get_qualified_module(st.get_type_id()),
202            ),
203            type_::Interface(interface) => do_branding(
204                ctx,
205                interface.get_type_id(),
206                interface.get_brand()?,
207                module,
208                &ctx.get_qualified_module(interface.get_type_id()),
209            ),
210            type_::List(ot1) => {
211                let element_type = ot1.get_element_type()?;
212                match element_type.which()? {
213                    type_::Struct(_) => {
214                        let inner = element_type.type_string(ctx, Leaf::Owned)?;
215                        Ok(fmt!(
216                            ctx,
217                            "{capnp}::struct_list::{}<{lifetime_comma}{inner}>",
218                            module.bare_name()
219                        ))
220                    }
221                    type_::Enum(_) => {
222                        let inner = element_type.type_string(ctx, Leaf::Owned)?;
223                        Ok(fmt!(
224                            ctx,
225                            "{capnp}::enum_list::{}<{lifetime_comma}{inner}>",
226                            module.bare_name()
227                        ))
228                    }
229                    type_::List(_) => {
230                        let inner = element_type.type_string(ctx, Leaf::Owned)?;
231                        Ok(fmt!(
232                            ctx,
233                            "{capnp}::list_list::{}<{lifetime_comma}{inner}>",
234                            module.bare_name()
235                        ))
236                    }
237                    type_::Text(()) => Ok(fmt!(ctx, "{capnp}::text_list::{module}")),
238                    type_::Data(()) => Ok(fmt!(ctx, "{capnp}::data_list::{module}")),
239                    type_::Interface(_) => {
240                        let inner = element_type.type_string(ctx, Leaf::Client)?;
241                        Ok(fmt!(
242                            ctx,
243                            "{capnp}::capability_list::{}<{lifetime_comma}{inner}>",
244                            module.bare_name()
245                        ))
246                    }
247                    type_::AnyPointer(_) => {
248                        Err(Error::failed("List(AnyPointer) is unsupported".to_string()))
249                    }
250                    _ => {
251                        let inner = element_type.type_string(ctx, Leaf::Owned)?;
252                        Ok(fmt!(
253                            ctx,
254                            "{capnp}::primitive_list::{}<{lifetime_comma}{inner}>",
255                            module.bare_name()
256                        ))
257                    }
258                }
259            }
260            type_::Enum(en) => Ok(ctx.get_qualified_module(en.get_type_id())),
261            type_::AnyPointer(pointer) => match pointer.which()? {
262                type_::any_pointer::Parameter(def) => {
263                    let the_struct = &ctx.node_map[&def.get_scope_id()];
264                    let parameters = the_struct.get_parameters()?;
265                    let parameter = parameters.get(u32::from(def.get_parameter_index()));
266                    let parameter_name = parameter.get_name()?.to_str()?;
267                    match module {
268                        Leaf::Owned => Ok(parameter_name.to_string()),
269                        Leaf::Reader(lifetime) => Ok(fmt!(
270                            ctx,
271                            "<{parameter_name} as {capnp}::traits::Owned>::Reader<{lifetime}>"
272                        )),
273                        Leaf::Builder(lifetime) => Ok(fmt!(
274                            ctx,
275                            "<{parameter_name} as {capnp}::traits::Owned>::Builder<{lifetime}>"
276                        )),
277                        Leaf::Pipeline => Ok(fmt!(
278                            ctx,
279                            "<{parameter_name} as {capnp}::traits::Pipelined>::Pipeline"
280                        )),
281                        _ => Err(Error::unimplemented(
282                            "unimplemented any_pointer leaf".to_string(),
283                        )),
284                    }
285                }
286                _ => match module {
287                    Leaf::Reader(lifetime) => {
288                        Ok(fmt!(ctx, "{capnp}::any_pointer::Reader<{lifetime}>"))
289                    }
290                    Leaf::Builder(lifetime) => {
291                        Ok(fmt!(ctx, "{capnp}::any_pointer::Builder<{lifetime}>"))
292                    }
293                    _ => Ok(fmt!(ctx, "{capnp}::any_pointer::{module}")),
294                },
295            },
296        }
297    }
298
299    fn is_parameter(&self) -> Result<bool, Error> {
300        match self.which()? {
301            type_::AnyPointer(pointer) => match pointer.which()? {
302                type_::any_pointer::Parameter(_) => Ok(true),
303                _ => Ok(false),
304            },
305            _ => Ok(false),
306        }
307    }
308
309    fn is_branded(&self) -> Result<bool, Error> {
310        match self.which()? {
311            type_::Struct(st) => {
312                let brand = st.get_brand()?;
313                let scopes = brand.get_scopes()?;
314                Ok(!scopes.is_empty())
315            }
316            _ => Ok(false),
317        }
318    }
319
320    #[inline(always)]
321    fn is_prim(&self) -> Result<bool, Error> {
322        match self.which()? {
323            type_::Int8(())
324            | type_::Int16(())
325            | type_::Int32(())
326            | type_::Int64(())
327            | type_::Uint8(())
328            | type_::Uint16(())
329            | type_::Uint32(())
330            | type_::Uint64(())
331            | type_::Float32(())
332            | type_::Float64(())
333            | type_::Void(())
334            | type_::Bool(()) => Ok(true),
335            _ => Ok(false),
336        }
337    }
338
339    #[inline(always)]
340    fn is_pointer(&self) -> Result<bool, Error> {
341        Ok(matches!(
342            self.which()?,
343            type_::Text(())
344                | type_::Data(())
345                | type_::List(_)
346                | type_::Struct(_)
347                | type_::Interface(_)
348                | type_::AnyPointer(_)
349        ))
350    }
351}
352
353pub fn do_branding(
354    ctx: &GeneratorContext,
355    node_id: u64,
356    brand: brand::Reader,
357    leaf: Leaf,
358    the_mod: &str,
359) -> Result<String, Error> {
360    let scopes = brand.get_scopes()?;
361    let mut brand_scopes = HashMap::new();
362    for scope in scopes {
363        brand_scopes.insert(scope.get_scope_id(), scope);
364    }
365    let brand_scopes = brand_scopes; // freeze
366    let mut current_node_id = node_id;
367    let mut accumulator: Vec<Vec<String>> = Vec::new();
368    loop {
369        let current_node = ctx.node_map[&current_node_id];
370        let params = current_node.get_parameters()?;
371        let mut arguments: Vec<String> = Vec::new();
372        match brand_scopes.get(&current_node_id) {
373            None => {
374                for _ in params {
375                    arguments.push(fmt!(ctx, "{capnp}::any_pointer::Owned"));
376                }
377            }
378            Some(scope) => match scope.which()? {
379                brand::scope::Inherit(()) => {
380                    for param in params {
381                        arguments.push(param.get_name()?.to_string()?);
382                    }
383                }
384                brand::scope::Bind(bindings_list_opt) => {
385                    let bindings_list = bindings_list_opt?;
386                    assert_eq!(bindings_list.len(), params.len());
387                    for binding in bindings_list {
388                        match binding.which()? {
389                            brand::binding::Unbound(()) => {
390                                arguments.push(fmt!(ctx, "{capnp}::any_pointer::Owned"));
391                            }
392                            brand::binding::Type(t) => {
393                                arguments.push(t?.type_string(ctx, Leaf::Owned)?);
394                            }
395                        }
396                    }
397                }
398            },
399        }
400        accumulator.push(arguments);
401        current_node_id = match ctx.node_parents.get(&current_node_id).copied() {
402            Some(0) | None => break,
403            Some(id) => id,
404        };
405    }
406
407    // Now add a lifetime parameter if the leaf has one.
408    match leaf {
409        Leaf::Reader(lt) => accumulator.push(vec![lt.to_string()]),
410        Leaf::Builder(lt) => accumulator.push(vec![lt.to_string()]),
411        Leaf::ServerDispatch => accumulator.push(vec!["_T".to_string()]), // HACK
412        _ => (),
413    }
414
415    accumulator.reverse();
416    let accumulated = accumulator.concat();
417
418    let arguments = if !accumulated.is_empty() {
419        format!("<{}>", accumulated.join(","))
420    } else {
421        "".to_string()
422    };
423
424    let maybe_colons = if leaf == Leaf::ServerDispatch || leaf == Leaf::GetType {
425        "::"
426    } else {
427        ""
428    }; // HACK
429    Ok(format!(
430        "{the_mod}::{leaf}{maybe_colons}{arguments}",
431        leaf = leaf.bare_name()
432    ))
433}
434
435pub fn get_type_parameters(ctx: &GeneratorContext, node_id: u64) -> Vec<String> {
436    let mut current_node_id = node_id;
437    let mut accumulator: Vec<Vec<String>> = Vec::new();
438    loop {
439        let current_node = ctx.node_map[&current_node_id];
440        let mut params = Vec::new();
441        for param in current_node.get_parameters().unwrap() {
442            params.push(param.get_name().unwrap().to_string().unwrap());
443        }
444
445        accumulator.push(params);
446        current_node_id = match ctx.node_parents.get(&current_node_id).copied() {
447            Some(0) | None => break,
448            Some(id) => id,
449        };
450    }
451
452    accumulator.reverse();
453    accumulator.concat()
454}