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#[derive(Clone, Debug)]
10pub enum SingleTypeArg {
11 None,
13 Infer,
15 Explicit(Box<Type>),
17}
18
19impl SingleTypeArg {
20 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 pub fn is_infer(&self) -> bool {
30 matches!(self, Self::Infer)
31 }
32}
33
34pub 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
86pub 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
156pub 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
163pub 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}