lance_tokenizer/
ascii_folding_filter.rs1use std::mem;
5
6use unicode_normalization::{UnicodeNormalization, char::is_combining_mark};
7
8use crate::{Token, TokenFilter, TokenStream, Tokenizer};
9
10#[derive(Clone)]
11pub struct AsciiFoldingFilter;
12
13impl TokenFilter for AsciiFoldingFilter {
14 type Tokenizer<T: Tokenizer> = AsciiFoldingFilterWrapper<T>;
15
16 fn transform<T: Tokenizer>(self, tokenizer: T) -> Self::Tokenizer<T> {
17 AsciiFoldingFilterWrapper {
18 tokenizer,
19 buffer: String::new(),
20 }
21 }
22}
23
24#[derive(Clone)]
25pub struct AsciiFoldingFilterWrapper<T> {
26 tokenizer: T,
27 buffer: String,
28}
29
30impl<T: Tokenizer> Tokenizer for AsciiFoldingFilterWrapper<T> {
31 type TokenStream<'a> = AsciiFoldingFilterTokenStream<'a, T::TokenStream<'a>>;
32
33 fn token_stream<'a>(&'a mut self, text: &'a str) -> Self::TokenStream<'a> {
34 self.buffer.clear();
35 AsciiFoldingFilterTokenStream {
36 buffer: &mut self.buffer,
37 tail: self.tokenizer.token_stream(text),
38 }
39 }
40}
41
42pub struct AsciiFoldingFilterTokenStream<'a, T> {
43 buffer: &'a mut String,
44 tail: T,
45}
46
47impl<T: TokenStream> TokenStream for AsciiFoldingFilterTokenStream<'_, T> {
48 fn advance(&mut self) -> bool {
49 if !self.tail.advance() {
50 return false;
51 }
52 let token = self.tail.token_mut();
53 if !token.text.is_ascii() {
54 to_ascii(&token.text, self.buffer);
55 mem::swap(&mut token.text, self.buffer);
56 }
57 true
58 }
59
60 fn token(&self) -> &Token {
61 self.tail.token()
62 }
63
64 fn token_mut(&mut self) -> &mut Token {
65 self.tail.token_mut()
66 }
67}
68
69fn to_ascii(text: &str, output: &mut String) {
70 output.clear();
71 output.reserve(text.len());
72 for ch in text.chars() {
73 if ch.is_ascii() {
74 output.push(ch);
75 continue;
76 }
77
78 if let Some(mapped) = fold_char(ch) {
79 output.push_str(mapped);
80 continue;
81 }
82
83 let original_len = output.len();
84 for decomposed in ch.nfkd() {
85 if decomposed.is_ascii() {
86 output.push(decomposed);
87 } else if is_combining_mark(decomposed) {
88 continue;
89 } else if let Some(mapped) = fold_char(decomposed) {
90 output.push_str(mapped);
91 }
92 }
93
94 if output.len() == original_len {
95 output.push(ch);
96 }
97 }
98}
99
100fn fold_char(ch: char) -> Option<&'static str> {
101 match ch {
102 'ß' => Some("ss"),
103 'ẞ' => Some("SS"),
104 'Æ' => Some("AE"),
105 'æ' => Some("ae"),
106 'Œ' => Some("OE"),
107 'œ' => Some("oe"),
108 'Ø' => Some("O"),
109 'ø' => Some("o"),
110 'Ł' => Some("L"),
111 'ł' => Some("l"),
112 'Đ' | 'Ð' => Some("D"),
113 'đ' | 'ð' => Some("d"),
114 'Þ' => Some("TH"),
115 'þ' => Some("th"),
116 'Ħ' => Some("H"),
117 'ħ' => Some("h"),
118 'Ŧ' => Some("T"),
119 'ŧ' => Some("t"),
120 'Ŋ' => Some("N"),
121 'ŋ' => Some("n"),
122 'ı' => Some("i"),
123 'ĸ' => Some("k"),
124 'ſ' => Some("s"),
125 _ => None,
126 }
127}
128
129#[cfg(test)]
130mod tests {
131 use crate::{AsciiFoldingFilter, RawTokenizer, TextAnalyzer, Token};
132
133 fn collect_tokens(text: &str) -> Vec<Token> {
134 let mut analyzer = TextAnalyzer::builder(RawTokenizer::default())
135 .filter(AsciiFoldingFilter)
136 .build();
137 let mut stream = analyzer.token_stream(text);
138 let mut tokens = Vec::new();
139 stream.process(&mut |token| tokens.push(token.clone()));
140 tokens
141 }
142
143 #[test]
144 fn test_ascii_folding_accents() {
145 let tokens = collect_tokens("café");
146 assert_eq!(tokens[0].text, "cafe");
147 }
148
149 #[test]
150 fn test_ascii_folding_sharp_s() {
151 let tokens = collect_tokens("straße");
152 assert_eq!(tokens[0].text, "strasse");
153 }
154
155 #[test]
156 fn test_ascii_folding_cjk_unchanged() {
157 let tokens = collect_tokens("こんにちは世界");
158 assert_eq!(tokens[0].text, "こんにちは世界");
159 }
160}