Skip to main content

attribute_dsl/
infer.rs

1use syn::spanned::Spanned as _;
2use syn::visit_mut::{self, VisitMut as _};
3use syn::{
4    AngleBracketedGenericArguments, Error, Expr, GenericArgument, Path, PathArguments, Result,
5    Token, Type,
6};
7
8/// Single terminal type argument split from a path.
9#[derive(Clone, Debug)]
10pub enum SingleTypeArg {
11    /// The path's final segment had no generic arguments.
12    None,
13    /// The path's final segment used `_` as its only type argument.
14    Infer,
15    /// The path's final segment used one explicit type argument.
16    Explicit(Box<Type>),
17}
18
19impl SingleTypeArg {
20    /// Return the explicit type argument, if one was present.
21    pub fn explicit_type(&self) -> Option<&Type> {
22        match self {
23            Self::Explicit(ty) => Some(ty),
24            Self::None | Self::Infer => None,
25        }
26    }
27
28    /// Return whether the terminal type argument was `_`.
29    pub fn is_infer(&self) -> bool {
30        matches!(self, Self::Infer)
31    }
32}
33
34/// Split a path's final generic argument into a normalized single type arg.
35///
36/// This is useful for DSLs where `Thing::<_>` means "infer the field type" and
37/// `Thing::<T>` pins an explicit target type.
38///
39/// # Errors
40///
41/// Returns [`syn::Error`] when the path has no final segment, the final segment
42/// has more than one generic argument, the argument is not a type, or the final
43/// segment uses parenthesized generic arguments.
44pub fn split_terminal_single_type_arg(
45    mut path: Path,
46    subject: &str,
47) -> Result<(Path, SingleTypeArg)> {
48    let path_span = path.span();
49    let last_segment = path
50        .segments
51        .last_mut()
52        .ok_or_else(|| Error::new(path_span, format!("expected {subject} path")))?;
53
54    let args = std::mem::replace(&mut last_segment.arguments, PathArguments::None);
55    let type_arg = match args {
56        PathArguments::None => SingleTypeArg::None,
57        PathArguments::AngleBracketed(mut angle_args) => {
58            if angle_args.args.len() != 1 {
59                return Err(Error::new(
60                    angle_args.span(),
61                    format!("{subject} type syntax expects exactly one type argument"),
62                ));
63            }
64
65            let arg = angle_args.args.pop().expect("len checked");
66            match arg {
67                GenericArgument::Type(Type::Infer(_)) => SingleTypeArg::Infer,
68                GenericArgument::Type(ty) => SingleTypeArg::Explicit(Box::new(ty)),
69                _ => Err(Error::new(
70                    arg.span(),
71                    format!("{subject} type syntax expects a type argument"),
72                ))?,
73            }
74        },
75        PathArguments::Parenthesized(args) => {
76            return Err(Error::new(
77                args.span(),
78                format!("{subject} path does not support parenthesized arguments"),
79            ));
80        },
81    };
82
83    Ok((path, type_arg))
84}
85
86/// Substitute `replacement` for every `_` occurrence inside a type.
87pub fn substitute_infer_in_type(ty: &Type, replacement: &Type) -> Type {
88    match ty {
89        Type::Infer(_) => replacement.clone(),
90        Type::Path(type_path) => {
91            let mut type_path = type_path.clone();
92            type_path.path = substitute_infer_in_path(&type_path.path, replacement);
93            Type::Path(type_path)
94        },
95        Type::Array(array) => {
96            let mut array = array.clone();
97            array.elem = Box::new(substitute_infer_in_type(&array.elem, replacement));
98            Type::Array(array)
99        },
100        Type::Slice(slice) => {
101            let mut slice = slice.clone();
102            slice.elem = Box::new(substitute_infer_in_type(&slice.elem, replacement));
103            Type::Slice(slice)
104        },
105        Type::Ptr(ptr) => {
106            let mut ptr = ptr.clone();
107            ptr.elem = Box::new(substitute_infer_in_type(&ptr.elem, replacement));
108            Type::Ptr(ptr)
109        },
110        Type::FnPtr(fn_ptr) => {
111            let mut fn_ptr = fn_ptr.clone();
112            for input in &mut fn_ptr.inputs {
113                input.ty = substitute_infer_in_type(&input.ty, replacement);
114            }
115            substitute_infer_in_return_type(&mut fn_ptr.output, replacement);
116            Type::FnPtr(fn_ptr)
117        },
118        Type::TraitObject(trait_object) => {
119            let mut trait_object = trait_object.clone();
120            substitute_infer_in_bounds(&mut trait_object.bounds, replacement);
121            Type::TraitObject(trait_object)
122        },
123        Type::ImplTrait(impl_trait) => {
124            let mut impl_trait = impl_trait.clone();
125            substitute_infer_in_bounds(&mut impl_trait.bounds, replacement);
126            Type::ImplTrait(impl_trait)
127        },
128        Type::Tuple(tuple) => {
129            let mut tuple = tuple.clone();
130            tuple.elems = tuple
131                .elems
132                .iter()
133                .map(|ty| substitute_infer_in_type(ty, replacement))
134                .collect();
135            Type::Tuple(tuple)
136        },
137        Type::Paren(paren) => {
138            let mut paren = paren.clone();
139            paren.elem = Box::new(substitute_infer_in_type(&paren.elem, replacement));
140            Type::Paren(paren)
141        },
142        Type::Group(group) => {
143            let mut group = group.clone();
144            group.elem = Box::new(substitute_infer_in_type(&group.elem, replacement));
145            Type::Group(group)
146        },
147        Type::Reference(reference) => {
148            let mut reference = reference.clone();
149            *reference.elem = substitute_infer_in_type(&reference.elem, replacement);
150            Type::Reference(reference)
151        },
152        _ => ty.clone(),
153    }
154}
155
156/// Substitute `replacement` for every `_` occurrence inside an expression.
157pub fn substitute_infer_in_expr(expr: &Expr, replacement: &Type) -> Expr {
158    let mut expr = expr.clone();
159    InferSubstitutor { replacement }.visit_expr_mut(&mut expr);
160    expr
161}
162
163/// Substitute `replacement` for every `_` occurrence inside path arguments.
164pub fn substitute_infer_in_path(path: &Path, replacement: &Type) -> Path {
165    let mut path = path.clone();
166
167    for segment in &mut path.segments {
168        substitute_infer_in_path_arguments(&mut segment.arguments, replacement);
169    }
170
171    path
172}
173
174struct InferSubstitutor<'a> {
175    replacement: &'a Type,
176}
177
178impl visit_mut::VisitMut for InferSubstitutor<'_> {
179    fn visit_type_mut(&mut self, node: &mut Type) {
180        *node = substitute_infer_in_type(node, self.replacement);
181    }
182
183    fn visit_path_mut(&mut self, node: &mut Path) {
184        *node = substitute_infer_in_path(node, self.replacement);
185    }
186}
187
188fn substitute_infer_in_return_type(return_type: &mut syn::ReturnType, replacement: &Type) {
189    if let syn::ReturnType::Type(_, ty) = return_type {
190        **ty = substitute_infer_in_type(ty, replacement);
191    }
192}
193
194fn substitute_infer_in_bounds(
195    bounds: &mut syn::punctuated::Punctuated<syn::TypeParamBound, Token![+]>,
196    replacement: &Type,
197) {
198    for bound in bounds {
199        if let syn::TypeParamBound::Trait(trait_bound) = bound {
200            trait_bound.path = substitute_infer_in_path(&trait_bound.path, replacement);
201        }
202    }
203}
204
205fn substitute_infer_in_path_arguments(arguments: &mut PathArguments, replacement: &Type) {
206    match arguments {
207        PathArguments::AngleBracketed(args) => {
208            substitute_infer_in_angle_bracketed_arguments(args, replacement);
209        },
210        PathArguments::Parenthesized(args) => {
211            for input in &mut args.inputs {
212                input.ty = substitute_infer_in_type(&input.ty, replacement);
213            }
214            substitute_infer_in_return_type(&mut args.output, replacement);
215        },
216        PathArguments::None => {},
217    }
218}
219
220fn substitute_infer_in_angle_bracketed_arguments(
221    args: &mut AngleBracketedGenericArguments,
222    replacement: &Type,
223) {
224    for arg in &mut args.args {
225        match arg {
226            GenericArgument::Type(ty) => {
227                *ty = substitute_infer_in_type(ty, replacement);
228            },
229            GenericArgument::AssocType(assoc_type) => {
230                if let Some(generics) = &mut assoc_type.generics {
231                    substitute_infer_in_angle_bracketed_arguments(generics, replacement);
232                }
233                assoc_type.ty = substitute_infer_in_type(&assoc_type.ty, replacement);
234            },
235            GenericArgument::Constraint(constraint) => {
236                if let Some(generics) = &mut constraint.generics {
237                    substitute_infer_in_angle_bracketed_arguments(generics, replacement);
238                }
239                substitute_infer_in_bounds(&mut constraint.bounds, replacement);
240            },
241            _ => {},
242        }
243    }
244}
245
246#[cfg(test)]
247mod tests {
248    use super::*;
249    use syn::{Type, parse_quote};
250
251    fn compact(tokens: impl quote::ToTokens) -> String {
252        tokens
253            .to_token_stream()
254            .to_string()
255            .chars()
256            .filter(|ch| !ch.is_whitespace())
257            .collect()
258    }
259
260    fn parenthesized_path(output: Type) -> Path {
261        let mut inputs = syn::punctuated::Punctuated::new();
262        inputs.push(syn::NamedArg {
263            attrs: Vec::new(),
264            name: None,
265            ty: parse_quote!(_),
266        });
267
268        Path::from(syn::PathSegment {
269            ident: parse_quote!(FnOnce),
270            arguments: PathArguments::Parenthesized(syn::ParenthesizedGenericArguments {
271                paren_token: Default::default(),
272                inputs,
273                output: syn::ReturnType::Type(Default::default(), Box::new(output)),
274            }),
275        })
276    }
277
278    #[test]
279    fn splits_terminal_single_type_arg() {
280        let path: Path = parse_quote!(crate::RangeValidation::<_>);
281        let (path, arg) = split_terminal_single_type_arg(path, "validator").expect("valid path");
282        assert_eq!(compact(&path), "crate::RangeValidation");
283        assert!(arg.is_infer());
284
285        let path: Path = parse_quote!(crate::RangeValidation::<i32>);
286        let (_, arg) = split_terminal_single_type_arg(path, "validator").expect("valid path");
287        assert_eq!(compact(arg.explicit_type().expect("explicit type")), "i32");
288    }
289
290    #[test]
291    fn splits_absent_terminal_type_arg_and_rejects_invalid_args() {
292        let path: Path = parse_quote!(crate::RangeValidation);
293        let (path, arg) = split_terminal_single_type_arg(path, "validator").expect("valid path");
294        assert_eq!(compact(&path), "crate::RangeValidation");
295        assert!(!arg.is_infer());
296        assert!(arg.explicit_type().is_none());
297
298        let path: Path = parse_quote!(crate::RangeValidation::<i32, String>);
299        let err = split_terminal_single_type_arg(path, "validator").expect_err("too many args");
300        assert!(
301            err.to_string()
302                .contains("validator type syntax expects exactly one type argument"),
303            "{err}"
304        );
305
306        let path: Path = parse_quote!(crate::RangeValidation::<3>);
307        let err = split_terminal_single_type_arg(path, "validator").expect_err("const arg");
308        assert!(
309            err.to_string()
310                .contains("validator type syntax expects a type argument"),
311            "{err}"
312        );
313
314        let path = parenthesized_path(parse_quote!(i32));
315        let err = split_terminal_single_type_arg(path, "validator").expect_err("function args");
316        assert!(
317            err.to_string()
318                .contains("validator path does not support parenthesized arguments"),
319            "{err}"
320        );
321    }
322
323    #[test]
324    fn substitutes_infer_in_paths_types_and_exprs() {
325        let replacement: Type = parse_quote!(String);
326        let path: Path = parse_quote!(crate::Input<Option<_>>);
327        assert_eq!(
328            compact(substitute_infer_in_path(&path, &replacement)),
329            "crate::Input<Option<String>>"
330        );
331
332        let ty: Type = parse_quote!(fn([_; 2], &[_]) -> Option<_>);
333        assert_eq!(
334            compact(substitute_infer_in_type(&ty, &replacement)),
335            "fn([String;2],&[String])->Option<String>"
336        );
337
338        let expr: Expr = parse_quote!(crate::Select::<_>.searchable(true));
339        assert_eq!(
340            compact(substitute_infer_in_expr(&expr, &replacement)),
341            "crate::Select::<String>.searchable(true)"
342        );
343    }
344
345    #[test]
346    fn substitutes_infer_in_additional_type_forms() {
347        let replacement: Type = parse_quote!(String);
348
349        let ptr: Type = parse_quote!(*const _);
350        assert_eq!(
351            compact(substitute_infer_in_type(&ptr, &replacement)),
352            "*constString"
353        );
354
355        let trait_object: Type = parse_quote!(dyn Iterator<Item = _> + Send);
356        assert_eq!(
357            compact(substitute_infer_in_type(&trait_object, &replacement)),
358            "dynIterator<Item=String>+Send"
359        );
360
361        let impl_trait: Type = parse_quote!(impl Into<_> + Send);
362        assert_eq!(
363            compact(substitute_infer_in_type(&impl_trait, &replacement)),
364            "implInto<String>+Send"
365        );
366
367        let tuple: Type = parse_quote!((_, Option<_>));
368        assert_eq!(
369            compact(substitute_infer_in_type(&tuple, &replacement)),
370            "(String,Option<String>)"
371        );
372
373        let paren: Type = parse_quote!((Option<_>));
374        assert_eq!(
375            compact(substitute_infer_in_type(&paren, &replacement)),
376            "(Option<String>)"
377        );
378
379        let group = Type::Group(syn::TypeGroup {
380            attrs: Vec::new(),
381            group_token: Default::default(),
382            elem: Box::new(parse_quote!(Option<_>)),
383        });
384        assert_eq!(
385            compact(substitute_infer_in_type(&group, &replacement)),
386            "Option<String>"
387        );
388
389        let never: Type = parse_quote!(!);
390        assert_eq!(compact(substitute_infer_in_type(&never, &replacement)), "!");
391    }
392
393    #[test]
394    fn substitutes_infer_in_path_argument_variants() {
395        let replacement: Type = parse_quote!(String);
396
397        let parenthesized = parenthesized_path(parse_quote!(_));
398        assert_eq!(
399            compact(substitute_infer_in_path(&parenthesized, &replacement)),
400            "FnOnce(String)->String"
401        );
402
403        let assoc_type: Path = parse_quote!(Trait<Assoc<_> = Result<_, _>>);
404        assert_eq!(
405            compact(substitute_infer_in_path(&assoc_type, &replacement)),
406            "Trait<Assoc<String>=Result<String,String>>"
407        );
408
409        let constraint: Path = parse_quote!(Trait<Assoc<_>: Into<_> + From<_>>);
410        assert_eq!(
411            compact(substitute_infer_in_path(&constraint, &replacement)),
412            "Trait<Assoc<String>:Into<String>+From<String>>"
413        );
414
415        let lifetime_and_const: Path = parse_quote!(Trait<'static, 3, _>);
416        assert_eq!(
417            compact(substitute_infer_in_path(&lifetime_and_const, &replacement)),
418            "Trait<'static,3,String>"
419        );
420    }
421
422    #[test]
423    fn substitutes_infer_inside_expression_types() {
424        let replacement: Type = parse_quote!(String);
425        let expr: Expr = parse_quote!(value as *const _);
426
427        assert_eq!(
428            compact(substitute_infer_in_expr(&expr, &replacement)),
429            "valueas*constString"
430        );
431    }
432}