use cargo_metadata::diagnostic::{Diagnostic, DiagnosticLevel};
pub fn is_cgp_diagnostic(diagnostic: &Diagnostic) -> bool {
let cgp_patterns = [
"CanUseComponent",
"IsProviderFor",
"HasField",
"cgp_impl",
"cgp_component",
"cgp_auto_getter",
"delegate_components",
"check_components",
];
if cgp_patterns.iter().any(|p| diagnostic.message.contains(p)) {
return true;
}
for child in &diagnostic.children {
if cgp_patterns.iter().any(|p| child.message.contains(p)) {
return true;
}
}
false
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ComponentInfo {
pub component_type: String,
pub provider_trait: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct FieldInfo {
pub field_name: String,
pub is_complete: bool,
pub has_unknown_chars: bool,
pub target_type: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ProviderRelationship {
pub provider_type: String,
pub component: String,
pub context: String,
}
pub fn extract_component_from_can_use(message: &str) -> Option<ComponentInfo> {
let start = message.find("CanUseComponent<")?;
let after_start = start + "CanUseComponent<".len();
let component_type = extract_balanced_generic(message, after_start)?;
let provider_trait = derive_provider_trait_name(&component_type);
Some(ComponentInfo {
component_type,
provider_trait,
})
}
pub fn extract_component_info(message: &str) -> Option<ComponentInfo> {
if let Some(info) = extract_component_from_can_use(message) {
return Some(info);
}
if let Some(start) = message.find("IsProviderFor<") {
let after_start = start + "IsProviderFor<".len();
if let Some(comma_pos) = message[after_start..].find(',') {
let component_type = message[after_start..after_start + comma_pos].trim();
if component_type.contains("Component") {
let provider_trait = derive_provider_trait_name(component_type);
return Some(ComponentInfo {
component_type: component_type.to_string(),
provider_trait,
});
}
}
}
for word in message.split_whitespace() {
let clean_word =
word.trim_matches(|c: char| !c.is_alphanumeric() && c != '<' && c != '>' && c != ',');
if clean_word.contains("IsProviderFor") {
continue;
}
if clean_word.contains("Component") {
if let Some(component_type) = extract_component_type_name(clean_word) {
let provider_trait = derive_provider_trait_name(&component_type);
return Some(ComponentInfo {
component_type,
provider_trait,
});
}
}
}
None
}
fn extract_component_type_name(text: &str) -> Option<String> {
if text.ends_with("Component") && !text.contains('<') {
return Some(text.to_string());
}
if let Some(component_pos) = text.rfind("Component") {
let before_component = &text[..component_pos + "Component".len()];
let mut depth = 0;
let mut start_idx = 0;
for (i, ch) in before_component.char_indices().rev() {
if ch == '>' {
depth += 1;
} else if ch == '<' {
depth -= 1;
} else if depth == 0 && !ch.is_alphanumeric() && ch != '_' {
start_idx = i + 1;
break;
}
}
return Some(before_component[start_idx..].to_string());
}
None
}
pub fn derive_provider_trait_name(component_name: &str) -> Option<String> {
if let Some(stripped) = component_name.strip_suffix("Component") {
if !stripped.is_empty() {
return Some(stripped.to_string());
}
}
if component_name.contains("Component") {
if let Some(pos) = component_name.rfind("Component") {
let before = &component_name[..pos];
if !before.is_empty() {
return Some(before.to_string());
}
}
}
None
}
pub fn extract_field_info(diagnostic: &Diagnostic) -> Option<FieldInfo> {
for child in &diagnostic.children {
if matches!(child.level, DiagnosticLevel::Help) {
let message = &child.message;
if message.contains("HasField") && message.contains("is not implemented for") {
let field_name_result = extract_field_name_from_symbol(message)?;
let target_type = extract_type_from_not_implemented(message)?;
return Some(FieldInfo {
field_name: field_name_result.0,
is_complete: field_name_result.1,
has_unknown_chars: field_name_result.2,
target_type,
});
}
}
}
None
}
fn extract_field_name_from_symbol(message: &str) -> Option<(String, bool, bool)> {
let relevant_part = if let Some(pos) = message.find("but trait") {
&message[..pos]
} else {
message
};
let expected_length = extract_symbol_length(relevant_part)?;
let (chars, has_unknown) = extract_chars_from_pattern(relevant_part);
if chars.is_empty() {
return None;
}
let field_name: String = chars.iter().collect();
let is_complete = field_name.len() == expected_length;
Some((field_name, is_complete, has_unknown))
}
fn extract_symbol_length(text: &str) -> Option<usize> {
let start = text.find("Symbol<")?;
let after_symbol = &text[start + 7..];
let comma_pos = after_symbol.find(',')?;
after_symbol[..comma_pos].trim().parse::<usize>().ok()
}
fn extract_chars_from_pattern(text: &str) -> (Vec<char>, bool) {
let mut chars = Vec::new();
let mut has_unknown = false;
let mut idx = 0;
while idx < text.len() {
if text[idx..].starts_with("Chars<'") {
let char_start = idx + 7;
if let Some(ch) = text[char_start..].chars().next() {
if ch != '\'' {
chars.push(ch);
}
}
} else if text[idx..].starts_with("Chars<_") {
chars.push('\u{FFFD}'); has_unknown = true;
}
idx += 1;
}
(chars, has_unknown)
}
fn extract_type_from_not_implemented(message: &str) -> Option<String> {
let start = message.find("is not implemented for `")?;
let after_start = start + "is not implemented for `".len();
let end = message[after_start..].find('`')?;
let full_name = &message[after_start..after_start + end];
let simple_name = full_name.split("::").last().unwrap_or(full_name);
Some(simple_name.to_string())
}
pub fn extract_provider_relationship(message: &str) -> Option<ProviderRelationship> {
if !message.contains("IsProviderFor") {
return None;
}
let provider_type = extract_type_from_for_to_implement(message)?;
let start = message.find("IsProviderFor<")?;
let after_start = start + "IsProviderFor<".len();
let comma_pos = find_comma_at_depth(after_start, message)?;
let component = message[after_start..comma_pos].trim().to_string();
let after_comma = comma_pos + 1;
let context = extract_balanced_generic(message, after_comma)?;
Some(ProviderRelationship {
provider_type,
component,
context,
})
}
fn extract_type_from_for_to_implement(message: &str) -> Option<String> {
let start = message.find("for `")?;
let after_start = start + 5;
let end = message[after_start..].find("` to")?;
let full_name = &message[after_start..after_start + end];
let simple_name = full_name.split("::").last().unwrap_or(full_name);
Some(simple_name.to_string())
}
fn find_comma_at_depth(start_pos: usize, text: &str) -> Option<usize> {
let mut depth = 0;
for (i, ch) in text[start_pos..].char_indices() {
match ch {
'<' => depth += 1,
'>' => depth -= 1,
',' if depth == 0 => return Some(start_pos + i),
_ => {}
}
}
None
}
fn extract_balanced_generic(text: &str, start_pos: usize) -> Option<String> {
let mut depth = 1; let mut end_pos = start_pos;
for (i, ch) in text[start_pos..].char_indices() {
match ch {
'<' => depth += 1,
'>' => {
depth -= 1;
if depth == 0 {
end_pos = start_pos + i;
break;
}
}
_ => {}
}
}
if depth == 0 {
Some(text[start_pos..end_pos].trim().to_string())
} else {
Some(text[start_pos..].trim_end_matches('>').trim().to_string())
}
}
pub fn extract_check_trait(message: &str) -> Option<String> {
let start = message.find("required by a bound in `")?;
let after_start = start + "required by a bound in `".len();
let end = message[after_start..].find('`')?;
Some(message[after_start..after_start + end].to_string())
}
pub fn has_other_hasfield_implementations(diagnostic: &Diagnostic) -> bool {
for child in &diagnostic.children {
if matches!(child.level, DiagnosticLevel::Help) {
if child.message.contains("but trait `HasField")
|| child
.message
.contains("the following other types implement trait")
{
return true;
}
}
}
false
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ConsumerTraitDependency {
pub trait_name: String,
pub context_type: String,
pub component_name: Option<String>,
}
pub fn extract_consumer_trait_dependency(note: &str) -> Option<ConsumerTraitDependency> {
if let Some(for_pos) = note.find("required for `") {
let after_for = for_pos + "required for `".len();
if let Some(context_end) = note[after_for..].find('`') {
let context_type = ¬e[after_for..after_for + context_end];
if let Some(implement_pos) = note[after_for + context_end..].find("to implement `") {
let trait_start = after_for + context_end + implement_pos + "to implement `".len();
if let Some(trait_end) = note[trait_start..].find('`') {
let trait_name = ¬e[trait_start..trait_start + trait_end];
let cleaned_trait = strip_module_prefixes(trait_name);
if cleaned_trait.starts_with("Can")
&& !cleaned_trait.contains("CanUseComponent")
&& !cleaned_trait.starts_with("IsProviderFor")
{
let component_name = derive_component_from_consumer_trait(&cleaned_trait);
return Some(ConsumerTraitDependency {
trait_name: cleaned_trait,
context_type: strip_module_prefixes(context_type),
component_name,
});
}
}
}
}
}
None
}
pub fn derive_component_from_consumer_trait(consumer_trait: &str) -> Option<String> {
if let Some(action_part) = consumer_trait.strip_prefix("Can") {
Some(format!("{}Component", action_part))
} else {
None
}
}
pub fn strip_module_prefixes(message: &str) -> String {
let mut result = message.to_string();
for _ in 0..5 {
result = result.replace("cgp::prelude::", "");
result = result.replace("cgp::", "");
}
if result.starts_with("IsProviderFor<") {
if let Some(start) = result.find('<') {
let after_start = start + 1;
result = result[after_start..].to_string();
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_derive_provider_trait_name() {
assert_eq!(
derive_provider_trait_name("AreaCalculatorComponent"),
Some("AreaCalculator".to_string())
);
assert_eq!(
derive_provider_trait_name("FooComponent"),
Some("Foo".to_string())
);
assert_eq!(derive_provider_trait_name("Component"), None);
assert_eq!(derive_provider_trait_name("NoSuffix"), None);
}
#[test]
fn test_extract_symbol_length() {
let text = "Symbol<6, Chars<'h', Chars<'e', ...>>>";
assert_eq!(extract_symbol_length(text), Some(6));
let text2 = "Symbol<5, Chars<'w', ...>>";
assert_eq!(extract_symbol_length(text2), Some(5));
}
#[test]
fn test_extract_chars_from_pattern() {
let text = "Chars<'h', Chars<'e', Chars<'i', Chars<'g', Chars<'h', Chars<'t', Nil>>>>>>";
let (chars, has_unknown) = extract_chars_from_pattern(text);
assert_eq!(chars, vec!['h', 'e', 'i', 'g', 'h', 't']);
assert!(!has_unknown);
let text2 = "Chars<'w', Chars<'i', Chars<'d', Chars<_, Chars<'h', Nil>>>>>";
let (chars2, has_unknown2) = extract_chars_from_pattern(text2);
assert_eq!(chars2, vec!['w', 'i', 'd', '\u{FFFD}', 'h']);
assert!(has_unknown2);
}
#[test]
fn test_derive_component_from_consumer_trait() {
assert_eq!(
derive_component_from_consumer_trait("CanCalculateArea"),
Some("CalculateAreaComponent".to_string())
);
assert_eq!(
derive_component_from_consumer_trait("CanFoo"),
Some("FooComponent".to_string())
);
assert_eq!(
derive_component_from_consumer_trait("NotAConsumerTrait"),
None
);
}
#[test]
fn test_extract_consumer_trait_dependency() {
let note = "required for `Rectangle` to implement `CanCalculateArea`";
let dep = extract_consumer_trait_dependency(note).unwrap();
assert_eq!(dep.trait_name, "CanCalculateArea");
assert_eq!(dep.context_type, "Rectangle");
assert_eq!(
dep.component_name,
Some("CalculateAreaComponent".to_string())
);
let note2 = "required for `Rectangle` to implement `CanUseComponent<Something>`";
assert!(extract_consumer_trait_dependency(note2).is_none());
}
}