Skip to main content

ktrs_compose/rules/
view_model_injection.rs

1//! Port of `rules/ViewModelInjection.kt`.
2
3use ktrs_ast::psi::{self, KtCallExpression, KtFunction, KtFunctionType, KtProperty, kt_psi_factory};
4use ktrs_ast::{Ast, NodeId};
5use ktrs_syntax::SyntaxKind::WHITE_SPACE;
6
7use crate::core::compose_kt_config::ComposeKtConfig;
8use crate::core::compose_kt_visitor::ComposeKtVisitor;
9use crate::core::emitter::Emitter;
10use crate::core::util::ast_nodes::{first_child_leaf_or_self, last_child_leaf_or_self, next_code_sibling};
11use crate::core::util::kt_functions::{defined_in_interface, is_override};
12use crate::core::util::psi_elements::{ChildrenByClass, find_direct_children_by_class, find_direct_first_child_by_class};
13
14pub struct ViewModelInjection;
15
16impl ComposeKtVisitor for ViewModelInjection {
17    fn visit_composable(&self, ast: &mut Ast, function: KtFunction, emitter: &mut dyn Emitter, config: &dyn ComposeKtConfig) {
18        if is_override(ast, function) || defined_in_interface(ast, function) {
19            return;
20        }
21        let Some(body_block) = function.body_block_expression(ast) else { return };
22        let mut known_view_model_factories: Vec<String> = DEFAULT_KNOWN_VIEW_MODEL_FACTORIES.iter().map(|s| (*s).to_owned()).collect();
23        for factory in config.get_set("viewModelFactories", &[]) {
24            if !known_view_model_factories.contains(&factory) {
25                known_view_model_factories.push(factory);
26            }
27        }
28
29        // Lazy upstream, and the fix deletes the property: step the walk.
30        let mut properties = ChildrenByClass::new(body_block.node());
31        while let Some(property) = properties.next::<KtProperty>(ast, |_, _| true) {
32            let matches: Vec<String> = find_direct_children_by_class::<KtCallExpression>(ast, property.node())
33                .into_iter()
34                .filter_map(|it| {
35                    let callee = ast.text(it.callee_expression(ast)?);
36                    known_view_model_factories.contains(&callee).then_some((it, callee))
37                })
38                .filter(|&(it, _)| !is_navigation(ast, it, body_block.node()))
39                .map(|(_, callee)| callee)
40                .collect();
41            for view_model_factory_name in matches {
42                emitter.report(ast, property.node(), &error_message(&view_model_factory_name), true).if_fix(|| {
43                    fix(ast, function, property, &view_model_factory_name);
44                });
45            }
46        }
47    }
48}
49
50fn fix(ast: &mut Ast, composable: KtFunction, property: KtProperty, view_model_factory_name: &str) {
51    let variable_name = property.name(ast).unwrap_or_else(|| "null".to_owned());
52    let Some(call_expression) = find_direct_first_child_by_class::<KtCallExpression>(ast, property.node()) else { return };
53    let Some(argument_list) = call_expression.value_argument_list(ast) else { return };
54    if !call_expression.value_arguments(ast).is_empty() {
55        return;
56    }
57    let view_model_type_reference = match property.type_reference(ast) {
58        Some(type_reference) => type_reference.node(),
59        None => match call_expression.type_arguments(ast).as_slice() {
60            [single] => *single,
61            _ => return,
62        },
63    };
64
65    let raw_view_model_type = ast.text(view_model_type_reference);
66    let raw_argument_list = ast.text(argument_list.node());
67    let value_parameters = composable.value_parameters(ast);
68    let last_parameters = &value_parameters[value_parameters.len().saturating_sub(2)..];
69    let Some(parameter_list) = composable.value_parameter_list(ast) else { return };
70
71    let new_code = format!("{variable_name}: {raw_view_model_type} = {view_model_factory_name}{raw_argument_list}");
72    let new_param = kt_psi_factory::create_parameter(ast, &new_code);
73    let new_param_text = new_param.text(ast);
74
75    let last_is_function_type = last_parameters
76        .last()
77        .and_then(|it| it.type_reference(ast))
78        .and_then(|it| it.type_element(ast))
79        .is_some_and(|it| KtFunctionType::is(ast, it.node()));
80    if last_parameters.is_empty() {
81        let last_token = leaf(ast, last_child_leaf_or_self(ast, parameter_list.node()));
82        ast.raw_replace_with_text(last_token, &format!("{new_param_text})"));
83    } else if last_is_function_type {
84        if last_parameters.len() == 1 {
85            let first_token = leaf(ast, first_child_leaf_or_self(ast, parameter_list.node()));
86            ast.raw_replace_with_text(first_token, &format!("({new_code}, "));
87        } else {
88            let comma = next_code_sibling(ast, last_parameters[0].node()).expect("NullPointerException: nextCodeSibling()!!");
89            let last_token = leaf(ast, last_child_leaf_or_self(ast, comma));
90            let text = format!("{} {new_code},", ast.text(last_token));
91            ast.raw_replace_with_text(last_token, &text);
92        }
93    } else {
94        let last_parameter = value_parameters.last().expect("NoSuchElementException: List is empty.");
95        let has_trailing_comma = next_code_sibling(ast, last_parameter.node()).is_some_and(|it| ast.text(it) == ",");
96        let pre_comma_if_needed = if has_trailing_comma { "" } else { "," };
97        let trailing_comma_if_needed = if has_trailing_comma { "," } else { "" };
98        let last_token = leaf(ast, last_child_leaf_or_self(ast, parameter_list.node()));
99        ast.raw_replace_with_text(last_token, &format!("{pre_comma_if_needed}{new_param_text}{trailing_comma_if_needed})"));
100    }
101
102    if let Some(previous) = ast.tree_prev(property.node()).filter(|&it| ast.element_type(it) == WHITE_SPACE) {
103        psi::delete(ast, previous);
104    }
105    psi::delete(ast, property.node());
106}
107
108/// `node as LeafPsiElement`.
109fn leaf(ast: &Ast, node: NodeId) -> NodeId {
110    assert!(ast.is_leaf_element(node), "ClassCastException: {:?} cannot be cast to LeafPsiElement", ast.element_type(node));
111    node
112}
113
114fn is_navigation(ast: &Ast, call: KtCallExpression, stop_at: NodeId) -> bool {
115    ast.parents(call.node())
116        .take_while(|&it| it != stop_at)
117        .filter_map(|it| KtCallExpression::cast(ast, it))
118        .any(|it| it.callee_expression(ast).is_some_and(|c| KNOWN_NAVIGATION_CALL_EXPRESSIONS.contains(&ast.text(c).as_str())))
119}
120
121const KNOWN_NAVIGATION_CALL_EXPRESSIONS: &[&str] = &["composable", "NavHost"];
122
123const DEFAULT_KNOWN_VIEW_MODEL_FACTORIES: &[&str] = &[
124    "viewModel",
125    "weaverViewModel",
126    "hiltViewModel",
127    "injectedViewModel",
128    "mavericksViewModel",
129    "tangleViewModel",
130    "metroViewModel",
131    "anvilViewModel",
132];
133
134pub fn error_message(factory_name: &str) -> String {
135    format!(
136        "Implicit dependencies of composables should be made explicit.\n\
137         Usages of {factory_name} to acquire a ViewModel should be done in composable default parameters, so that it is more testable and flexible.\n\
138         See https://mrmans0n.github.io/compose-rules/rules/#viewmodels for more information."
139    )
140}