Skip to main content

hir_expand/builtin/
derive_macro.rs

1//! Builtin derives.
2
3use base_db::SourceDatabase;
4use either::Either;
5use intern::sym;
6use itertools::{Itertools, izip};
7use parser::SyntaxKind;
8use rustc_hash::FxHashSet;
9use span::{Edition, Span};
10use stdx::never;
11use syntax_bridge::DocCommentDesugarMode;
12use tracing::debug;
13
14use crate::{
15    ExpandError, ExpandResult, MacroCallId,
16    builtin::quote::dollar_crate,
17    hygiene::span_with_def_site_ctxt,
18    name::{self, AsName, Name},
19    span_map::ExpansionSpanMap,
20    tt,
21};
22use syntax::{
23    ast::{
24        self, AstNode, FieldList, HasAttrs, HasGenericArgs, HasGenericParams, HasModuleItem,
25        HasName, HasTypeBounds,
26    },
27    syntax_editor::{GetOrCreateWhereClause, SyntaxEditor},
28};
29
30macro_rules! register_builtin {
31    ( $($trait:ident => $expand:ident),* $(,)? ) => {
32        #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
33        pub enum BuiltinDeriveExpander {
34            $($trait),*
35        }
36
37        impl BuiltinDeriveExpander {
38            pub fn expander(&self) -> fn(&dyn SourceDatabase, Span, &tt::TopSubtree) -> ExpandResult<tt::TopSubtree>  {
39                match *self {
40                    $( BuiltinDeriveExpander::$trait => $expand, )*
41                }
42            }
43
44            fn find_by_name(name: &name::Name) -> Option<Self> {
45                match name {
46                    $( id if id == &sym::$trait => Some(BuiltinDeriveExpander::$trait), )*
47                     _ => None,
48                }
49            }
50        }
51    };
52}
53
54impl BuiltinDeriveExpander {
55    pub fn expand(
56        &self,
57        db: &dyn SourceDatabase,
58        id: MacroCallId,
59        tt: &tt::TopSubtree,
60        span: Span,
61    ) -> ExpandResult<tt::TopSubtree> {
62        let span = span_with_def_site_ctxt(db, span, id.into(), Edition::CURRENT);
63        self.expander()(db, span, tt)
64    }
65}
66
67register_builtin! {
68    Copy => copy_expand,
69    Clone => clone_expand,
70    Default => default_expand,
71    Debug => debug_expand,
72    Hash => hash_expand,
73    Ord => ord_expand,
74    PartialOrd => partial_ord_expand,
75    Eq => eq_expand,
76    PartialEq => partial_eq_expand,
77    CoercePointee => coerce_pointee_expand,
78}
79
80pub fn find_builtin_derive(ident: &name::Name) -> Option<BuiltinDeriveExpander> {
81    BuiltinDeriveExpander::find_by_name(ident)
82}
83
84#[derive(Clone)]
85enum VariantShape {
86    Struct(Vec<tt::Ident>),
87    Tuple(usize),
88    Unit,
89}
90
91fn tuple_field_iterator(span: Span, n: usize) -> impl Iterator<Item = tt::Ident> {
92    (0..n).map(move |it| tt::Ident::new(&format!("f{it}"), span))
93}
94
95impl VariantShape {
96    fn as_pattern(&self, path: tt::TopSubtree, span: Span) -> tt::TopSubtree {
97        self.as_pattern_map(path, span, |it| quote!(span => #it))
98    }
99
100    fn field_names(&self, span: Span) -> Vec<tt::Ident> {
101        match self {
102            VariantShape::Struct(s) => s.clone(),
103            VariantShape::Tuple(n) => tuple_field_iterator(span, *n).collect(),
104            VariantShape::Unit => vec![],
105        }
106    }
107
108    fn as_pattern_map(
109        &self,
110        path: tt::TopSubtree,
111        span: Span,
112        field_map: impl Fn(&tt::Ident) -> tt::TopSubtree,
113    ) -> tt::TopSubtree {
114        match self {
115            VariantShape::Struct(fields) => {
116                let fields = fields.iter().map(|it| {
117                    let mapped = field_map(it);
118                    quote! {span => #it : #mapped , }
119                });
120                quote! {span =>
121                    #path { # #fields }
122                }
123            }
124            &VariantShape::Tuple(n) => {
125                let fields = tuple_field_iterator(span, n).map(|it| {
126                    let mapped = field_map(&it);
127                    quote! {span =>
128                        #mapped ,
129                    }
130                });
131                quote! {span =>
132                    #path ( # #fields )
133                }
134            }
135            VariantShape::Unit => path,
136        }
137    }
138
139    fn from(
140        call_site: Span,
141        tm: &ExpansionSpanMap,
142        value: Option<FieldList>,
143    ) -> Result<Self, ExpandError> {
144        let r = match value {
145            None => VariantShape::Unit,
146            Some(FieldList::RecordFieldList(it)) => VariantShape::Struct(
147                it.fields()
148                    .map(|it| it.name())
149                    .map(|it| name_to_token(call_site, tm, it))
150                    .collect::<Result<_, _>>()?,
151            ),
152            Some(FieldList::TupleFieldList(it)) => VariantShape::Tuple(it.fields().count()),
153        };
154        Ok(r)
155    }
156}
157
158#[derive(Clone)]
159enum AdtShape {
160    Struct(VariantShape),
161    Enum { variants: Vec<(tt::Ident, VariantShape)>, default_variant: Option<usize> },
162    Union,
163}
164
165impl AdtShape {
166    fn as_pattern(&self, span: Span, name: &tt::Ident) -> Vec<tt::TopSubtree> {
167        self.as_pattern_map(name, |it| quote!(span =>#it), span)
168    }
169
170    fn field_names(&self, span: Span) -> Vec<Vec<tt::Ident>> {
171        match self {
172            AdtShape::Struct(s) => {
173                vec![s.field_names(span)]
174            }
175            AdtShape::Enum { variants, .. } => {
176                variants.iter().map(|(_, fields)| fields.field_names(span)).collect()
177            }
178            AdtShape::Union => {
179                never!("using fields of union in derive is always wrong");
180                vec![]
181            }
182        }
183    }
184
185    fn as_pattern_map(
186        &self,
187        name: &tt::Ident,
188        field_map: impl Fn(&tt::Ident) -> tt::TopSubtree,
189        span: Span,
190    ) -> Vec<tt::TopSubtree> {
191        match self {
192            AdtShape::Struct(s) => {
193                vec![s.as_pattern_map(quote! {span => #name }, span, field_map)]
194            }
195            AdtShape::Enum { variants, .. } => variants
196                .iter()
197                .map(|(v, fields)| {
198                    fields.as_pattern_map(quote! {span => #name :: #v }, span, &field_map)
199                })
200                .collect(),
201            AdtShape::Union => {
202                never!("pattern matching on union is always wrong");
203                vec![quote! {span => un }]
204            }
205        }
206    }
207}
208
209#[derive(Clone)]
210struct BasicAdtInfo {
211    name: tt::Ident,
212    shape: AdtShape,
213    /// first field is the name, and
214    /// second field is `Some(ty)` if it's a const param of type `ty`, `None` if it's a type param.
215    /// third fields is where bounds, if any
216    param_types: Vec<AdtParam>,
217    where_clause: Vec<tt::TopSubtree>,
218    associated_types: Vec<tt::TopSubtree>,
219}
220
221#[derive(Clone)]
222struct AdtParam {
223    name: tt::TopSubtree,
224    /// `None` if this is a type parameter.
225    const_ty: Option<tt::TopSubtree>,
226    bounds: Option<tt::TopSubtree>,
227}
228
229// FIXME: This whole thing needs a refactor. Each derive requires its special values, and the result is a mess.
230fn parse_adt(
231    db: &dyn SourceDatabase,
232    tt: &tt::TopSubtree,
233    call_site: Span,
234) -> Result<BasicAdtInfo, ExpandError> {
235    let (adt, tm) = to_adt_syntax(db, tt, call_site)?;
236    parse_adt_from_syntax(&adt, &tm, call_site)
237}
238
239fn parse_adt_from_syntax(
240    adt: &ast::Adt,
241    tm: &span::SpanMap,
242    call_site: Span,
243) -> Result<BasicAdtInfo, ExpandError> {
244    let (name, generic_param_list, where_clause, shape) = match &adt {
245        ast::Adt::Struct(it) => (
246            it.name(),
247            it.generic_param_list(),
248            it.where_clause(),
249            AdtShape::Struct(VariantShape::from(call_site, tm, it.field_list())?),
250        ),
251        ast::Adt::Enum(it) => {
252            let default_variant = it
253                .variant_list()
254                .into_iter()
255                .flat_map(|it| it.variants())
256                .position(|it| it.attrs().any(|it| it.simple_name() == Some("default".into())));
257            (
258                it.name(),
259                it.generic_param_list(),
260                it.where_clause(),
261                AdtShape::Enum {
262                    default_variant,
263                    variants: it
264                        .variant_list()
265                        .into_iter()
266                        .flat_map(|it| it.variants())
267                        .map(|it| {
268                            Ok((
269                                name_to_token(call_site, tm, it.name())?,
270                                VariantShape::from(call_site, tm, it.field_list())?,
271                            ))
272                        })
273                        .collect::<Result<_, ExpandError>>()?,
274                },
275            )
276        }
277        ast::Adt::Union(it) => {
278            (it.name(), it.generic_param_list(), it.where_clause(), AdtShape::Union)
279        }
280    };
281
282    let mut param_type_set: FxHashSet<Name> = FxHashSet::default();
283    let param_types = generic_param_list
284        .into_iter()
285        .flat_map(|param_list| param_list.type_or_const_params())
286        .map(|param| {
287            let name = {
288                let this = param.name();
289                match this {
290                    Some(it) => {
291                        param_type_set.insert(it.as_name());
292                        syntax_bridge::syntax_node_to_token_tree(
293                            it.syntax(),
294                            tm,
295                            call_site,
296                            DocCommentDesugarMode::ProcMacro,
297                        )
298                    }
299                    None => {
300                        tt::TopSubtree::empty(::tt::DelimSpan { open: call_site, close: call_site })
301                    }
302                }
303            };
304            let bounds = match &param {
305                ast::TypeOrConstParam::Type(it) => it.type_bound_list().map(|it| {
306                    syntax_bridge::syntax_node_to_token_tree(
307                        it.syntax(),
308                        tm,
309                        call_site,
310                        DocCommentDesugarMode::ProcMacro,
311                    )
312                }),
313                ast::TypeOrConstParam::Const(_) => None,
314            };
315            let const_ty = if let ast::TypeOrConstParam::Const(param) = param {
316                let ty = param
317                    .ty()
318                    .map(|ty| {
319                        syntax_bridge::syntax_node_to_token_tree(
320                            ty.syntax(),
321                            tm,
322                            call_site,
323                            DocCommentDesugarMode::ProcMacro,
324                        )
325                    })
326                    .unwrap_or_else(|| {
327                        tt::TopSubtree::empty(::tt::DelimSpan { open: call_site, close: call_site })
328                    });
329                Some(ty)
330            } else {
331                None
332            };
333            AdtParam { name, const_ty, bounds }
334        })
335        .collect();
336
337    let where_clause = if let Some(w) = where_clause {
338        w.predicates()
339            .map(|it| {
340                syntax_bridge::syntax_node_to_token_tree(
341                    it.syntax(),
342                    tm,
343                    call_site,
344                    DocCommentDesugarMode::ProcMacro,
345                )
346            })
347            .collect()
348    } else {
349        vec![]
350    };
351
352    // For a generic parameter `T`, when shorthand associated type `T::Assoc` appears in field
353    // types (of any variant for enums), we generate trait bound for it. It sounds reasonable to
354    // also generate trait bound for qualified associated type `<T as Trait>::Assoc`, but rustc
355    // does not do that for some unknown reason.
356    //
357    // See the analogous function in rustc [find_type_parameters()] and rust-lang/rust#50730.
358    // [find_type_parameters()]: https://github.com/rust-lang/rust/blob/1.70.0/compiler/rustc_builtin_macros/src/deriving/generic/mod.rs#L378
359
360    // It's cumbersome to deal with the distinct structures of ADTs, so let's just get untyped
361    // `SyntaxNode` that contains fields and look for descendant `ast::PathType`s. Of note is that
362    // we should not inspect `ast::PathType`s in parameter bounds and where clauses.
363    let field_list = match adt {
364        ast::Adt::Enum(it) => it.variant_list().map(|list| list.syntax().clone()),
365        ast::Adt::Struct(it) => it.field_list().map(|list| list.syntax().clone()),
366        ast::Adt::Union(it) => it.record_field_list().map(|list| list.syntax().clone()),
367    };
368    let associated_types = field_list
369        .into_iter()
370        .flat_map(|it| it.descendants())
371        .filter_map(ast::PathType::cast)
372        .filter_map(|p| {
373            let name = p.path()?.qualifier()?.as_single_name_ref()?.as_name();
374            param_type_set.contains(&name).then_some(p)
375        })
376        .map(|it| {
377            syntax_bridge::syntax_node_to_token_tree(
378                it.syntax(),
379                tm,
380                call_site,
381                DocCommentDesugarMode::ProcMacro,
382            )
383        })
384        .collect();
385    let name_token = name_to_token(call_site, tm, name)?;
386    Ok(BasicAdtInfo { name: name_token, shape, param_types, where_clause, associated_types })
387}
388
389fn to_adt_syntax(
390    db: &dyn SourceDatabase,
391    tt: &tt::TopSubtree,
392    call_site: Span,
393) -> Result<(ast::Adt, span::SpanMap), ExpandError> {
394    let (parsed, tm) = crate::token_tree_to_syntax_node(db, tt, crate::ExpandTo::Items);
395    let macro_items = ast::MacroItems::cast(parsed.syntax_node())
396        .ok_or_else(|| ExpandError::other(call_site, "invalid item definition"))?;
397    let item =
398        macro_items.items().next().ok_or_else(|| ExpandError::other(call_site, "no item found"))?;
399    let adt = ast::Adt::cast(item.syntax().clone())
400        .ok_or_else(|| ExpandError::other(call_site, "expected struct, enum or union"))?;
401    Ok((adt, tm))
402}
403
404fn name_to_token(
405    call_site: Span,
406    token_map: &ExpansionSpanMap,
407    name: Option<ast::Name>,
408) -> Result<tt::Ident, ExpandError> {
409    let name = name.ok_or_else(|| {
410        debug!("parsed item has no name");
411        ExpandError::other(call_site, "missing name")
412    })?;
413    let span = token_map.span_at(name.syntax().text_range().start());
414
415    let name_token = tt::Ident::new(name.text().as_ref(), span);
416    Ok(name_token)
417}
418
419/// Given that we are deriving a trait `DerivedTrait` for a type like:
420///
421/// ```ignore (only-for-syntax-highlight)
422/// struct Struct<'a, ..., 'z, A, B: DeclaredTrait, C, ..., Z> where C: WhereTrait {
423///     a: A,
424///     b: B::Item,
425///     b1: <B as DeclaredTrait>::Item,
426///     c1: <C as WhereTrait>::Item,
427///     c2: Option<<C as WhereTrait>::Item>,
428///     ...
429/// }
430/// ```
431///
432/// create an impl like:
433///
434/// ```ignore (only-for-syntax-highlight)
435/// impl<'a, ..., 'z, A, B: DeclaredTrait, C, ... Z> where
436///     C:                       WhereTrait,
437///     A: DerivedTrait + B1 + ... + BN,
438///     B: DerivedTrait + B1 + ... + BN,
439///     C: DerivedTrait + B1 + ... + BN,
440///     B::Item:                 DerivedTrait + B1 + ... + BN,
441///     <C as WhereTrait>::Item: DerivedTrait + B1 + ... + BN,
442///     ...
443/// {
444///     ...
445/// }
446/// ```
447///
448/// where B1, ..., BN are the bounds given by `bounds_paths`. Z is a phantom type, and
449/// therefore does not get bound by the derived trait.
450fn expand_simple_derive(
451    db: &dyn SourceDatabase,
452    invoc_span: Span,
453    tt: &tt::TopSubtree,
454    trait_path: tt::TopSubtree,
455    allow_unions: bool,
456    make_trait_body: impl FnOnce(&BasicAdtInfo) -> tt::TopSubtree,
457) -> ExpandResult<tt::TopSubtree> {
458    let info = match parse_adt(db, tt, invoc_span) {
459        Ok(info) => info,
460        Err(e) => {
461            return ExpandResult::new(
462                tt::TopSubtree::empty(tt::DelimSpan { open: invoc_span, close: invoc_span }),
463                e,
464            );
465        }
466    };
467    if !allow_unions && matches!(info.shape, AdtShape::Union) {
468        return ExpandResult::new(
469            tt::TopSubtree::empty(tt::DelimSpan::from_single(invoc_span)),
470            ExpandError::other(invoc_span, "this trait cannot be derived for unions"),
471        );
472    }
473    ExpandResult::ok(expand_simple_derive_with_parsed(
474        invoc_span,
475        info,
476        trait_path,
477        make_trait_body,
478        true,
479        tt::TopSubtree::empty(tt::DelimSpan::from_single(invoc_span)),
480    ))
481}
482
483fn expand_simple_derive_with_parsed(
484    invoc_span: Span,
485    info: BasicAdtInfo,
486    trait_path: tt::TopSubtree,
487    make_trait_body: impl FnOnce(&BasicAdtInfo) -> tt::TopSubtree,
488    constrain_to_trait: bool,
489    extra_impl_params: tt::TopSubtree,
490) -> tt::TopSubtree {
491    let trait_body = make_trait_body(&info);
492    let mut where_block: Vec<_> =
493        info.where_clause.into_iter().map(|w| quote! {invoc_span => #w , }).collect();
494    let (params, args): (Vec<_>, Vec<_>) = info
495        .param_types
496        .into_iter()
497        .map(|param| {
498            let ident = param.name;
499            if let Some(b) = param.bounds {
500                let ident2 = ident.clone();
501                where_block.push(quote! {invoc_span => #ident2 : #b , });
502            }
503            if let Some(ty) = param.const_ty {
504                let ident2 = ident.clone();
505                (quote! {invoc_span => const #ident : #ty , }, quote! {invoc_span => #ident2 , })
506            } else {
507                let bound = trait_path.clone();
508                let ident2 = ident.clone();
509                let param = if constrain_to_trait {
510                    quote! {invoc_span => #ident : #bound , }
511                } else {
512                    quote! {invoc_span => #ident , }
513                };
514                (param, quote! {invoc_span => #ident2 , })
515            }
516        })
517        .unzip();
518
519    if constrain_to_trait {
520        where_block.extend(info.associated_types.iter().map(|it| {
521            let it = it.clone();
522            let bound = trait_path.clone();
523            quote! {invoc_span => #it : #bound , }
524        }));
525    }
526
527    let name = info.name;
528    quote! {invoc_span =>
529        impl < # #params #extra_impl_params > #trait_path for #name < # #args > where # #where_block { #trait_body }
530    }
531}
532
533fn copy_expand(
534    db: &dyn SourceDatabase,
535    span: Span,
536    tt: &tt::TopSubtree,
537) -> ExpandResult<tt::TopSubtree> {
538    let krate = dollar_crate(span);
539    expand_simple_derive(
540        db,
541        span,
542        tt,
543        quote! {span => #krate::marker::Copy },
544        true,
545        |_| quote! {span =>},
546    )
547}
548
549fn clone_expand(
550    db: &dyn SourceDatabase,
551    span: Span,
552    tt: &tt::TopSubtree,
553) -> ExpandResult<tt::TopSubtree> {
554    let krate = dollar_crate(span);
555    expand_simple_derive(db, span, tt, quote! {span => #krate::clone::Clone }, true, |adt| {
556        if matches!(adt.shape, AdtShape::Union) {
557            let star = tt::Punct { char: '*', spacing: ::tt::Spacing::Alone, span };
558            return quote! {span =>
559                fn clone(&self) -> Self {
560                    #star self
561                }
562            };
563        }
564        if matches!(&adt.shape, AdtShape::Enum { variants, .. } if variants.is_empty()) {
565            let star = tt::Punct { char: '*', spacing: ::tt::Spacing::Alone, span };
566            return quote! {span =>
567                fn clone(&self) -> Self {
568                    match #star self {}
569                }
570            };
571        }
572        let name = &adt.name;
573        let patterns = adt.shape.as_pattern(span, name);
574        let exprs = adt.shape.as_pattern_map(name, |it| quote! {span => #it .clone() }, span);
575        let arms = patterns.into_iter().zip(exprs).map(|(pat, expr)| {
576            let fat_arrow = fat_arrow(span);
577            quote! {span =>
578                #pat #fat_arrow #expr,
579            }
580        });
581
582        quote! {span =>
583            fn clone(&self) -> Self {
584                match self {
585                    # #arms
586                }
587            }
588        }
589    })
590}
591
592/// This function exists since `quote! {span => => }` doesn't work.
593fn fat_arrow(span: Span) -> tt::TopSubtree {
594    let eq = tt::Punct { char: '=', spacing: ::tt::Spacing::Joint, span };
595    quote! {span => #eq> }
596}
597
598/// This function exists since `quote! {span => && }` doesn't work.
599fn and_and(span: Span) -> tt::TopSubtree {
600    let and = tt::Punct { char: '&', spacing: ::tt::Spacing::Joint, span };
601    quote! {span => #and& }
602}
603
604fn default_expand(
605    db: &dyn SourceDatabase,
606    span: Span,
607    tt: &tt::TopSubtree,
608) -> ExpandResult<tt::TopSubtree> {
609    let krate = &dollar_crate(span);
610    let adt = match parse_adt(db, tt, span) {
611        Ok(info) => info,
612        Err(e) => {
613            return ExpandResult::new(
614                tt::TopSubtree::empty(tt::DelimSpan { open: span, close: span }),
615                e,
616            );
617        }
618    };
619    let (body, constrain_to_trait) = match &adt.shape {
620        AdtShape::Struct(fields) => {
621            let name = &adt.name;
622            let body = fields.as_pattern_map(
623                quote!(span =>#name),
624                span,
625                |_| quote!(span =>#krate::default::Default::default()),
626            );
627            (body, true)
628        }
629        AdtShape::Enum { default_variant, variants } => {
630            if let Some(d) = default_variant {
631                let (name, fields) = &variants[*d];
632                let adt_name = &adt.name;
633                let body = fields.as_pattern_map(
634                    quote!(span =>#adt_name :: #name),
635                    span,
636                    |_| quote!(span =>#krate::default::Default::default()),
637                );
638                (body, false)
639            } else {
640                return ExpandResult::new(
641                    tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
642                    ExpandError::other(span, "`#[derive(Default)]` on enum with no `#[default]`"),
643                );
644            }
645        }
646        AdtShape::Union => {
647            return ExpandResult::new(
648                tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
649                ExpandError::other(span, "this trait cannot be derived for unions"),
650            );
651        }
652    };
653    ExpandResult::ok(expand_simple_derive_with_parsed(
654        span,
655        adt,
656        quote! {span => #krate::default::Default },
657        |_adt| {
658            quote! {span =>
659                fn default() -> Self {
660                    #body
661                }
662            }
663        },
664        constrain_to_trait,
665        tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
666    ))
667}
668
669fn debug_expand(
670    db: &dyn SourceDatabase,
671    span: Span,
672    tt: &tt::TopSubtree,
673) -> ExpandResult<tt::TopSubtree> {
674    let krate = &dollar_crate(span);
675    expand_simple_derive(db, span, tt, quote! {span => #krate::fmt::Debug }, false, |adt| {
676        let for_variant = |name: String, v: &VariantShape| match v {
677            VariantShape::Struct(fields) => {
678                let for_fields = fields.iter().map(|it| {
679                    let x_string = it.to_string();
680                    quote! {span =>
681                        .field(#x_string, & #it)
682                    }
683                });
684                quote! {span =>
685                    f.debug_struct(#name) # #for_fields .finish()
686                }
687            }
688            VariantShape::Tuple(n) => {
689                let for_fields = tuple_field_iterator(span, *n).map(|it| {
690                    quote! {span =>
691                        .field( & #it)
692                    }
693                });
694                quote! {span =>
695                    f.debug_tuple(#name) # #for_fields .finish()
696                }
697            }
698            VariantShape::Unit => quote! {span =>
699                f.write_str(#name)
700            },
701        };
702        if matches!(&adt.shape, AdtShape::Enum { variants, .. } if variants.is_empty()) {
703            let star = tt::Punct { char: '*', spacing: ::tt::Spacing::Alone, span };
704            return quote! {span =>
705                fn fmt(&self, f: &mut #krate::fmt::Formatter) -> #krate::fmt::Result {
706                    match #star self {}
707                }
708            };
709        }
710        let arms = match &adt.shape {
711            AdtShape::Struct(fields) => {
712                let fat_arrow = fat_arrow(span);
713                let name = &adt.name;
714                let pat = fields.as_pattern(quote!(span =>#name), span);
715                let expr = for_variant(name.to_string(), fields);
716                vec![quote! {span => #pat #fat_arrow #expr }]
717            }
718            AdtShape::Enum { variants, .. } => variants
719                .iter()
720                .map(|(name, v)| {
721                    let fat_arrow = fat_arrow(span);
722                    let adt_name = &adt.name;
723                    let pat = v.as_pattern(quote!(span =>#adt_name :: #name), span);
724                    let expr = for_variant(name.to_string(), v);
725                    quote! {span =>
726                        #pat #fat_arrow #expr ,
727                    }
728                })
729                .collect(),
730            AdtShape::Union => unreachable!(),
731        };
732        quote! {span =>
733            fn fmt(&self, f: &mut #krate::fmt::Formatter) -> #krate::fmt::Result {
734                match self {
735                    # #arms
736                }
737            }
738        }
739    })
740}
741
742fn hash_expand(
743    db: &dyn SourceDatabase,
744    span: Span,
745    tt: &tt::TopSubtree,
746) -> ExpandResult<tt::TopSubtree> {
747    let krate = &dollar_crate(span);
748    expand_simple_derive(db, span, tt, quote! {span => #krate::hash::Hash }, false, |adt| {
749        if matches!(&adt.shape, AdtShape::Enum { variants, .. } if variants.is_empty()) {
750            let star = tt::Punct { char: '*', spacing: ::tt::Spacing::Alone, span };
751            return quote! {span =>
752                fn hash<H: #krate::hash::Hasher>(&self, ra_expand_state: &mut H) {
753                    match #star self {}
754                }
755            };
756        }
757        let arms =
758            adt.shape.as_pattern(span, &adt.name).into_iter().zip(adt.shape.field_names(span)).map(
759                |(pat, names)| {
760                    let expr = {
761                        let it =
762                            names.iter().map(|it| quote! {span => #it . hash(ra_expand_state); });
763                        quote! {span => {
764                            # #it
765                        } }
766                    };
767                    let fat_arrow = fat_arrow(span);
768                    quote! {span =>
769                        #pat #fat_arrow #expr ,
770                    }
771                },
772            );
773        let check_discriminant = if matches!(&adt.shape, AdtShape::Enum { .. }) {
774            quote! {span => #krate::mem::discriminant(self).hash(ra_expand_state); }
775        } else {
776            quote! {span =>}
777        };
778        quote! {span =>
779            fn hash<H: #krate::hash::Hasher>(&self, ra_expand_state: &mut H) {
780                #check_discriminant
781                match self {
782                    # #arms
783                }
784            }
785        }
786    })
787}
788
789fn eq_expand(
790    db: &dyn SourceDatabase,
791    span: Span,
792    tt: &tt::TopSubtree,
793) -> ExpandResult<tt::TopSubtree> {
794    let krate = dollar_crate(span);
795    expand_simple_derive(
796        db,
797        span,
798        tt,
799        quote! {span => #krate::cmp::Eq },
800        true,
801        |_| quote! {span =>},
802    )
803}
804
805fn partial_eq_expand(
806    db: &dyn SourceDatabase,
807    span: Span,
808    tt: &tt::TopSubtree,
809) -> ExpandResult<tt::TopSubtree> {
810    let krate = dollar_crate(span);
811    expand_simple_derive(db, span, tt, quote! {span => #krate::cmp::PartialEq }, false, |adt| {
812        let name = &adt.name;
813
814        let (self_patterns, other_patterns) = self_and_other_patterns(adt, name, span);
815        let arms = izip!(self_patterns, other_patterns, adt.shape.field_names(span)).map(
816            |(pat1, pat2, names)| {
817                let fat_arrow = fat_arrow(span);
818                let body = match &*names {
819                    [] => {
820                        quote!(span =>true)
821                    }
822                    [first, rest @ ..] => {
823                        let rest = rest.iter().map(|it| {
824                            let t1 = tt::Ident::new(&format!("{}_self", it.sym), it.span);
825                            let t2 = tt::Ident::new(&format!("{}_other", it.sym), it.span);
826                            let and_and = and_and(span);
827                            quote!(span =>#and_and #t1 .eq( #t2 ))
828                        });
829                        let first = {
830                            let t1 = tt::Ident::new(&format!("{}_self", first.sym), first.span);
831                            let t2 = tt::Ident::new(&format!("{}_other", first.sym), first.span);
832                            quote!(span =>#t1 .eq( #t2 ))
833                        };
834                        quote!(span =>#first # #rest)
835                    }
836                };
837                quote! {span => ( #pat1 , #pat2 ) #fat_arrow #body , }
838            },
839        );
840
841        let fat_arrow = fat_arrow(span);
842        quote! {span =>
843            fn eq(&self, other: &Self) -> bool {
844                match (self, other) {
845                    # #arms
846                    _unused #fat_arrow false
847                }
848            }
849        }
850    })
851}
852
853fn self_and_other_patterns(
854    adt: &BasicAdtInfo,
855    name: &tt::Ident,
856    span: Span,
857) -> (Vec<tt::TopSubtree>, Vec<tt::TopSubtree>) {
858    let self_patterns = adt.shape.as_pattern_map(
859        name,
860        |it| {
861            let t = tt::Ident::new(&format!("{}_self", it.sym), it.span);
862            quote!(span =>#t)
863        },
864        span,
865    );
866    let other_patterns = adt.shape.as_pattern_map(
867        name,
868        |it| {
869            let t = tt::Ident::new(&format!("{}_other", it.sym), it.span);
870            quote!(span =>#t)
871        },
872        span,
873    );
874    (self_patterns, other_patterns)
875}
876
877fn ord_expand(
878    db: &dyn SourceDatabase,
879    span: Span,
880    tt: &tt::TopSubtree,
881) -> ExpandResult<tt::TopSubtree> {
882    let krate = &dollar_crate(span);
883    expand_simple_derive(db, span, tt, quote! {span => #krate::cmp::Ord }, false, |adt| {
884        fn compare(
885            krate: &tt::Ident,
886            left: tt::TopSubtree,
887            right: tt::TopSubtree,
888            rest: tt::TopSubtree,
889            span: Span,
890        ) -> tt::TopSubtree {
891            let fat_arrow1 = fat_arrow(span);
892            let fat_arrow2 = fat_arrow(span);
893            quote! {span =>
894                match #left.cmp(&#right) {
895                    #krate::cmp::Ordering::Equal #fat_arrow1 {
896                        #rest
897                    }
898                    c #fat_arrow2 return c,
899                }
900            }
901        }
902        let (self_patterns, other_patterns) = self_and_other_patterns(adt, &adt.name, span);
903        let arms = izip!(self_patterns, other_patterns, adt.shape.field_names(span)).map(
904            |(pat1, pat2, fields)| {
905                let mut body = quote!(span =>#krate::cmp::Ordering::Equal);
906                for f in fields.into_iter().rev() {
907                    let t1 = tt::Ident::new(&format!("{}_self", f.sym), f.span);
908                    let t2 = tt::Ident::new(&format!("{}_other", f.sym), f.span);
909                    body = compare(krate, quote!(span =>#t1), quote!(span =>#t2), body, span);
910                }
911                let fat_arrow = fat_arrow(span);
912                quote! {span => ( #pat1 , #pat2 ) #fat_arrow #body , }
913            },
914        );
915        let fat_arrow = fat_arrow(span);
916        let mut body = quote! {span =>
917            match (self, other) {
918                # #arms
919                _unused #fat_arrow #krate::cmp::Ordering::Equal
920            }
921        };
922        if matches!(&adt.shape, AdtShape::Enum { .. }) {
923            let left = quote!(span =>#krate::intrinsics::discriminant_value(self));
924            let right = quote!(span =>#krate::intrinsics::discriminant_value(other));
925            body = compare(krate, left, right, body, span);
926        }
927        quote! {span =>
928            fn cmp(&self, other: &Self) -> #krate::cmp::Ordering {
929                #body
930            }
931        }
932    })
933}
934
935fn partial_ord_expand(
936    db: &dyn SourceDatabase,
937    span: Span,
938    tt: &tt::TopSubtree,
939) -> ExpandResult<tt::TopSubtree> {
940    let krate = &dollar_crate(span);
941    expand_simple_derive(db, span, tt, quote! {span => #krate::cmp::PartialOrd }, false, |adt| {
942        fn compare(
943            krate: &tt::Ident,
944            left: tt::TopSubtree,
945            right: tt::TopSubtree,
946            rest: tt::TopSubtree,
947            span: Span,
948        ) -> tt::TopSubtree {
949            let fat_arrow1 = fat_arrow(span);
950            let fat_arrow2 = fat_arrow(span);
951            quote! {span =>
952                match #left.partial_cmp(&#right) {
953                    #krate::option::Option::Some(#krate::cmp::Ordering::Equal) #fat_arrow1 {
954                        #rest
955                    }
956                    c #fat_arrow2 return c,
957                }
958            }
959        }
960        let left = quote!(span =>#krate::intrinsics::discriminant_value(self));
961        let right = quote!(span =>#krate::intrinsics::discriminant_value(other));
962
963        let (self_patterns, other_patterns) = self_and_other_patterns(adt, &adt.name, span);
964        let arms = izip!(self_patterns, other_patterns, adt.shape.field_names(span)).map(
965            |(pat1, pat2, fields)| {
966                let mut body =
967                    quote!(span =>#krate::option::Option::Some(#krate::cmp::Ordering::Equal));
968                for f in fields.into_iter().rev() {
969                    let t1 = tt::Ident::new(&format!("{}_self", f.sym), f.span);
970                    let t2 = tt::Ident::new(&format!("{}_other", f.sym), f.span);
971                    body = compare(krate, quote!(span =>#t1), quote!(span =>#t2), body, span);
972                }
973                let fat_arrow = fat_arrow(span);
974                quote! {span => ( #pat1 , #pat2 ) #fat_arrow #body , }
975            },
976        );
977        let fat_arrow = fat_arrow(span);
978        let body = compare(
979            krate,
980            left,
981            right,
982            quote! {span =>
983                match (self, other) {
984                    # #arms
985                    _unused #fat_arrow #krate::option::Option::Some(#krate::cmp::Ordering::Equal)
986                }
987            },
988            span,
989        );
990        quote! {span =>
991            fn partial_cmp(&self, other: &Self) -> #krate::option::Option<#krate::cmp::Ordering> {
992                #body
993            }
994        }
995    })
996}
997
998fn coerce_pointee_expand(
999    db: &dyn SourceDatabase,
1000    span: Span,
1001    tt: &tt::TopSubtree,
1002) -> ExpandResult<tt::TopSubtree> {
1003    let (adt, _span_map) = match to_adt_syntax(db, tt, span) {
1004        Ok(it) => it,
1005        Err(err) => {
1006            return ExpandResult::new(tt::TopSubtree::empty(tt::DelimSpan::from_single(span)), err);
1007        }
1008    };
1009    let (editor, adt) = SyntaxEditor::with_ast_node(&adt);
1010    let make = editor.make();
1011    let ast::Adt::Struct(strukt) = &adt else {
1012        return ExpandResult::new(
1013            tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1014            ExpandError::other(span, "`CoercePointee` can only be derived on `struct`s"),
1015        );
1016    };
1017    let has_at_least_one_field = strukt
1018        .field_list()
1019        .map(|it| match it {
1020            ast::FieldList::RecordFieldList(it) => it.fields().next().is_some(),
1021            ast::FieldList::TupleFieldList(it) => it.fields().next().is_some(),
1022        })
1023        .unwrap_or(false);
1024    if !has_at_least_one_field {
1025        return ExpandResult::new(
1026            tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1027            ExpandError::other(
1028                span,
1029                "`CoercePointee` can only be derived on `struct`s with at least one field",
1030            ),
1031        );
1032    }
1033    let is_repr_transparent = strukt.attrs().any(|attr| {
1034        attr.as_simple_call().is_some_and(|(name, tt)| {
1035            name == "repr"
1036                && tt.syntax().children_with_tokens().any(|it| {
1037                    it.into_token().is_some_and(|it| {
1038                        it.kind() == SyntaxKind::IDENT && it.text() == "transparent"
1039                    })
1040                })
1041        })
1042    });
1043    if !is_repr_transparent {
1044        return ExpandResult::new(
1045            tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1046            ExpandError::other(
1047                span,
1048                "`CoercePointee` can only be derived on `struct`s with `#[repr(transparent)]`",
1049            ),
1050        );
1051    }
1052    let type_params = strukt
1053        .generic_param_list()
1054        .into_iter()
1055        .flat_map(|generics| {
1056            generics.generic_params().filter_map(|param| match param {
1057                ast::GenericParam::TypeParam(param) => Some(param),
1058                _ => None,
1059            })
1060        })
1061        .collect_vec();
1062    if type_params.is_empty() {
1063        return ExpandResult::new(
1064            tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1065            ExpandError::other(
1066                span,
1067                "`CoercePointee` can only be derived on `struct`s that are generic over at least one type",
1068            ),
1069        );
1070    }
1071    let (pointee_param, pointee_param_idx) = if type_params.len() == 1 {
1072        // Regardless of the only type param being designed as `#[pointee]` or not, we can just use it as such.
1073        (type_params[0].clone(), 0)
1074    } else {
1075        let mut pointees = type_params.iter().cloned().enumerate().filter(|(_, param)| {
1076            param.attrs().any(|attr| {
1077                let is_pointee = attr.as_simple_atom().is_some_and(|name| name == "pointee");
1078                if is_pointee {
1079                    // Remove the `#[pointee]` attribute so it won't be present in the generated
1080                    // impls (where we cannot resolve it).
1081                    editor.delete(attr.syntax());
1082                }
1083                is_pointee
1084            })
1085        });
1086        match (pointees.next(), pointees.next()) {
1087            (Some((pointee_idx, pointee)), None) => (pointee, pointee_idx),
1088            (None, _) => {
1089                return ExpandResult::new(
1090                    tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1091                    ExpandError::other(
1092                        span,
1093                        "exactly one generic type parameter must be marked \
1094                                as `#[pointee]` to derive `CoercePointee` traits",
1095                    ),
1096                );
1097            }
1098            (Some(_), Some(_)) => {
1099                return ExpandResult::new(
1100                    tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1101                    ExpandError::other(
1102                        span,
1103                        "only one type parameter can be marked as `#[pointee]` \
1104                                when deriving `CoercePointee` traits",
1105                    ),
1106                );
1107            }
1108        }
1109    };
1110    let (Some(struct_name), Some(pointee_param_name)) = (strukt.name(), pointee_param.name())
1111    else {
1112        return ExpandResult::new(
1113            tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1114            ExpandError::other(span, "invalid item"),
1115        );
1116    };
1117
1118    {
1119        let mut pointee_has_maybe_sized_bound = false;
1120        if let Some(bounds) = pointee_param.type_bound_list() {
1121            pointee_has_maybe_sized_bound |= bounds.bounds().any(is_maybe_sized_bound);
1122        }
1123        if let Some(where_clause) = strukt.where_clause() {
1124            pointee_has_maybe_sized_bound |= where_clause.predicates().any(|pred| {
1125                let Some(ast::Type::PathType(ty)) = pred.ty() else { return false };
1126                let is_not_pointee = ty.path().is_none_or(|path| {
1127                    let is_pointee = path
1128                        .as_single_name_ref()
1129                        .is_some_and(|name| name.text() == pointee_param_name.text());
1130                    !is_pointee
1131                });
1132                if is_not_pointee {
1133                    return false;
1134                }
1135                pred.type_bound_list()
1136                    .is_some_and(|bounds| bounds.bounds().any(is_maybe_sized_bound))
1137            })
1138        }
1139        if !pointee_has_maybe_sized_bound {
1140            return ExpandResult::new(
1141                tt::TopSubtree::empty(tt::DelimSpan::from_single(span)),
1142                ExpandError::other(
1143                    span,
1144                    format!(
1145                        "`derive(CoercePointee)` requires `{pointee_param_name}` to be marked `?Sized`"
1146                    ),
1147                ),
1148            );
1149        }
1150    }
1151
1152    const ADDED_PARAM: &str = "__S";
1153
1154    let mut new_predicates: Vec<ast::WherePred> = Vec::new();
1155
1156    {
1157        // # Rewrite generic parameter bounds
1158        // For each bound `U: ..` in `struct<U: ..>`, make a new bound with `__S` in place of `#[pointee]`
1159        // Example:
1160        // ```
1161        // struct<
1162        //     U: Trait<T>,
1163        //     #[pointee] T: Trait<T> + ?Sized,
1164        //     V: Trait<T>> ...
1165        // ```
1166        // ... generates this `impl` generic parameters
1167        // ```
1168        // impl<
1169        //     U: Trait<T>,
1170        //     T: Trait<T> + ?Sized,
1171        //     V: Trait<T>
1172        // >
1173        // where
1174        //     U: Trait<__S>,
1175        //     __S: Trait<__S> + ?Sized,
1176        //     V: Trait<__S> ...
1177        // ```
1178        for param in &type_params {
1179            let Some(param_name) = param.name() else { continue };
1180            if let Some(bounds) = param.type_bound_list() {
1181                // If the target type is the pointee, duplicate the bound as whole.
1182                // Otherwise, duplicate only bounds that mention the pointee.
1183                let is_pointee = param_name.text() == pointee_param_name.text();
1184                let new_bounds = bounds.bounds().filter_map(|bound| {
1185                    let new_bound = substitute_type_bound(
1186                        bound.clone(),
1187                        &pointee_param_name.text(),
1188                        ADDED_PARAM,
1189                    );
1190
1191                    if is_pointee {
1192                        return new_bound.or(Some(bound));
1193                    }
1194                    new_bound
1195                });
1196
1197                let new_bounds_target = if is_pointee {
1198                    make.name_ref(ADDED_PARAM)
1199                } else {
1200                    make.name_ref(&param_name.text())
1201                };
1202                new_predicates.push(make.where_pred(
1203                    Either::Right(
1204                        make.ty_path_from_segments([make.path_segment(new_bounds_target)], false),
1205                    ),
1206                    new_bounds,
1207                ));
1208            }
1209        }
1210
1211        // # Rewrite `where` clauses
1212        //
1213        // Move on to `where` clauses.
1214        // Example:
1215        // ```
1216        // struct MyPointer<#[pointee] T, ..>
1217        // where
1218        //   U: Trait<V> + Trait<T>,
1219        //   Companion<T>: Trait<T>,
1220        //   T: Trait<T> + ?Sized,
1221        // { .. }
1222        // ```
1223        // ... will have a impl prelude like so
1224        // ```
1225        // impl<..> ..
1226        // where
1227        //   U: Trait<V> + Trait<T>,
1228        //   U: Trait<__S>,
1229        //   Companion<T>: Trait<T>,
1230        //   Companion<__S>: Trait<__S>,
1231        //   T: Trait<T> + ?Sized,
1232        //   __S: Trait<__S> + ?Sized,
1233        // ```
1234        //
1235        // We should also write a few new `where` bounds from `#[pointee] T` to `__S`
1236        // as well as any bound that indirectly involves the `#[pointee] T` type.
1237        for predicate in strukt.where_clause().into_iter().flat_map(|wc| wc.predicates()) {
1238            let Some(pred_target) = predicate.ty() else { continue };
1239
1240            // If the target type references the pointee, duplicate the bound as whole.
1241            // Otherwise, duplicate only bounds that mention the pointee.
1242            if let Some(predicate_with_substituted_target) =
1243                substitute_where_pred(&predicate, &pointee_param_name.text(), ADDED_PARAM)
1244            {
1245                new_predicates.push(predicate_with_substituted_target);
1246            } else if let Some(bounds) = predicate.type_bound_list() {
1247                let new_bounds = bounds.bounds().filter_map(|bound| {
1248                    substitute_type_bound(bound, &pointee_param_name.text(), ADDED_PARAM)
1249                });
1250                new_predicates.push(make.where_pred(Either::Right(pred_target), new_bounds));
1251            }
1252        }
1253    }
1254
1255    {
1256        // # Add `Unsize<__S>` bound to `#[pointee]` at the generic parameter location
1257        //
1258        // Find the `#[pointee]` parameter and add an `Unsize<__S>` bound to it.
1259        new_predicates.push(
1260            make.where_pred(
1261                Either::Right(make.ty_path_from_segments(
1262                    [make.path_segment(make.name_ref(&pointee_param_name.text()))],
1263                    false,
1264                )),
1265                [make.type_bound(
1266                    make.ty_path_from_segments(
1267                        [
1268                            make.path_segment(make.name_ref("core")),
1269                            make.path_segment(make.name_ref("marker")),
1270                            make.generic_ty_path_segment(
1271                                make.name_ref("Unsize"),
1272                                [make
1273                                    .type_arg(make.ty_path_from_segments(
1274                                        [make.path_segment(make.name_ref(ADDED_PARAM))],
1275                                        false,
1276                                    ))
1277                                    .into()],
1278                            ),
1279                        ],
1280                        true,
1281                    ),
1282                )],
1283            ),
1284        );
1285    }
1286
1287    let self_for_traits = {
1288        // Replace the `#[pointee]` with `__S`.
1289        let mut type_param_idx = 0;
1290        let self_params_for_traits = strukt
1291            .generic_param_list()
1292            .into_iter()
1293            .flat_map(|params| params.generic_params())
1294            .filter_map(|param| {
1295                Some(match param {
1296                    ast::GenericParam::ConstParam(param) => {
1297                        ast::GenericArg::ConstArg(make.expr_const_value(&param.name()?.text()))
1298                    }
1299                    ast::GenericParam::LifetimeParam(param) => {
1300                        make.lifetime_arg(param.lifetime()?).into()
1301                    }
1302                    ast::GenericParam::TypeParam(param) => {
1303                        let name = if pointee_param_idx == type_param_idx {
1304                            make.name_ref(ADDED_PARAM)
1305                        } else {
1306                            make.name_ref(&param.name()?.text())
1307                        };
1308                        type_param_idx += 1;
1309                        make.type_arg(make.ty_path_from_segments([make.path_segment(name)], false))
1310                            .into()
1311                    }
1312                })
1313            });
1314
1315        make.path_from_segments(
1316            [make.generic_ty_path_segment(
1317                make.name_ref(&struct_name.text()),
1318                self_params_for_traits,
1319            )],
1320            false,
1321        )
1322    };
1323
1324    strukt.get_or_create_where_clause(&editor, new_predicates.into_iter());
1325    let edit = editor.finish();
1326    let strukt = ast::Struct::cast(edit.new_root().clone()).unwrap();
1327    let adt = ast::Adt::Struct(strukt.clone());
1328
1329    let mut span_map = span::SpanMap::empty();
1330    // One span for them all.
1331    span_map.push(adt.syntax().text_range().end(), span);
1332
1333    let self_for_traits = syntax_bridge::syntax_node_to_token_tree(
1334        self_for_traits.syntax(),
1335        &span_map,
1336        span,
1337        DocCommentDesugarMode::ProcMacro,
1338    );
1339    let info = match parse_adt_from_syntax(&adt, &span_map, span) {
1340        Ok(it) => it,
1341        Err(err) => {
1342            return ExpandResult::new(tt::TopSubtree::empty(tt::DelimSpan::from_single(span)), err);
1343        }
1344    };
1345
1346    let self_for_traits2 = self_for_traits.clone();
1347    let krate = dollar_crate(span);
1348    let krate2 = krate.clone();
1349    let dispatch_from_dyn = expand_simple_derive_with_parsed(
1350        span,
1351        info.clone(),
1352        quote! {span => #krate2::ops::DispatchFromDyn<#self_for_traits2> },
1353        |_adt| quote! {span => },
1354        false,
1355        quote! {span => __S },
1356    );
1357    let coerce_unsized = expand_simple_derive_with_parsed(
1358        span,
1359        info,
1360        quote! {span => #krate::ops::CoerceUnsized<#self_for_traits> },
1361        |_adt| quote! {span => },
1362        false,
1363        quote! {span => __S },
1364    );
1365    return ExpandResult::ok(quote! {span => #dispatch_from_dyn #coerce_unsized });
1366
1367    fn is_maybe_sized_bound(bound: ast::TypeBound) -> bool {
1368        if bound.question_mark_token().is_none() {
1369            return false;
1370        }
1371        let Some(ast::Type::PathType(ty)) = bound.ty() else {
1372            return false;
1373        };
1374        let Some(path) = ty.path() else {
1375            return false;
1376        };
1377        return segments_eq(&path, &["Sized"])
1378            || segments_eq(&path, &["core", "marker", "Sized"])
1379            || segments_eq(&path, &["std", "marker", "Sized"]);
1380
1381        fn segments_eq(path: &ast::Path, expected: &[&str]) -> bool {
1382            path.segments().zip_longest(expected.iter().copied()).all(|value| {
1383                value.both().is_some_and(|(segment, expected)| {
1384                    segment.name_ref().is_some_and(|name| name.text() == expected)
1385                })
1386            })
1387        }
1388    }
1389
1390    /// Returns true if any substitution was performed.
1391    fn substitute_type_bound(
1392        bound: ast::TypeBound,
1393        param_name: &str,
1394        replacement: &str,
1395    ) -> Option<ast::TypeBound> {
1396        let (editor, bound) = SyntaxEditor::with_ast_node(&bound);
1397        let substituted = bound
1398            .ty()
1399            .is_some_and(|ty| substitute_type_in_bound(&editor, ty, param_name, replacement));
1400        if !substituted {
1401            return None;
1402        }
1403
1404        let edit = editor.finish();
1405        Some(ast::TypeBound::cast(edit.new_root().clone()).unwrap())
1406    }
1407
1408    fn substitute_where_pred(
1409        predicate: &ast::WherePred,
1410        param_name: &str,
1411        replacement: &str,
1412    ) -> Option<ast::WherePred> {
1413        let (editor, predicate) = SyntaxEditor::with_ast_node(predicate);
1414        let substituted = predicate
1415            .ty()
1416            .is_some_and(|ty| substitute_type_in_bound(&editor, ty, param_name, replacement));
1417        if substituted && let Some(bounds) = predicate.type_bound_list() {
1418            for bound in bounds.bounds() {
1419                if let Some(ty) = bound.ty() {
1420                    substitute_type_in_bound(&editor, ty, param_name, replacement);
1421                }
1422            }
1423        }
1424        if !substituted {
1425            return None;
1426        }
1427
1428        let edit = editor.finish();
1429        Some(ast::WherePred::cast(edit.new_root().clone()).unwrap())
1430    }
1431
1432    fn substitute_type_in_bound(
1433        editor: &SyntaxEditor,
1434        ty: ast::Type,
1435        param_name: &str,
1436        replacement: &str,
1437    ) -> bool {
1438        let make = editor.make();
1439        return match ty {
1440            ast::Type::ArrayType(ty) => ty
1441                .ty()
1442                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1443            ast::Type::DynTraitType(ty) => {
1444                go_bounds(editor, ty.type_bound_list(), param_name, replacement)
1445            }
1446            ast::Type::FnPtrType(ty) => any_long(
1447                ty.param_list()
1448                    .into_iter()
1449                    .flat_map(|params| params.params().filter_map(|param| param.ty()))
1450                    .chain(ty.ret_type().and_then(|it| it.ty())),
1451                |ty| substitute_type_in_bound(editor, ty, param_name, replacement),
1452            ),
1453            ast::Type::ForType(ty) => ty
1454                .ty()
1455                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1456            ast::Type::ImplTraitType(ty) => {
1457                go_bounds(editor, ty.type_bound_list(), param_name, replacement)
1458            }
1459            ast::Type::ParenType(ty) => ty
1460                .ty()
1461                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1462            ast::Type::PathType(ty) => ty.path().is_some_and(|path| {
1463                if path.as_single_name_ref().is_some_and(|name| name.text() == param_name) {
1464                    editor.replace(
1465                        path.syntax(),
1466                        make.path_from_segments(
1467                            [make.path_segment(make.name_ref(replacement))],
1468                            false,
1469                        )
1470                        .syntax(),
1471                    );
1472                    return true;
1473                }
1474
1475                any_long(
1476                    path.segments()
1477                        .filter_map(|segment| segment.generic_arg_list())
1478                        .flat_map(|it| it.generic_args())
1479                        .filter_map(|generic_arg| match generic_arg {
1480                            ast::GenericArg::TypeArg(ty) => ty.ty(),
1481                            _ => None,
1482                        }),
1483                    |ty| substitute_type_in_bound(editor, ty, param_name, replacement),
1484                )
1485            }),
1486            ast::Type::PtrType(ty) => ty
1487                .ty()
1488                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1489            ast::Type::RefType(ty) => ty
1490                .ty()
1491                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1492            ast::Type::SliceType(ty) => ty
1493                .ty()
1494                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1495            ast::Type::TupleType(ty) => any_long(ty.fields(), |ty| {
1496                substitute_type_in_bound(editor, ty, param_name, replacement)
1497            }),
1498            ast::Type::PatternType(ty) => ty
1499                .ty()
1500                .is_some_and(|ty| substitute_type_in_bound(editor, ty, param_name, replacement)),
1501            ast::Type::InferType(_) | ast::Type::MacroType(_) | ast::Type::NeverType(_) => false,
1502        };
1503
1504        fn go_bounds(
1505            editor: &SyntaxEditor,
1506            bounds: Option<ast::TypeBoundList>,
1507            param_name: &str,
1508            replacement: &str,
1509        ) -> bool {
1510            bounds.is_some_and(|bounds| {
1511                any_long(bounds.bounds(), |bound| {
1512                    bound.ty().is_some_and(|ty| {
1513                        substitute_type_in_bound(editor, ty, param_name, replacement)
1514                    })
1515                })
1516            })
1517        }
1518
1519        /// Like [`Iterator::any()`], but not short-circuiting.
1520        fn any_long<I: Iterator, F: FnMut(I::Item) -> bool>(iter: I, mut f: F) -> bool {
1521            let mut result = false;
1522            iter.for_each(|item| result |= f(item));
1523            result
1524        }
1525    }
1526}