Skip to main content

icu_segmenter/complex/
dictionary.rs

1// This file is part of ICU4X. For terms of use, please see the file
2// called LICENSE at the top level of the ICU4X source tree
3// (online at: https://github.com/unicode-org/icu4x/blob/main/LICENSE ).
4
5use crate::grapheme::*;
6use crate::indices::Utf16Indices;
7use crate::provider::*;
8use crate::scaffold::{Utf16, Utf8};
9use core::str::CharIndices;
10use icu_collections::char16trie::{Char16Trie, TrieResult};
11
12/// A trait for dictionary based iterator
13trait DictionaryType {
14    /// The iterator over characters.
15    type IterAttr<'s>: Iterator<Item = (usize, Self::CharType)> + Clone;
16
17    /// The character type.
18    type CharType: Copy + Into<u32>;
19
20    fn to_char(c: Self::CharType) -> char;
21    fn char_len(c: Self::CharType) -> usize;
22}
23
24struct DictionaryBreakIterator<
25    'l,
26    's,
27    Y: DictionaryType + ?Sized,
28    X: Iterator<Item = usize> + ?Sized,
29> {
30    trie: Char16Trie<'l>,
31    iter: Y::IterAttr<'s>,
32    len: usize,
33    grapheme_iter: X,
34    // TODO transform value for byte trie
35}
36
37/// Implement the [`Iterator`] trait over the segmenter break opportunities of the given string.
38/// Please see the [module-level documentation](crate) for its usages.
39///
40/// Lifetimes:
41/// - `'l` = lifetime of the segmenter object from which this iterator was created
42/// - `'s` = lifetime of the string being segmented
43///
44/// [`Iterator`]: core::iter::Iterator
45impl<Y: DictionaryType + ?Sized, X: Iterator<Item = usize> + ?Sized> Iterator
46    for DictionaryBreakIterator<'_, '_, Y, X>
47{
48    type Item = usize;
49
50    fn next(&mut self) -> Option<Self::Item> {
51        let mut trie_iter = self.trie.iter();
52        let mut intermediate_length = 0;
53        let mut not_match = false;
54        let mut previous_match = None;
55        let mut last_grapheme_offset = 0;
56
57        while let Some(next) = self.iter.next() {
58            let ch = Y::to_char(next.1);
59            match trie_iter.next(ch) {
60                TrieResult::FinalValue(_) => {
61                    return Some(next.0 + Y::char_len(next.1));
62                }
63                TrieResult::Intermediate(_) => {
64                    // Dictionary has to match with grapheme cluster segment.
65                    // If not, we ignore it.
66                    while last_grapheme_offset < next.0 + Y::char_len(next.1) {
67                        if let Some(offset) = self.grapheme_iter.next() {
68                            last_grapheme_offset = offset;
69                            continue;
70                        }
71                        last_grapheme_offset = self.len;
72                        break;
73                    }
74                    if last_grapheme_offset != next.0 + Y::char_len(next.1) {
75                        continue;
76                    }
77
78                    intermediate_length = next.0 + Y::char_len(next.1);
79                    previous_match = Some(self.iter.clone());
80                }
81                TrieResult::NoMatch => {
82                    if intermediate_length > 0 {
83                        if let Some(previous_match) = previous_match {
84                            // Rewind previous match point
85                            self.iter = previous_match;
86                        }
87                        return Some(intermediate_length);
88                    }
89                    // Not found
90                    return Some(next.0 + Y::char_len(next.1));
91                }
92                TrieResult::NoValue => {
93                    // Prefix string is matched
94                    not_match = true;
95                }
96            }
97        }
98
99        if intermediate_length > 0 {
100            Some(intermediate_length)
101        } else if not_match {
102            // no match by scanning text
103            Some(self.len)
104        } else {
105            None
106        }
107    }
108}
109
110impl DictionaryType for u32 {
111    type IterAttr<'s> = Utf16Indices<'s>;
112    type CharType = u32;
113
114    fn to_char(c: u32) -> char {
115        char::from_u32(c).unwrap_or(char::REPLACEMENT_CHARACTER)
116    }
117
118    fn char_len(c: u32) -> usize {
119        if c >= 0x10000 {
120            2
121        } else {
122            1
123        }
124    }
125}
126
127impl DictionaryType for char {
128    type IterAttr<'s> = CharIndices<'s>;
129    type CharType = char;
130
131    fn to_char(c: char) -> char {
132        c
133    }
134
135    fn char_len(c: char) -> usize {
136        c.len_utf8()
137    }
138}
139
140pub(super) struct DictionarySegmenter<'l> {
141    dict: &'l UCharDictionaryBreakData<'l>,
142    grapheme: GraphemeClusterSegmenterBorrowed<'l>,
143}
144
145impl<'l> DictionarySegmenter<'l> {
146    pub(super) fn new(
147        dict: &'l UCharDictionaryBreakData<'l>,
148        grapheme: GraphemeClusterSegmenterBorrowed<'l>,
149    ) -> Self {
150        // TODO: no way to verify trie data
151        Self { dict, grapheme }
152    }
153
154    /// Create a dictionary based break iterator for an `str` (a UTF-8 string).
155    pub(super) fn segment_str(&'l self, input: &'l str) -> impl Iterator<Item = usize> + 'l {
156        let grapheme_iter = self.grapheme.segment_str(input);
157        DictionaryBreakIterator::<char, GraphemeClusterBreakIterator<Utf8>> {
158            trie: Char16Trie::new(self.dict.trie_data.clone()),
159            iter: input.char_indices(),
160            len: input.len(),
161            grapheme_iter,
162        }
163    }
164
165    /// Create a dictionary based break iterator for a UTF-16 string.
166    pub(super) fn segment_utf16(&'l self, input: &'l [u16]) -> impl Iterator<Item = usize> + 'l {
167        let grapheme_iter = self.grapheme.segment_utf16(input);
168        DictionaryBreakIterator::<u32, GraphemeClusterBreakIterator<Utf16>> {
169            trie: Char16Trie::new(self.dict.trie_data.clone()),
170            iter: Utf16Indices::new(input),
171            len: input.len(),
172            grapheme_iter,
173        }
174    }
175}
176
177#[cfg(test)]
178#[cfg(feature = "serde")]
179mod tests {
180    use super::*;
181    use crate::{GraphemeClusterSegmenter, LineSegmenter, WordSegmenter};
182    use icu_provider::prelude::*;
183
184    #[test]
185    fn burmese_dictionary_test() {
186        let segmenter = LineSegmenter::new_dictionary(Default::default());
187        // From css/css-text/word-break/word-break-normal-my-000.html
188        let s = "မြန်မာစာမြန်မာစာမြန်မာစာ";
189        let result: Vec<usize> = segmenter.segment_str(s).collect();
190        assert_eq!(result, vec![0, 18, 24, 42, 48, 66, 72]);
191
192        let s_utf16: Vec<u16> = s.encode_utf16().collect();
193        let result: Vec<usize> = segmenter.segment_utf16(&s_utf16).collect();
194        assert_eq!(result, vec![0, 6, 8, 14, 16, 22, 24]);
195    }
196
197    #[test]
198    fn cj_dictionary_test() {
199        let response: DataResponse<SegmenterDictionaryAutoV1> = crate::provider::Baked
200            .load(DataRequest {
201                id: DataIdentifierBorrowed::for_marker_attributes(
202                    DataMarkerAttributes::from_str_or_panic("cjdict"),
203                ),
204                ..Default::default()
205            })
206            .unwrap();
207        let word_segmenter = WordSegmenter::new_dictionary(Default::default());
208        let dict_segmenter =
209            DictionarySegmenter::new(response.payload.get(), GraphemeClusterSegmenter::new());
210
211        // Match case
212        let s = "龟山岛龟山岛";
213        let result: Vec<usize> = dict_segmenter.segment_str(s).collect();
214        assert_eq!(result, vec![9, 18]);
215
216        let result: Vec<usize> = word_segmenter.segment_str(s).collect();
217        assert_eq!(result, vec![0, 9, 18]);
218
219        let s_utf16: Vec<u16> = s.encode_utf16().collect();
220        let result: Vec<usize> = dict_segmenter.segment_utf16(&s_utf16).collect();
221        assert_eq!(result, vec![3, 6]);
222
223        let result: Vec<usize> = word_segmenter.segment_utf16(&s_utf16).collect();
224        assert_eq!(result, vec![0, 3, 6]);
225
226        // Match case, then no match case
227        let s = "エディターエディ";
228        let result: Vec<usize> = dict_segmenter.segment_str(s).collect();
229        assert_eq!(result, vec![15, 24]);
230
231        // TODO(#3236): Why is WordSegmenter not returning the middle segment?
232        let result: Vec<usize> = word_segmenter.segment_str(s).collect();
233        assert_eq!(result, vec![0, 24]);
234
235        let s_utf16: Vec<u16> = s.encode_utf16().collect();
236        let result: Vec<usize> = dict_segmenter.segment_utf16(&s_utf16).collect();
237        assert_eq!(result, vec![5, 8]);
238
239        // TODO(#3236): Why is WordSegmenter not returning the middle segment?
240        let result: Vec<usize> = word_segmenter.segment_utf16(&s_utf16).collect();
241        assert_eq!(result, vec![0, 8]);
242    }
243
244    #[test]
245    fn khmer_dictionary_test() {
246        let segmenter = LineSegmenter::new_dictionary(Default::default());
247        let s = "ភាសាខ្មែរភាសាខ្មែរភាសាខ្មែរ";
248        let result: Vec<usize> = segmenter.segment_str(s).collect();
249        assert_eq!(result, vec![0, 27, 54, 81]);
250
251        let s_utf16: Vec<u16> = s.encode_utf16().collect();
252        let result: Vec<usize> = segmenter.segment_utf16(&s_utf16).collect();
253        assert_eq!(result, vec![0, 9, 18, 27]);
254    }
255
256    #[test]
257    fn lao_dictionary_test() {
258        let segmenter = LineSegmenter::new_dictionary(Default::default());
259        let s = "ພາສາລາວພາສາລາວພາສາລາວ";
260        let r: Vec<usize> = segmenter.segment_str(s).collect();
261        assert_eq!(r, vec![0, 12, 21, 33, 42, 54, 63]);
262
263        let s_utf16: Vec<u16> = s.encode_utf16().collect();
264        let r: Vec<usize> = segmenter.segment_utf16(&s_utf16).collect();
265        assert_eq!(r, vec![0, 4, 7, 11, 14, 18, 21]);
266    }
267}