Skip to main content

rpfm_extensions/lua/
check.rs

1//---------------------------------------------------------------------------//
2// Copyright (c) 2017-2026 Ismael Gutiérrez González. All rights reserved.
3//
4// This file is part of the Rusted PackFile Manager (RPFM) project,
5// which can be found here: https://github.com/Frodo45127/rpfm.
6//
7// This file is licensed under the MIT license, which can be found here:
8// https://github.com/Frodo45127/rpfm/blob/master/LICENSE.
9//---------------------------------------------------------------------------//
10
11//! Static checks of Lua scripts against the scripting API.
12//!
13//! Types are followed along call chains starting from known globals (`cm`, `core`, tables of functions like
14//! `common`), from the context of event listeners registered with `core:add_listener`, and from locals
15//! assigned from such chains. When a type can't be resolved, the chain is skipped instead of guessed, so
16//! unresolved code never gets reported.
17
18use full_moon::ast::{Ast, Call, Expression, FunctionArgs, FunctionBody, FunctionCall, Index, LocalAssignment, Assignment, FunctionDeclaration, Parameter, Prefix, Suffix, Var};
19use full_moon::node::Node;
20use full_moon::tokenizer::{Position, StringLiteralQuoteType, Symbol, TokenReference, TokenType};
21use full_moon::visitors::Visitor;
22use full_moon::LuaVersion;
23use getset::Getters;
24
25use std::collections::{HashMap, HashSet};
26
27use super::{LuaApi, LuaCallStyle, LuaFunction, LuaHoverTarget, LuaType};
28
29/// Globals holding objects whose methods come from several documented owners.
30const GLOBAL_OWNERS: [(&str, &[&str]); 2] = [
31    ("cm", &["cm", "campaign_manager"]),
32    ("core", &["core"]),
33];
34
35/// Prefix of the events scripts trigger themselves. These are not part of the game's event list.
36const SCRIPT_EVENT_PREFIX: &str = "ScriptEvent";
37
38/// Methods used to trigger custom events.
39const TRIGGER_EVENT_METHODS: [&str; 2] = ["trigger_event", "trigger_custom_event"];
40
41/// Range in a script, as `((start line, start column), (end line, end column))`. All values are 0-based.
42pub type LuaRange = ((u64, u64), (u64, u64));
43
44//---------------------------------------------------------------------------//
45//                              Enum & Structs
46//---------------------------------------------------------------------------//
47
48/// A problem found in a script.
49#[derive(Clone, Debug, PartialEq, Getters)]
50#[getset(get = "pub")]
51pub struct LuaIssue {
52
53    /// What the problem is.
54    kind: LuaIssueKind,
55
56    /// Where the problem is.
57    range: LuaRange,
58}
59
60/// Kinds of problems found in scripts.
61#[derive(Clone, Debug, PartialEq)]
62pub enum LuaIssueKind {
63
64    /// The script can't be parsed. Contains the parser's message.
65    SyntaxError(String),
66
67    /// A method not documented for the type it's called on. Contains the type and the method.
68    UnknownMethod(String, String),
69
70    /// A documented function called with the wrong amount of arguments. Contains the function,
71    /// the minimum and maximum (if any) amount of arguments it takes, and the amount passed.
72    WrongArgumentCount(String, usize, Option<usize>, usize),
73
74    /// A listener for an event that doesn't exist. Contains the event.
75    UnknownEvent(String),
76}
77
78/// A string literal passed to a parameter that expects a key from a DB table.
79#[derive(Clone, Debug, PartialEq, Getters)]
80#[getset(get = "pub")]
81pub struct LuaKeyReference {
82
83    /// DB table, with the `_tables` suffix.
84    table: String,
85
86    /// The key used in the script.
87    key: String,
88
89    /// Where the key is.
90    range: LuaRange,
91}
92
93/// A range of a script the docs describe.
94#[derive(Clone, Debug, PartialEq, Getters)]
95#[getset(get = "pub")]
96pub struct LuaHover {
97
98    /// What the docs describe. Use [`LuaApi::hover_html`] to get its docs.
99    target: LuaHoverTarget,
100
101    /// Where it is.
102    range: LuaRange,
103}
104
105/// Results of checking a script.
106#[derive(Clone, Debug, Default, PartialEq, Getters)]
107#[getset(get = "pub")]
108pub struct LuaCheckResult {
109
110    /// Problems found in the script.
111    issues: Vec<LuaIssue>,
112
113    /// DB keys used by the script, to be validated against the DB.
114    key_references: Vec<LuaKeyReference>,
115
116    /// Documented functions, accessors and events used by the script.
117    hovers: Vec<LuaHover>,
118}
119
120/// Things scripts define themselves, which must not be reported as unknown.
121#[derive(Clone, Debug, Default, PartialEq, Getters)]
122#[getset(get = "pub")]
123pub struct LuaDefinitions {
124
125    /// Members defined on globals, as `(owner, name)`, like `function cm:my_function()` or `cm.my_value = ...`.
126    members: HashSet<(String, String)>,
127
128    /// Custom events triggered by the scripts.
129    events: HashSet<String>,
130}
131
132/// Type the checker has resolved for a value.
133#[derive(Clone, Debug, PartialEq)]
134enum ResolvedType {
135
136    /// An object whose methods are documented under these owners, and whose own name is the first of them.
137    Owners(Vec<String>),
138
139    /// The context received by listeners of an event.
140    EventContext(String),
141}
142
143/// Visitor doing the checks.
144struct Checker<'a> {
145    api: &'a LuaApi,
146    definitions: &'a LuaDefinitions,
147
148    /// Known types of variables, one map per function. `None` means the variable exists but its type is unknown.
149    scopes: Vec<HashMap<String, Option<ResolvedType>>>,
150
151    /// Types of the first parameter of listener functions, by the start byte of their body.
152    listener_contexts: HashMap<usize, ResolvedType>,
153
154    result: LuaCheckResult,
155}
156
157/// Visitor collecting the definitions of a script.
158struct DefinitionsCollector<'a> {
159    definitions: &'a mut LuaDefinitions,
160}
161
162/// An argument of a call.
163#[derive(Clone, Copy)]
164enum Argument<'a> {
165    Expression(&'a Expression),
166
167    /// String passed without parentheses, like `f "text"`.
168    String(&'a TokenReference),
169
170    /// Table passed without parentheses, like `f { ... }`.
171    Table,
172}
173
174//---------------------------------------------------------------------------//
175//                             Implementations
176//---------------------------------------------------------------------------//
177
178impl LuaDefinitions {
179
180    /// This function adds the definitions of a script.
181    ///
182    /// Scripts that can't be parsed are ignored.
183    ///
184    /// # Arguments
185    ///
186    /// * `source` - Code of the script.
187    pub fn add_script(&mut self, source: &str) {
188        if let Ok(ast) = full_moon::parse_fallible(source, LuaVersion::lua51()).into_result() {
189            DefinitionsCollector { definitions: self }.visit_ast(&ast);
190        }
191    }
192
193    /// This function adds the definitions of another set to this one.
194    ///
195    /// # Arguments
196    ///
197    /// * `other` - Definitions to add.
198    pub fn extend(&mut self, other: Self) {
199        self.members.extend(other.members);
200        self.events.extend(other.events);
201    }
202}
203
204/// This function checks a script.
205///
206/// # Arguments
207///
208/// * `source` - Code of the script.
209/// * `api` - API of the game. Without it, only the syntax is checked.
210/// * `definitions` - Definitions from all the scripts the checked one can see, including itself.
211///
212/// # Returns
213///
214/// The problems and DB keys found in the script.
215pub fn check_script(source: &str, api: Option<&LuaApi>, definitions: &LuaDefinitions) -> LuaCheckResult {
216    let ast = match full_moon::parse_fallible(source, LuaVersion::lua51()).into_result() {
217        Ok(ast) => ast,
218        Err(errors) => {
219            let issues = errors.iter()
220                .map(|error| LuaIssue {
221                    kind: LuaIssueKind::SyntaxError(error.error_message().to_string()),
222                    range: (position(error.range().0), position(error.range().1)),
223                })
224                .collect();
225
226            return LuaCheckResult { issues, ..Default::default() };
227        }
228    };
229
230    match api {
231        Some(api) => check_ast(&ast, api, definitions),
232        None => LuaCheckResult::default(),
233    }
234}
235
236/// This function runs the API checks over a parsed script.
237///
238/// # Arguments
239///
240/// * `ast` - The parsed script.
241/// * `api` - API of the game.
242/// * `definitions` - Definitions from all the scripts the checked one can see.
243///
244/// # Returns
245///
246/// The problems and DB keys found in the script.
247fn check_ast(ast: &Ast, api: &LuaApi, definitions: &LuaDefinitions) -> LuaCheckResult {
248    let mut checker = Checker {
249        api,
250        definitions,
251        scopes: vec![HashMap::new()],
252        listener_contexts: HashMap::new(),
253        result: LuaCheckResult::default(),
254    };
255
256    checker.visit_ast(ast);
257    checker.result
258}
259
260impl Checker<'_> {
261
262    /// This function returns the type of a variable, looking from the innermost function outwards, then at the globals.
263    fn variable_type(&self, name: &str) -> Option<ResolvedType> {
264        if let Some(scope) = self.scopes.iter().rev().find(|scope| scope.contains_key(name)) {
265            return scope.get(name).cloned().flatten();
266        }
267
268        if let Some((_, owners)) = GLOBAL_OWNERS.iter().find(|(global, _)| *global == name) {
269            return Some(ResolvedType::Owners(owners.iter().map(|owner| owner.to_string()).collect()));
270        }
271
272        // Owners with only `.` functions are global tables of functions, like `common`. Other owners in the
273        // docs are named after example variables (`unit`, `army`, ...), which don't exist as globals.
274        self.api.owners().get(name)
275            .filter(|functions| !functions.is_empty() && functions.values().all(|function| *function.call_style() == LuaCallStyle::Field))
276            .map(|_| ResolvedType::Owners(vec![name.to_owned()]))
277    }
278
279    /// This function turns a documented type into a type the checker can follow.
280    fn resolve(&self, lua_type: &LuaType) -> Option<ResolvedType> {
281        let owner = match lua_type {
282            LuaType::Interface(interface) => interface.to_owned(),
283            LuaType::Object(object) => {
284                let interface = format!("{}_SCRIPT_INTERFACE", object.to_uppercase());
285                if self.api.owners().contains_key(&interface) {
286                    interface
287                } else {
288                    object.to_owned()
289                }
290            }
291            _ => return None,
292        };
293
294        if self.api.owners().contains_key(&owner) {
295            Some(ResolvedType::Owners(vec![owner]))
296        } else {
297            None
298        }
299    }
300
301    /// This function follows a call chain, reporting problems found on it if `report` is true.
302    ///
303    /// # Returns
304    ///
305    /// The type of the value the chain evaluates to, if it can be resolved.
306    fn follow_chain<'s>(&mut self, prefix: &Prefix, suffixes: impl Iterator<Item = &'s Suffix>, report: bool) -> Option<ResolvedType> {
307        let Prefix::Name(root) = prefix else {
308            return None;
309        };
310
311        let root_name = root.token().to_string();
312        let suffixes = suffixes.collect::<Vec<_>>();
313
314        let mut current = match self.variable_type(&root_name) {
315            Some(current) => current,
316
317            // Calls to global functions only get their arguments checked.
318            None => {
319                let is_local = self.scopes.iter().any(|scope| scope.contains_key(&root_name));
320                let global = if is_local { None } else { self.api.function("", &root_name) };
321                return match (global, suffixes.first()) {
322                    (Some(function), Some(Suffix::Call(Call::AnonymousCall(args)))) => {
323                        if report {
324                            self.push_hover(LuaHoverTarget::Function(String::new(), root_name.to_owned()), root);
325                            if suffixes.len() == 1 {
326                                self.check_call(function, &root_name, args, root);
327                            }
328                        }
329                        None
330                    }
331                    _ => None,
332                };
333            }
334        };
335
336        let mut index = 0;
337        while index < suffixes.len() {
338            let (name_token, args, via_dot) = match (suffixes[index], suffixes.get(index + 1)) {
339                (Suffix::Call(Call::MethodCall(method_call)), _) => (method_call.name(), method_call.args(), false),
340                (Suffix::Index(Index::Dot { name, .. }), Some(Suffix::Call(Call::AnonymousCall(args)))) => {
341                    index += 1;
342                    (name, args, true)
343                },
344                _ => return None,
345            };
346
347            let name = name_token.token().to_string();
348            current = match current {
349                ResolvedType::EventContext(event) => {
350                    let accessors = self.api.events().get(&event)?;
351                    match accessors.get(&name) {
352                        Some(accessor) => {
353                            if report {
354                                self.push_hover(LuaHoverTarget::Accessor(event.to_owned(), name.to_owned()), name_token);
355                            }
356                            self.resolve(accessor.lua_type())?
357                        }
358                        None => {
359                            if report && !via_dot && !accessors.is_empty() {
360                                self.push_issue(LuaIssueKind::UnknownMethod(format!("{event} context"), name), name_token);
361                            }
362                            return None;
363                        }
364                    }
365                }
366
367                ResolvedType::Owners(owners) => {
368                    let function = owners.iter().find_map(|owner| self.api.function(owner, &name).map(|function| (owner, function)));
369                    match function {
370                        Some((owner, function)) => {
371                            if report {
372                                self.push_hover(LuaHoverTarget::Function(owner.to_owned(), name.to_owned()), name_token);
373                            }
374
375                            // Calls using the other syntax pass `self` differently, so their arguments can't be counted.
376                            let expected_via_dot = *function.call_style() == LuaCallStyle::Field;
377                            if report && via_dot == expected_via_dot {
378                                self.check_call(function, &name, args, name_token);
379                            }
380
381                            self.resolve(function.returns().first()?.lua_type())?
382                        }
383                        None => {
384
385                            // Without docs for the object there is nothing to compare against.
386                            if !owners.iter().any(|owner| self.api.owners().contains_key(owner)) {
387                                return None;
388                            }
389
390                            let defined = owners.iter()
391                                .chain(std::iter::once(&root_name))
392                                .any(|owner| {
393                                    let member = (owner.to_owned(), name.to_owned());
394                                    self.definitions.members.contains(&member) || self.api.script_definitions().members.contains(&member)
395                                });
396
397                            if report && !via_dot && !defined {
398                                self.push_issue(LuaIssueKind::UnknownMethod(owners[0].to_owned(), name), name_token);
399                            }
400                            return None;
401                        }
402                    }
403                }
404            };
405
406            index += 1;
407        }
408
409        Some(current)
410    }
411
412    /// This function checks the arguments passed to a documented function, and collects the DB keys passed to it.
413    fn check_call(&mut self, function: &LuaFunction, name: &str, args: &FunctionArgs, name_token: &TokenReference) {
414        let Some(parameters) = function.parameters() else {
415            return;
416        };
417
418        let arguments = call_arguments(args);
419        let minimum = parameters.iter().rposition(|parameter| !parameter.optional() && !parameter.variadic()).map_or(0, |position| position + 1);
420        let maximum = if parameters.iter().any(|parameter| *parameter.variadic()) { None } else { Some(parameters.len()) };
421
422        // A call or `...` as last argument can expand to any amount of values.
423        let open_ended = match arguments.last() {
424            Some(Argument::Expression(Expression::FunctionCall(_))) => true,
425            Some(Argument::Expression(Expression::Symbol(symbol))) => matches!(symbol.token_type(), TokenType::Symbol { symbol: Symbol::Ellipsis }),
426            _ => false,
427        };
428
429        let too_few = !open_ended && arguments.len() < minimum;
430        let too_many = maximum.is_some_and(|maximum| arguments.len() > maximum);
431        if too_few || too_many {
432            self.push_issue(LuaIssueKind::WrongArgumentCount(name.to_owned(), minimum, maximum, arguments.len()), name_token);
433        }
434
435        for (parameter, argument) in parameters.iter().zip(arguments.iter()) {
436            if *parameter.variadic() {
437                break;
438            }
439
440            // Empty strings are how many functions take "no key", so they're never a reference.
441            if let (Some(table), Some((key, token))) = (parameter.db_table(), string_literal(argument)) {
442                if key.is_empty() {
443                    continue;
444                }
445
446                if let Some(range) = string_literal_range(token) {
447                    self.result.key_references.push(LuaKeyReference { table: table.to_owned(), key, range });
448                }
449            }
450        }
451    }
452
453    /// This function checks the event of a `core:add_listener` call, and types the context received by its functions.
454    fn check_listener(&mut self, call: &FunctionCall) {
455        let Prefix::Name(root) = call.prefix() else {
456            return;
457        };
458
459        let suffixes = call.suffixes().collect::<Vec<_>>();
460        let [Suffix::Call(Call::MethodCall(method_call))] = suffixes.as_slice() else {
461            return;
462        };
463
464        if root.token().to_string() != "core" || method_call.name().token().to_string() != "add_listener" || self.variable_type("core").is_none() {
465            return;
466        }
467
468        let arguments = call_arguments(method_call.args());
469        let Some((event, event_token)) = arguments.get(1).and_then(string_literal) else {
470            return;
471        };
472
473        if self.api.events().contains_key(&event) {
474            if let Some(range) = string_literal_range(event_token) {
475                self.result.hovers.push(LuaHover { target: LuaHoverTarget::Event(event.to_owned()), range });
476            }
477        } else {
478            let defined = self.definitions.events.contains(&event) || self.api.script_definitions().events.contains(&event);
479            if !event.starts_with(SCRIPT_EVENT_PREFIX) && !defined {
480                self.push_issue(LuaIssueKind::UnknownEvent(event), event_token);
481            }
482            return;
483        }
484
485        for argument in arguments.iter().skip(2).take(2) {
486            if let Argument::Expression(Expression::Function(function)) = argument {
487                if let Some(start) = Node::start_position(function.body()) {
488                    self.listener_contexts.insert(start.bytes(), ResolvedType::EventContext(event.to_owned()));
489                }
490            }
491        }
492    }
493
494    fn push_hover(&mut self, target: LuaHoverTarget, node: &impl Node) {
495        if let Some(range) = node_range(node) {
496            self.result.hovers.push(LuaHover { target, range });
497        }
498    }
499
500    fn push_issue(&mut self, kind: LuaIssueKind, node: &impl Node) {
501        if let Some(range) = node_range(node) {
502            self.result.issues.push(LuaIssue { kind, range });
503        }
504    }
505}
506
507impl Visitor for Checker<'_> {
508    fn visit_function_body(&mut self, body: &FunctionBody) {
509        let mut scope = HashMap::new();
510        let context = Node::start_position(body).and_then(|start| self.listener_contexts.remove(&start.bytes()));
511        for (index, parameter) in body.parameters().iter().enumerate() {
512            if let Parameter::Name(name) = parameter {
513                let parameter_type = if index == 0 { context.clone() } else { None };
514                scope.insert(name.token().to_string(), parameter_type);
515            }
516        }
517
518        self.scopes.push(scope);
519    }
520
521    fn visit_function_body_end(&mut self, _body: &FunctionBody) {
522        self.scopes.pop();
523    }
524
525    fn visit_function_call(&mut self, call: &FunctionCall) {
526        self.check_listener(call);
527        self.follow_chain(call.prefix(), call.suffixes(), true);
528    }
529
530    fn visit_local_assignment_end(&mut self, assignment: &LocalAssignment) {
531        let values = assignment.expressions().iter().collect::<Vec<_>>();
532        for (index, name) in assignment.names().iter().enumerate() {
533            let value_type = match values.get(index) {
534                Some(Expression::FunctionCall(call)) => self.follow_chain(call.prefix(), call.suffixes(), false),
535                Some(Expression::Var(Var::Name(other))) => self.variable_type(&other.token().to_string()),
536                _ => None,
537            };
538
539            if let Some(scope) = self.scopes.last_mut() {
540                scope.insert(name.token().to_string(), value_type);
541            }
542        }
543    }
544
545    fn visit_assignment(&mut self, assignment: &Assignment) {
546        for variable in assignment.variables() {
547            if let Var::Name(name) = variable {
548                let name = name.token().to_string();
549                if let Some(scope) = self.scopes.iter_mut().rev().find(|scope| scope.contains_key(&name)) {
550                    scope.insert(name, None);
551                }
552            }
553        }
554    }
555}
556
557impl Visitor for DefinitionsCollector<'_> {
558    fn visit_function_declaration(&mut self, declaration: &FunctionDeclaration) {
559        let names = declaration.name().names().iter().map(|name| name.token().to_string()).collect::<Vec<_>>();
560        let member = match declaration.name().method_name() {
561            Some(method) => names.last().map(|owner| (owner.to_owned(), method.token().to_string())),
562            None if names.len() >= 2 => Some((names[names.len() - 2].to_owned(), names[names.len() - 1].to_owned())),
563            None => None,
564        };
565
566        if let Some(member) = member {
567            self.definitions.members.insert(member);
568        }
569    }
570
571    fn visit_assignment(&mut self, assignment: &Assignment) {
572        for variable in assignment.variables() {
573            if let Var::Expression(expression) = variable {
574                let suffixes = expression.suffixes().collect::<Vec<_>>();
575                if let (Prefix::Name(owner), [Suffix::Index(Index::Dot { name, .. })]) = (expression.prefix(), suffixes.as_slice()) {
576                    self.definitions.members.insert((owner.token().to_string(), name.token().to_string()));
577                }
578            }
579        }
580    }
581
582    fn visit_method_call(&mut self, method_call: &full_moon::ast::MethodCall) {
583        if TRIGGER_EVENT_METHODS.contains(&method_call.name().token().to_string().as_str()) {
584            if let Some((event, _)) = call_arguments(method_call.args()).first().and_then(string_literal) {
585                self.definitions.events.insert(event);
586            }
587        }
588    }
589}
590
591//---------------------------------------------------------------------------//
592//                                 Helpers
593//---------------------------------------------------------------------------//
594
595/// This function returns the arguments of a call.
596fn call_arguments(args: &FunctionArgs) -> Vec<Argument<'_>> {
597    match args {
598        FunctionArgs::Parentheses { arguments, .. } => arguments.iter().map(Argument::Expression).collect(),
599        FunctionArgs::String(string) => vec![Argument::String(string)],
600        FunctionArgs::TableConstructor(_) => vec![Argument::Table],
601        _ => vec![],
602    }
603}
604
605/// This function returns the range of the text of a string literal, without its quotes.
606fn string_literal_range(token: &TokenReference) -> Option<LuaRange> {
607    let ((start_line, start_column), (end_line, end_column)) = node_range(token)?;
608    match token.token_type() {
609        TokenType::StringLiteral { quote_type: StringLiteralQuoteType::Double | StringLiteralQuoteType::Single, .. } =>
610            Some(((start_line, start_column + 1), (end_line, end_column.saturating_sub(1)))),
611        _ => Some(((start_line, start_column), (end_line, end_column))),
612    }
613}
614
615/// This function returns the value and token of an argument, if it's a string literal.
616fn string_literal<'a>(argument: &Argument<'a>) -> Option<(String, &'a TokenReference)> {
617    let token = match *argument {
618        Argument::Expression(Expression::String(token)) | Argument::String(token) => token,
619        _ => return None,
620    };
621
622    match token.token_type() {
623        TokenType::StringLiteral { literal, .. } => Some((literal.to_string(), token)),
624        _ => None,
625    }
626}
627
628/// This function returns the range of a node.
629fn node_range(node: &impl Node) -> Option<LuaRange> {
630    Some((position(Node::start_position(node)?), position(Node::end_position(node)?)))
631}
632
633/// This function turns a parser position (1-based) into a 0-based `(line, column)`.
634fn position(position: Position) -> (u64, u64) {
635    (position.line().saturating_sub(1) as u64, position.character().saturating_sub(1) as u64)
636}