use std::collections::HashMap;
use std::ops::Not;
use cairo_lang_syntax::node::SyntaxNode;
use itertools::Itertools;
use lsp_types::{
CodeAction, CodeActionKind, CodeActionOrCommand, CodeActionParams, CodeActionResponse,
Diagnostic, NumberOrString, Range, TextEdit, Url, WorkspaceEdit,
};
use tracing::{debug, warn};
use crate::lang::analysis_context::AnalysisContext;
use crate::lang::db::{AnalysisDatabase, LsSyntaxGroup};
use crate::lang::lsp::{LsProtoGroup, ToCairo};
use crate::project::ConfigsRegistry;
mod add_missing_trait;
mod cairo_lint;
mod create_module_file;
mod expand_macro;
mod fill_struct_fields;
mod fill_trait_members;
mod make_variable_mutable;
mod missing_import;
mod rename_unused_variable;
mod suggest_similar_identifier;
mod suggest_similar_member;
mod suggest_similar_method;
pub fn code_actions(
params: CodeActionParams,
config_registry: &ConfigsRegistry,
db: &AnalysisDatabase,
) -> Option<CodeActionResponse> {
let mut actions = Vec::with_capacity(params.context.diagnostics.len());
actions.extend(
get_code_actions_for_diagnostics(db, config_registry, ¶ms)
.into_iter()
.map(CodeActionOrCommand::from),
);
let node = node_on_range_start(db, ¶ms.text_document.uri, ¶ms.range)?;
actions.extend(expand_macro::expand_macro(db, node).into_iter().map(CodeActionOrCommand::from));
Some(actions)
}
fn get_code_actions_for_diagnostics(
db: &AnalysisDatabase,
config_registry: &ConfigsRegistry,
params: &CodeActionParams,
) -> Vec<CodeAction> {
let uri = ¶ms.text_document.uri;
let resolved_diagnostics: Vec<_> = params
.context
.diagnostics
.iter()
.map(|diagnostic| (extract_code(diagnostic), diagnostic))
.dedup_by(|(code1, diagnostic1), (code2, diagnostic2)| {
code1 == code2 && diagnostic1.range == diagnostic2.range
})
.filter_map(|(code, diagnostic)| {
let node = node_on_range_start(db, uri, &diagnostic.range)?;
let ctx = AnalysisContext::from_node(db, node)?;
Some((code, diagnostic, ctx))
})
.collect();
let mut result = get_lint_code_actions(db, config_registry, &resolved_diagnostics);
result.extend(resolved_diagnostics.iter().flat_map(|(code, diagnostic, ctx)| {
match code {
Some("E2200") => vec![],
Some("E0001") => rename_unused_variable::rename_unused_variable(
db,
&ctx.node,
(*diagnostic).clone(),
uri.clone(),
)
.to_vec(),
Some("E0002") => {
let fixes = add_missing_trait::add_missing_trait(db, ctx, uri.clone());
if let Some(fixes) = fixes
&& fixes.is_empty().not()
{
fixes
} else {
suggest_similar_method::suggest_similar_method(db, ctx, uri.clone())
.unwrap_or_default()
}
}
Some("E0003") => fill_struct_fields::fill_struct_fields(db, ctx.node, params).to_vec(),
Some("E0004") => fill_trait_members::fill_trait_members(db, ctx, params).to_vec(),
Some("E0005") => {
create_module_file::create_module_file(db, ctx.node, uri.clone()).to_vec()
}
Some("E0006") => {
let fixes = missing_import::missing_import(db, ctx, uri.clone());
if let Some(fixes) = fixes
&& fixes.is_empty().not()
{
fixes
} else {
suggest_similar_identifier::suggest_similar_identifier(db, ctx, uri)
.unwrap_or_default()
}
}
Some("E0007") => suggest_similar_member::suggest_similar_member(db, ctx, uri.clone())
.unwrap_or_default(),
Some("E2083") => {
make_variable_mutable::make_variable_mutable(db, ctx.node, uri.clone()).to_vec()
}
Some("E2080") => {
make_variable_mutable::make_ref_variable_mutable(db, ctx.node, uri.clone()).to_vec()
}
Some(code) => {
debug!("no code actions for diagnostic code: {code}");
vec![]
}
None => Default::default(),
}
}));
let changes = result
.iter()
.filter(|action| action.is_preferred.is_some_and(|a| a))
.flat_map(|action| action.edit.as_ref())
.flat_map(|edit| edit.changes.as_ref())
.cloned()
.fold(HashMap::<Url, Vec<TextEdit>>::new(), |mut acc, changes| {
for (url, edits) in changes {
let entry = acc.entry(url).or_default();
entry.extend(edits);
}
acc
});
if !changes.is_empty() {
result.push(CodeAction {
title: "Fix All".to_string(),
kind: Some(CodeActionKind::SOURCE_FIX_ALL),
edit: Some(WorkspaceEdit { changes: Some(changes), ..Default::default() }),
..Default::default()
});
}
result
}
fn get_lint_code_actions(
db: &AnalysisDatabase,
config_registry: &ConfigsRegistry,
resolved_diagnostics: &[(Option<&str>, &Diagnostic, AnalysisContext<'_>)],
) -> Vec<CodeAction> {
resolved_diagnostics
.iter()
.filter(|(code, ..)| *code == Some("E2200"))
.into_group_map_by(|(.., ctx)| ctx.module_id)
.into_values()
.flat_map(|ctxs| {
let nodes: Vec<_> = ctxs.iter().map(|(.., ctx)| ctx.node).collect();
let (.., ctx) = ctxs[0];
let tool_metadata = cairo_lint::get_tool_metadata(db, ctx.node, config_registry);
cairo_lint::cairo_lint(db, ctx.module_id, &nodes, tool_metadata).unwrap_or_default()
})
.collect()
}
trait VecExt<T> {
fn to_vec(self) -> Vec<T>;
}
impl<T> VecExt<T> for Option<T> {
fn to_vec(self) -> Vec<T> {
self.map(|result| vec![result]).unwrap_or_default()
}
}
fn extract_code(diagnostic: &Diagnostic) -> Option<&str> {
match &diagnostic.code {
Some(NumberOrString::String(code)) => Some(code),
Some(NumberOrString::Number(code)) => {
warn!("diagnostic code is not a string: `{code}`");
None
}
None => None,
}
}
fn node_on_range_start<'db>(
db: &'db AnalysisDatabase,
uri: &Url,
range: &Range,
) -> Option<SyntaxNode<'db>> {
let file_id = db.file_for_url(uri)?;
db.find_syntax_node_at_position(file_id, range.start.to_cairo())
}