use std::borrow::Cow;
use std::collections::HashMap;
use std::ops::Range;
use std::sync::Arc;
use tree_sitter_highlight::{
Highlight, HighlightConfiguration, HighlightEvent, Highlighter as TSHighlighter,
};
use crate::tokenize;
pub const HIGHLIGHT_NAMES: &[&str] = &[
"attribute",
"boolean",
"comment",
"comment.doc",
"constant",
"embedded",
"function",
"function.definition",
"function.method",
"function.special",
"function.special.definition",
"keyword",
"keyword.control",
"lifetime",
"number",
"operator",
"property",
"punctuation.bracket",
"punctuation.delimiter",
"punctuation.special",
"string",
"string.escape",
"type",
"type.builtin",
"type.interface",
"variable",
"variable.parameter",
"variable.special",
];
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct HighlightSpan {
pub range: Range<usize>,
pub highlight_id: usize,
}
pub fn highlight_id(name: &str) -> usize {
HIGHLIGHT_NAMES
.iter()
.position(|&n| n == name)
.unwrap_or(usize::MAX)
}
fn normalized_lang(lang: &str) -> Cow<'_, str> {
if lang.bytes().any(|b| b.is_ascii_uppercase()) {
Cow::Owned(lang.to_lowercase())
} else {
Cow::Borrowed(lang)
}
}
enum Backend {
TreeSitter(Box<HighlightConfiguration>),
Tokenizer(fn(&str) -> Vec<HighlightSpan>),
}
pub struct Highlighter {
inner: TSHighlighter,
languages: HashMap<String, Arc<Backend>>,
}
impl Default for Highlighter {
fn default() -> Self {
Self::new()
}
}
impl Highlighter {
pub fn new() -> Self {
let inner = TSHighlighter::new();
let mut languages = HashMap::new();
let mut register = |aliases: &[&str], backend: Backend| {
let backend = Arc::new(backend);
for a in aliases {
languages.insert((*a).to_string(), Arc::clone(&backend));
}
};
let rust = tree_sitter_rust::LANGUAGE.into();
if let Some(config) =
Self::create_config(rust, "rust", include_str!("../queries/rust/highlights.scm"))
{
register(&["rust", "rs"], Backend::TreeSitter(Box::new(config)));
}
let bash = tree_sitter_bash::LANGUAGE.into();
if let Some(config) = Self::create_config(bash, "bash", tree_sitter_bash::HIGHLIGHT_QUERY) {
register(
&["bash", "sh", "shell"],
Backend::TreeSitter(Box::new(config)),
);
}
register(
&["mermaid"],
Backend::Tokenizer(tokenize::highlight_mermaid),
);
register(
&["latex", "tex"],
Backend::Tokenizer(tokenize::highlight_latex),
);
Self { inner, languages }
}
fn create_config(
language: tree_sitter::Language,
name: &str,
highlights_query: &str,
) -> Option<HighlightConfiguration> {
let mut config = match HighlightConfiguration::new(language, name, highlights_query, "", "")
{
Ok(c) => c,
Err(e) => {
eprintln!("Failed to create {name} highlight config: {e}");
return None;
}
};
config.configure(HIGHLIGHT_NAMES);
Some(config)
}
pub fn supports_language(&self, lang: &str) -> bool {
self.languages.contains_key(normalized_lang(lang).as_ref())
}
pub fn highlight(&mut self, code: &str, language: &str) -> Vec<HighlightSpan> {
let Some(backend) = self.languages.get(normalized_lang(language).as_ref()) else {
return Vec::new();
};
let config = match backend.as_ref() {
Backend::Tokenizer(f) => return f(code),
Backend::TreeSitter(config) => config.as_ref(),
};
let highlights = match self.inner.highlight(
config,
code.as_bytes(),
None, |_| None, ) {
Ok(h) => h,
Err(_) => return Vec::new(),
};
let mut spans: Vec<HighlightSpan> = Vec::new();
let mut current_highlight: Option<usize> = None;
for event in highlights {
match event {
Ok(HighlightEvent::Source { start, end }) => {
if let Some(highlight_id) = current_highlight {
match spans.last_mut() {
Some(last)
if last.highlight_id == highlight_id && last.range.end == start =>
{
last.range.end = end;
}
_ => spans.push(HighlightSpan {
range: start..end,
highlight_id,
}),
}
}
}
Ok(HighlightEvent::HighlightStart(Highlight(id))) => {
current_highlight = Some(id);
}
Ok(HighlightEvent::HighlightEnd) => {
current_highlight = None;
}
Err(_) => break,
}
}
spans
}
pub fn capture_name(highlight_id: usize) -> &'static str {
HIGHLIGHT_NAMES
.get(highlight_id)
.copied()
.unwrap_or("unknown")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_highlighter_creation() {
let highlighter = Highlighter::new();
assert!(highlighter.supports_language("rust"));
assert!(highlighter.supports_language("rs"));
assert!(highlighter.supports_language("Rust")); assert!(highlighter.supports_language("bash"));
assert!(highlighter.supports_language("sh"));
assert!(highlighter.supports_language("shell"));
assert!(!highlighter.supports_language("python"));
}
#[test]
fn test_highlight_rust_simple() {
let mut highlighter = Highlighter::new();
let code = "let x = 42;";
let spans = highlighter.highlight(code, "rust");
assert!(!spans.is_empty(), "Should produce some highlight spans");
for span in &spans {
eprintln!(
" {:?}: {} @ {:?}",
&code[span.range.clone()],
Highlighter::capture_name(span.highlight_id),
span.range
);
}
let let_spans: Vec<_> = spans
.iter()
.filter(|s| s.range == (0..3) && Highlighter::capture_name(s.highlight_id) == "keyword")
.collect();
assert!(!let_spans.is_empty(), "Should have @keyword for 'let'");
let num_spans: Vec<_> = spans
.iter()
.filter(|s| s.range == (8..10) && Highlighter::capture_name(s.highlight_id) == "number")
.collect();
assert!(!num_spans.is_empty(), "Should have @number for '42'");
}
#[test]
fn test_highlight_rust_function() {
let mut highlighter = Highlighter::new();
let code = "fn main() {}";
let spans = highlighter.highlight(code, "rust");
for span in &spans {
eprintln!(
" {:?}: {} @ {:?}",
&code[span.range.clone()],
Highlighter::capture_name(span.highlight_id),
span.range
);
}
let fn_spans: Vec<_> = spans
.iter()
.filter(|s| s.range == (0..2) && Highlighter::capture_name(s.highlight_id) == "keyword")
.collect();
assert!(!fn_spans.is_empty(), "Should have @keyword for 'fn'");
let main_spans: Vec<_> = spans
.iter()
.filter(|s| {
s.range == (3..7)
&& Highlighter::capture_name(s.highlight_id).starts_with("function")
})
.collect();
assert!(!main_spans.is_empty(), "Should have @function* for 'main'");
}
#[test]
fn test_highlight_no_overlap() {
let mut highlighter = Highlighter::new();
let code = "fn foo(x: i32) -> Result<i32, Error> { Ok(x) }";
let spans = highlighter.highlight(code, "rust");
for span in &spans {
eprintln!(
" {:?}: {} @ {:?}",
&code[span.range.clone()],
Highlighter::capture_name(span.highlight_id),
span.range
);
}
for (i, span1) in spans.iter().enumerate() {
for span2 in spans.iter().skip(i + 1) {
let overlaps =
span1.range.start < span2.range.end && span2.range.start < span1.range.end;
assert!(
!overlaps,
"Spans should not overlap: {:?} and {:?}",
span1, span2
);
}
}
}
}