Skip to main content

yazi_codegen/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3mod helper;
4use syn::{Data, DeriveInput, Fields, parse_macro_input};
5
6use crate::helper::{generics_with_de, has_serde_attr, ident_name, named_fields};
7
8#[proc_macro_derive(DeserializeOver)]
9pub fn deserialize_over(input: TokenStream) -> TokenStream {
10	let DeriveInput { ident, generics, .. } = parse_macro_input!(input as DeriveInput);
11	let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
12
13	quote! {
14		impl #impl_generics yazi_shim::toml::DeserializeOverHook for #ident #ty_generics #where_clause {}
15	}
16	.into()
17}
18
19#[proc_macro_derive(DeserializeOver1)]
20pub fn deserialize_over1(input: TokenStream) -> TokenStream {
21	let DeriveInput { ident, generics, data, .. } = parse_macro_input!(input as DeriveInput);
22	let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
23
24	let visitor_generics = generics_with_de(&generics);
25	let (impl_visitor_generics, ..) = visitor_generics.split_for_impl();
26
27	let (flatten_fields, normal_fields): (Vec<_>, Vec<_>) =
28		named_fields(data).into_iter().partition(|f| has_serde_attr(&f.attrs, "flatten"));
29
30	let field_hooks: Vec<_> = flatten_fields
31		.iter()
32		.chain(&normal_fields)
33		.map(|f| {
34			let ident = f.ident.as_ref().unwrap();
35			quote! { #ident: deserialized.#ident.deserialize_over_hook().map_err(Error::custom)? }
36		})
37		.collect();
38
39	let normal_arms = normal_fields.into_iter().map(|f| {
40		let ident = f.ident.unwrap();
41		let name = ident_name(&ident);
42		quote! { #name => self.0.#ident = map.next_value_seed(DeserializeOverSeed(self.0.#ident))? }
43	});
44
45	let flatten_arm = match flatten_fields.into_iter().next() {
46		Some(f) => {
47			let ident = f.ident.unwrap();
48			quote! { _ => self.0.#ident = self.0.#ident.deserialize_over_with(single_map_entry(&*key, &mut map))? }
49		}
50		None => quote! { _ => _ = map.next_value::<IgnoredAny>()? },
51	};
52
53	quote! {
54		impl #impl_generics yazi_shim::toml::DeserializeOverWith for #ident #ty_generics #where_clause {
55			fn deserialize_over_with<'__de, __D: serde::Deserializer<'__de>>(self, de: __D) -> Result<Self, __D::Error> {
56				use serde::de::{Error, IgnoredAny, MapAccess, Visitor};
57				use yazi_shared::KebabCasedKey;
58				use yazi_shim::{serde::single_map_entry, toml::{DeserializeOverHook, DeserializeOverSeed, DeserializeOverWith}};
59
60				struct V #impl_generics (#ident #ty_generics) #where_clause;
61
62				impl #impl_visitor_generics Visitor<'__de> for V #ty_generics #where_clause {
63					type Value = #ident #ty_generics;
64
65					fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
66						f.write_str("a map")
67					}
68
69					fn visit_map<__M: MapAccess<'__de>>(mut self, mut map: __M) -> Result<Self::Value, __M::Error> {
70						while let Some(key) = map.next_key::<KebabCasedKey>()? {
71							match key.as_ref() {
72								#(#normal_arms,)*
73								#flatten_arm
74							}
75						}
76						Ok(self.0)
77					}
78				}
79
80				let deserialized = de.deserialize_map(V(self))?;
81				Ok(Self { #(#field_hooks,)* })
82			}
83		}
84	}
85	.into()
86}
87
88#[proc_macro_derive(DeserializeOver2)]
89pub fn deserialize_over2(input: TokenStream) -> TokenStream {
90	let DeriveInput { ident, generics, data, .. } = parse_macro_input!(input as DeriveInput);
91	let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
92
93	let visitor_generics = generics_with_de(&generics);
94	let (impl_visitor_generics, ..) = visitor_generics.split_for_impl();
95
96	let mut normal_arms = vec![];
97	let mut flatten_arm = quote! { _ => _ = map.next_value::<IgnoredAny>()? };
98	for field in named_fields(data) {
99		let (field_ident, field_ty) = (field.ident, field.ty);
100		let field_name = ident_name(field_ident.as_ref().unwrap());
101
102		if has_serde_attr(&field.attrs, "skip") {
103			continue;
104		}
105
106		if has_serde_attr(&field.attrs, "flatten") {
107			flatten_arm = quote! { _ => self.0.#field_ident = self.0.#field_ident.deserialize_over_with(single_map_entry(&*key, &mut map))? };
108			continue;
109		}
110
111		let serde_attrs: Vec<_> = field.attrs.iter().filter(|a| a.path().is_ident("serde")).collect();
112		if serde_attrs.is_empty() {
113			normal_arms.push(quote! { #field_name => self.0.#field_ident = map.next_value()? });
114		} else {
115			normal_arms.push(quote! {
116				#field_name => {
117					#[derive(serde::Deserialize)]
118					struct H #impl_generics(#(#serde_attrs)* #field_ty,) #where_clause;
119					self.0.#field_ident = map.next_value::<H #ty_generics>()?.0;
120				}
121			});
122		}
123	}
124
125	quote! {
126		impl #impl_generics yazi_shim::toml::DeserializeOverWith for #ident #ty_generics #where_clause {
127			fn deserialize_over_with<'__de, __D: serde::Deserializer<'__de>>(self, de: __D) -> Result<Self, __D::Error> {
128				use serde::de::{Error, IgnoredAny, MapAccess, Visitor};
129				use std::borrow::Cow;
130				use yazi_shim::{serde::single_map_entry, toml::DeserializeOverWith};
131
132				struct V #impl_generics (#ident #ty_generics) #where_clause;
133
134				impl #impl_visitor_generics Visitor<'__de> for V #ty_generics #where_clause {
135					type Value = #ident #ty_generics;
136
137					fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
138						f.write_str("a map")
139					}
140
141					fn visit_map<__M: MapAccess<'__de>>(mut self, mut map: __M) -> Result<Self::Value, __M::Error> {
142						while let Some(key) = map.next_key::<Cow<str>>()? {
143							match key.as_ref() {
144								#(#normal_arms,)*
145								#flatten_arm
146							}
147						}
148
149						Ok(self.0)
150					}
151				}
152
153				de.deserialize_map(V(self))
154			}
155		}
156	}
157	.into()
158}
159
160#[proc_macro_derive(Overlay)]
161pub fn overlay(input: TokenStream) -> TokenStream {
162	let DeriveInput { ident, generics, data, .. } = parse_macro_input!(input as DeriveInput);
163	let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
164
165	let stmts: Vec<_> = match data {
166		Data::Struct(s) => match s.fields {
167			Fields::Named(fields) => fields
168				.named
169				.into_iter()
170				.map(|f| {
171					let field_ident = f.ident;
172					quote! { self.#field_ident.overlay(new.#field_ident); }
173				})
174				.collect(),
175			Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
176				vec![quote! { self.0.overlay(new.0); }]
177			}
178			_ => panic!("expected named fields or a single-field tuple struct"),
179		},
180		_ => panic!("expected struct"),
181	};
182
183	quote! {
184		impl #impl_generics yazi_shim::serde::Overlay for #ident #ty_generics #where_clause {
185			fn overlay(&self, new: Self) {
186				use yazi_shim::serde::Overlay;
187
188				#(#stmts)*
189			}
190		}
191	}
192	.into()
193}
194
195#[proc_macro_derive(FromLuaOwned)]
196pub fn from_lua(input: TokenStream) -> TokenStream {
197	let DeriveInput { ident, generics, .. } = parse_macro_input!(input as DeriveInput);
198
199	let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
200
201	quote! {
202		impl #impl_generics ::mlua::FromLua for #ident #ty_generics #where_clause {
203			#[inline]
204			fn from_lua(value: ::mlua::Value, lua: &::mlua::Lua) -> ::mlua::Result<Self> {
205				<::mlua::UserDataOwned<Self> as ::mlua::FromLua>::from_lua(value, lua).map(|ud| ud.0)
206			}
207		}
208	}
209	.into()
210}