Skip to main content

codegen/
gen_aura.rs

1//! Emission of `aura_<name>` builder functions from manifest aura definitions.
2
3use anyhow::{Context, bail};
4use proc_macro2::TokenStream;
5use quote::{format_ident, quote};
6use wowlab_manifest_schema::{Manifest, ManifestAuraDef, SpellGroupRule};
7use wowlab_types::sim::FastMap;
8
9use crate::{
10    gen_aura_behavior::{apply_aura_consume_effects, apply_refresh_behavior, apply_tick_behavior},
11    gen_aura_modifiers::{
12        apply_per_stack_scalars, apply_rating_effects, apply_spell_scoped_effects,
13    },
14    gen_aura_periodic::{
15        apply_aura_duration, apply_periodic_damage, apply_periodic_hook_source,
16        apply_periodic_hooks, periodic_behavior_spell_id,
17    },
18    gen_emit::{
19        require_get, scalar_ref_call, spell_idx, validated_max_stacks_call,
20        validated_milliseconds_call,
21    },
22    helpers::to_snake,
23    rust_source::{ExpressionRequirements, expression_requirements, lit_int, result_expression},
24};
25
26pub(crate) fn aura_functions(
27    manifest: &Manifest,
28    module_name: &str,
29) -> anyhow::Result<TokenStream> {
30    let aura_count = manifest
31        .spell_groups
32        .iter()
33        .map(|group| group.auras.len())
34        .sum();
35    let mut spell_group_map = FastMap::default();
36
37    spell_group_map.reserve(aura_count);
38
39    for (group_index, group) in manifest.spell_groups.iter().enumerate() {
40        let group_index =
41            u8::try_from(group_index).context("manifest has more than 256 spell groups")?;
42
43        for aura_name in &group.auras {
44            spell_group_map.insert(aura_name.as_str(), (group_index, group.rule));
45        }
46    }
47
48    let mut aura_idx_map = FastMap::default();
49
50    aura_idx_map.reserve(manifest.auras.len());
51
52    for (index, name) in manifest.auras.keys().enumerate() {
53        let index = u8::try_from(index).context("manifest has more than 256 auras")?;
54
55        aura_idx_map.insert(name.as_str(), index);
56    }
57
58    let spell_id_map = manifest_spell_id_map(manifest)?;
59    let functions = manifest
60        .auras
61        .iter()
62        .map(|(name, def)| {
63            let options = AuraEmitOptions {
64                spell_group: spell_group_map.get(name.as_str()).copied(),
65                aura_idx_map: &aura_idx_map,
66                spell_id_map: &spell_id_map,
67                emit_default_max_stacks: def.max_stacks.is_none(),
68                has_item_budget: false,
69                hook_module: module_name,
70                hook_suffix: "_hook",
71            };
72
73            emit_aura_fn(name, def, &options)
74        })
75        .collect::<anyhow::Result<Vec<_>>>()?;
76
77    Ok(quote!(#(#functions)*))
78}
79
80fn manifest_spell_id_map(manifest: &Manifest) -> anyhow::Result<FastMap<&str, u32>> {
81    let mut ids: FastMap<&str, u32> = manifest
82        .spells
83        .iter()
84        .map(|(name, spell)| (name.as_str(), spell.id))
85        .collect();
86
87    for (name, id) in &manifest.reported_spells {
88        ids.entry(name).or_insert(*id);
89    }
90
91    let hero_spell_count = manifest
92        .hero_talents
93        .values()
94        .map(|tree| tree.spells.len())
95        .sum();
96    let mut hero_ids = FastMap::default();
97
98    hero_ids.reserve(hero_spell_count);
99
100    for tree in manifest.hero_talents.values() {
101        for (name, id) in &tree.spells {
102            if let Some(previous) = hero_ids.insert(name.as_str(), *id) {
103                if previous != *id && !ids.contains_key(name.as_str()) {
104                    bail!(
105                        "hero spell binding {name:?} is ambiguous between ids {previous} and {id}"
106                    );
107                }
108            }
109
110            ids.entry(name).or_insert(*id);
111        }
112    }
113
114    Ok(ids)
115}
116
117/// Controls item- and spec-specific aura emission.
118pub(crate) struct AuraEmitOptions<'a> {
119    pub spell_group: Option<(u8, SpellGroupRule)>,
120    pub aura_idx_map: &'a FastMap<&'a str, u8>,
121    pub spell_id_map: &'a FastMap<&'a str, u32>,
122    pub emit_default_max_stacks: bool,
123    pub has_item_budget: bool,
124    pub hook_module: &'a str,
125    pub hook_suffix: &'a str,
126}
127
128struct GeneratedAuraBuilderCode {
129    expression: TokenStream,
130    requirements: ExpressionRequirements,
131    _stage: AuraBuilderStage,
132}
133
134struct AuraBuilderStage;
135
136impl GeneratedAuraBuilderCode {
137    fn uses_data(&self) -> bool {
138        self.requirements.data
139    }
140
141    fn uses_budget(&self) -> bool {
142        self.requirements.budget
143    }
144
145    fn into_result_expression(self) -> anyhow::Result<TokenStream> {
146        Ok(result_expression(self.expression)?)
147    }
148
149    #[cfg(test)]
150    fn body_string(&self) -> String {
151        self.expression.to_string()
152    }
153}
154
155fn lower_manifest_aura(
156    def: &ManifestAuraDef,
157    options: &AuraEmitOptions<'_>,
158) -> anyhow::Result<GeneratedAuraBuilderCode> {
159    let id = def.id;
160    let spell_id = spell_idx(id);
161    let mut expression = match def.on.as_str() {
162        "target" => quote!(a.on_target()),
163        "player" => quote!(a.on_player()),
164        "pet" => quote!(a.on_pet()),
165        other => bail!("unsupported aura target {other:?} on aura {id}"),
166    };
167
168    expression = apply_core_aura_options(expression, def, options, &spell_id)?;
169    expression = apply_tick_behavior(expression, def)?;
170
171    if let Some(remove_stack) = &def.periodic_remove_stack {
172        let tick_ms = validated_milliseconds_call(
173            &remove_stack.tick_ms,
174            id,
175            "periodic_remove_stack.tick_ms",
176        )?;
177
178        expression = quote!(#expression.periodic_remove_stack(#tick_ms));
179    }
180
181    if let Some(value) = &def.primary_stat_percent_per_stack {
182        let value = scalar_ref_call(value)?;
183
184        expression = quote!(#expression.primary_pct_per_stack(#value)?);
185    }
186
187    if let Some(value) = &def.attack_power_percent_per_stack {
188        let value = scalar_ref_call(value)?;
189
190        expression = quote!(#expression.attack_power_pct_per_stack(#value)?);
191    }
192
193    expression = apply_spell_scoped_effects(expression, def, options)?;
194    expression = apply_periodic_damage(expression, def, id, options.has_item_budget)?;
195
196    if let Some(resource) = &def.periodic_resource {
197        let tick_ms =
198            validated_milliseconds_call(&resource.tick_ms, id, "periodic_resource.tick_ms")?;
199        let amount = scalar_ref_call(&resource.amount)?;
200
201        expression = quote!(#expression.periodic_resource(#tick_ms, #amount));
202    }
203
204    if let Some(drain) = &def.periodic_resource_drain {
205        let tick_ms =
206            validated_milliseconds_call(&drain.tick_ms, id, "periodic_resource_drain.tick_ms")?;
207        let amount = scalar_ref_call(&drain.amount)?;
208
209        expression = quote!(#expression.periodic_resource_drain(#tick_ms, #amount));
210    }
211
212    if let Some(apply) = &def.periodic_apply_aura {
213        let index = options
214            .aura_idx_map
215            .get(apply.aura.as_str())
216            .copied()
217            .ok_or_else(|| {
218                anyhow::anyhow!(
219                    "periodic_apply_aura.aura {:?} not found in [auras.*]",
220                    apply.aura
221                )
222            })?;
223        let tick_ms =
224            validated_milliseconds_call(&apply.tick_ms, id, "periodic_apply_aura.tick_ms")?;
225        let index = lit_int(index);
226
227        expression = quote!(#expression.periodic_apply_aura(#tick_ms, #index));
228    }
229
230    expression = apply_periodic_hook_source(expression, def)?;
231
232    if let Some(source_spell_id) = periodic_behavior_spell_id(def) {
233        let source_spell_id = lit_int(source_spell_id);
234
235        expression = quote!(#expression.apply_periodic_behavior_from_data(data, #source_spell_id)?);
236    }
237
238    expression = apply_periodic_hooks(expression, def, options)?;
239    expression = apply_aura_consume_effects(expression, def, options)?;
240
241    if let Some((group_index, rule)) = options.spell_group {
242        let group_index = lit_int(group_index);
243        let rule = match rule {
244            SpellGroupRule::Exclusive => {
245                quote!(wowlab_types::combat::SpellGroupRule::Exclusive)
246            }
247        };
248
249        expression = quote!(#expression.spell_group(#group_index, #rule));
250    }
251
252    let requirements = expression_requirements(expression.clone())?;
253
254    Ok(GeneratedAuraBuilderCode {
255        expression,
256        requirements,
257        _stage: AuraBuilderStage,
258    })
259}
260
261fn apply_core_aura_options(
262    mut expression: TokenStream,
263    def: &ManifestAuraDef,
264    options: &AuraEmitOptions<'_>,
265    spell_id: &TokenStream,
266) -> anyhow::Result<TokenStream> {
267    let id = def.id;
268
269    expression = apply_aura_duration(expression, def)?;
270
271    if options.emit_default_max_stacks {
272        let max_stacks = require_get("aura_max_stacks", spell_id);
273
274        expression = quote!(#expression.max_stacks(#max_stacks));
275    }
276
277    expression = apply_refresh_behavior(expression, def)?;
278
279    if def.snapshot {
280        expression = quote!(#expression.snapshot());
281    }
282
283    if let Some(value) = &def.forwarded_effect_value {
284        let value = scalar_ref_call(value)?;
285
286        expression = quote!(#expression.forwarded_effect_value(data, #value)?);
287    }
288
289    if let Some(multiplier) = &def.damage_mult {
290        let multiplier = scalar_ref_call(multiplier)?;
291
292        expression = quote!(#expression.damage_multiplier(#multiplier)?);
293    }
294
295    if let Some(multiplier) = &def.damage_mult_stacking {
296        let initial = scalar_ref_call(&multiplier.initial)?;
297        let per_stack = scalar_ref_call(&multiplier.per_stack)?;
298
299        expression = quote!(#expression.damage_multiplier_stacking(#initial, #per_stack)?);
300    }
301
302    if let Some(multiplier) = &def.pet_damage_mult {
303        let multiplier = scalar_ref_call(multiplier)?;
304
305        expression = quote!(#expression.pet_damage_multiplier(#multiplier)?);
306    }
307
308    if let Some(stacks) = &def.max_stacks {
309        let stacks = validated_max_stacks_call(stacks, id, "aura.max_stacks")?;
310
311        expression = quote!(#expression.max_stacks(#stacks));
312    }
313
314    expression = apply_per_stack_scalars(expression, def)?;
315    expression = apply_rating_effects(expression, def, options)?;
316
317    if def.apply_at_max_stacks {
318        expression = quote!(#expression.apply_at_max_stacks());
319    }
320
321    if def.reverse {
322        expression = quote!(#expression.reverse());
323    }
324
325    if def.freeze_stacks {
326        expression = quote!(#expression.freeze_stacks());
327    }
328
329    if def.async_stacks {
330        expression = quote!(#expression.async_stacks());
331    }
332
333    if def.periodic_damage_scales_with_stacks {
334        expression = quote!(#expression.periodic_damage_scales_with_stacks());
335    }
336
337    if def.reapply_on_expire {
338        expression = quote!(#expression.reapply_on_expire());
339    }
340
341    Ok(expression)
342}
343
344/// Emit a single `aura_<name>` builder function.
345pub(crate) fn emit_aura_fn(
346    name: &str,
347    def: &ManifestAuraDef,
348    options: &AuraEmitOptions<'_>,
349) -> anyhow::Result<TokenStream> {
350    let function_name = format_ident!("aura_{}", to_snake(name));
351    let generated = lower_manifest_aura(def, options)?;
352    let data_parameter = if generated.uses_data() {
353        quote!(data)
354    } else {
355        quote!(_data)
356    };
357    let uses_budget = generated.uses_budget();
358    let expression = generated.into_result_expression()?;
359
360    Ok(if options.has_item_budget {
361        let budget_parameter = if uses_budget {
362            quote!(budget)
363        } else {
364            quote!(_budget)
365        };
366
367        quote! {
368            fn #function_name(
369                a: AuraDefinitionDraft,
370                #data_parameter: &ResolvedGameData,
371                #budget_parameter: ItemBudget,
372            ) -> Result<AuraDefinitionDraft, wowlab_engine_combat::BuilderError> {
373                #expression
374            }
375        }
376    } else {
377        quote! {
378            fn #function_name(
379                a: AuraDefinitionDraft,
380                #data_parameter: &ResolvedGameData,
381            ) -> Result<AuraDefinitionDraft, wowlab_engine_combat::BuilderError> {
382                #expression
383            }
384        }
385    })
386}
387
388#[cfg(test)]
389#[path = "gen_aura/periodic_tests.rs"]
390mod periodic_tests;
391
392#[cfg(test)]
393#[path = "gen_aura/tests.rs"]
394mod tests;