Skip to main content

yuru_core/
query.rs

1use std::collections::HashMap;
2
3use crate::{KeyKind, LangMode, LanguageBackend, QueryVariantKind};
4
5#[derive(Clone, Debug, Eq, PartialEq)]
6pub struct QueryVariant {
7    pub text: String,
8    pub kind: QueryVariantKind,
9    pub weight: i32,
10}
11
12impl QueryVariant {
13    pub fn original(text: impl Into<String>) -> Self {
14        Self {
15            text: text.into(),
16            kind: QueryVariantKind::Original,
17            weight: 500,
18        }
19    }
20
21    pub fn normalized(text: impl Into<String>) -> Self {
22        Self {
23            text: text.into(),
24            kind: QueryVariantKind::Normalized,
25            weight: 450,
26        }
27    }
28
29    pub fn kana(text: impl Into<String>) -> Self {
30        Self {
31            text: text.into(),
32            kind: QueryVariantKind::Kana,
33            weight: 350,
34        }
35    }
36
37    pub fn romaji_to_kana(text: impl Into<String>) -> Self {
38        Self {
39            text: text.into(),
40            kind: QueryVariantKind::RomajiToKana,
41            weight: 200,
42        }
43    }
44
45    pub fn pinyin(text: impl Into<String>) -> Self {
46        Self {
47            text: text.into(),
48            kind: QueryVariantKind::Pinyin,
49            weight: 250,
50        }
51    }
52
53    pub fn initials(text: impl Into<String>) -> Self {
54        Self {
55            text: text.into(),
56            kind: QueryVariantKind::Initials,
57            weight: 250,
58        }
59    }
60}
61
62#[derive(Clone, Debug, Default)]
63pub struct PlainBackend;
64
65impl LanguageBackend for PlainBackend {
66    fn mode(&self) -> LangMode {
67        LangMode::Plain
68    }
69
70    fn build_candidate_keys(&self, _text: &str) -> Vec<crate::SearchKey> {
71        Vec::new()
72    }
73
74    fn expand_query(&self, query: &str) -> Vec<QueryVariant> {
75        base_query_variants(query)
76    }
77}
78
79pub fn base_query_variants(query: &str) -> Vec<QueryVariant> {
80    let mut variants = vec![QueryVariant::original(query)];
81    let normalized = crate::normalize::normalize(query);
82    if normalized != query {
83        variants.push(QueryVariant::normalized(normalized));
84    }
85    variants
86}
87
88pub fn dedup_and_limit_variants(
89    variants: Vec<QueryVariant>,
90    max_query_variants: usize,
91) -> Vec<QueryVariant> {
92    let mut seen_coverage_by_text = HashMap::new();
93    let mut out = Vec::new();
94
95    for variant in variants {
96        let coverage = key_kind_coverage(variant.kind);
97        let seen_coverage = seen_coverage_by_text
98            .entry(variant.text.clone())
99            .or_insert(0u16);
100        if coverage & !*seen_coverage != 0 {
101            *seen_coverage |= coverage;
102            out.push(variant);
103        }
104        if out.len() >= max_query_variants {
105            break;
106        }
107    }
108
109    out
110}
111
112fn key_kind_coverage(kind: QueryVariantKind) -> u16 {
113    const ORIGINAL: u16 = 1 << 0;
114    const NORMALIZED: u16 = 1 << 1;
115    const KANA_READING: u16 = 1 << 2;
116    const ROMAJI_READING: u16 = 1 << 3;
117    const PINYIN_FULL: u16 = 1 << 4;
118    const PINYIN_JOINED: u16 = 1 << 5;
119    const PINYIN_INITIALS: u16 = 1 << 6;
120    const KOREAN_ROMANIZED: u16 = 1 << 7;
121    const KOREAN_INITIALS: u16 = 1 << 8;
122    const KOREAN_KEYBOARD: u16 = 1 << 9;
123    const LEARNED_ALIAS: u16 = 1 << 10;
124
125    match kind {
126        QueryVariantKind::Original | QueryVariantKind::Normalized => {
127            ORIGINAL
128                | NORMALIZED
129                | ROMAJI_READING
130                | PINYIN_FULL
131                | PINYIN_JOINED
132                | KOREAN_ROMANIZED
133                | KOREAN_INITIALS
134                | KOREAN_KEYBOARD
135                | LEARNED_ALIAS
136        }
137        QueryVariantKind::Kana | QueryVariantKind::RomajiToKana => KANA_READING,
138        QueryVariantKind::Pinyin => PINYIN_FULL | PINYIN_JOINED,
139        QueryVariantKind::Initials => PINYIN_INITIALS | KOREAN_INITIALS | LEARNED_ALIAS,
140    }
141}
142
143pub fn key_kind_allowed(variant: &QueryVariant, kind: KeyKind) -> bool {
144    match variant.kind {
145        QueryVariantKind::Original | QueryVariantKind::Normalized => matches!(
146            kind,
147            KeyKind::Original
148                | KeyKind::Normalized
149                | KeyKind::RomajiReading
150                | KeyKind::PinyinFull
151                | KeyKind::PinyinJoined
152                | KeyKind::KoreanRomanized
153                | KeyKind::KoreanInitials
154                | KeyKind::KoreanKeyboard
155                | KeyKind::LearnedAlias
156        ),
157        QueryVariantKind::Kana | QueryVariantKind::RomajiToKana => {
158            matches!(kind, KeyKind::KanaReading)
159        }
160        QueryVariantKind::Pinyin => matches!(kind, KeyKind::PinyinFull | KeyKind::PinyinJoined),
161        QueryVariantKind::Initials => {
162            matches!(
163                kind,
164                KeyKind::PinyinInitials | KeyKind::KoreanInitials | KeyKind::LearnedAlias
165            )
166        }
167    }
168}
169
170#[cfg(test)]
171mod tests {
172    use crate::{KeyKind, QueryVariantKind};
173
174    use super::*;
175
176    #[test]
177    fn plain_query_expansion_is_small() {
178        let vars = PlainBackend.expand_query("Tokyo");
179        assert!(vars.iter().any(|v| v.text == "Tokyo"));
180        assert!(vars.iter().any(|v| v.text == "tokyo"));
181        assert!(vars.len() <= 2);
182    }
183
184    #[test]
185    fn empty_query_does_not_panic() {
186        let vars = PlainBackend.expand_query("");
187        assert!(vars.len() <= 1);
188    }
189
190    #[test]
191    fn romaji_to_kana_variant_only_targets_kana_keys() {
192        let variant = QueryVariant::romaji_to_kana("とうきょう");
193        assert!(key_kind_allowed(&variant, KeyKind::KanaReading));
194        assert!(!key_kind_allowed(&variant, KeyKind::PinyinJoined));
195    }
196
197    #[test]
198    fn kana_variant_only_targets_kana_keys() {
199        let variant = QueryVariant::kana("はち");
200        assert!(key_kind_allowed(&variant, KeyKind::KanaReading));
201        assert!(!key_kind_allowed(&variant, KeyKind::RomajiReading));
202    }
203
204    #[test]
205    fn pinyin_initial_variant_only_targets_pinyin_initials_and_aliases() {
206        let variant = QueryVariant {
207            text: "bjdx".to_string(),
208            kind: QueryVariantKind::Initials,
209            weight: 0,
210        };
211
212        assert!(key_kind_allowed(&variant, KeyKind::PinyinInitials));
213        assert!(key_kind_allowed(&variant, KeyKind::KoreanInitials));
214        assert!(key_kind_allowed(&variant, KeyKind::LearnedAlias));
215        assert!(!key_kind_allowed(&variant, KeyKind::KanaReading));
216    }
217
218    #[test]
219    fn dedup_preserves_same_text_when_it_adds_key_coverage() {
220        let variants = dedup_and_limit_variants(
221            vec![
222                QueryVariant::original("bjdx"),
223                QueryVariant::initials("bjdx"),
224                QueryVariant::pinyin("bjdx"),
225                QueryVariant::initials("bjdx"),
226            ],
227            8,
228        );
229
230        assert_eq!(variants.len(), 2);
231        assert_eq!(variants[0].kind, QueryVariantKind::Original);
232        assert_eq!(variants[1].kind, QueryVariantKind::Initials);
233    }
234}