1#[derive(Debug, Clone, Copy, PartialEq, Eq)]
11pub enum Language {
12 English,
13 Russian,
14 Hindi,
15 Chinese,
16 Unknown,
17}
18
19#[cfg(not(target_arch = "wasm32"))]
30mod forced_language {
31 use super::Language;
32 use std::cell::Cell;
33
34 thread_local! {
35 static FORCED_LANGUAGE: Cell<Option<Language>> = const { Cell::new(None) };
36 }
37
38 pub(super) fn replace(language: Option<Language>) -> Option<Language> {
39 FORCED_LANGUAGE.with(|slot| slot.replace(language))
40 }
41
42 pub(super) fn set(language: Option<Language>) {
43 FORCED_LANGUAGE.with(|slot| slot.set(language));
44 }
45
46 pub(super) fn get() -> Option<Language> {
47 FORCED_LANGUAGE.with(Cell::get)
48 }
49}
50
51#[cfg(target_arch = "wasm32")]
52mod forced_language {
53 use super::Language;
54 use core::cell::Cell;
55
56 struct Slot(Cell<Option<Language>>);
57
58 unsafe impl Sync for Slot {}
61
62 static FORCED_LANGUAGE: Slot = Slot(Cell::new(None));
63
64 pub(super) fn replace(language: Option<Language>) -> Option<Language> {
65 FORCED_LANGUAGE.0.replace(language)
66 }
67
68 pub(super) fn set(language: Option<Language>) {
69 FORCED_LANGUAGE.0.set(language);
70 }
71
72 pub(super) fn get() -> Option<Language> {
73 FORCED_LANGUAGE.0.get()
74 }
75}
76
77#[must_use]
83pub fn set_forced_language(language: Option<Language>) -> ForcedLanguageGuard {
84 let previous = forced_language::replace(language);
85 ForcedLanguageGuard { previous }
86}
87
88#[must_use]
90pub fn from_slug(slug: &str) -> Option<Language> {
91 match slug {
92 "en" => Some(Language::English),
93 "ru" => Some(Language::Russian),
94 "hi" => Some(Language::Hindi),
95 "zh" => Some(Language::Chinese),
96 _ => None,
97 }
98}
99
100pub struct ForcedLanguageGuard {
102 previous: Option<Language>,
103}
104
105impl Drop for ForcedLanguageGuard {
106 fn drop(&mut self) {
107 forced_language::set(self.previous);
108 }
109}
110
111#[derive(Debug, Clone, Copy, PartialEq, Eq)]
112enum Script {
113 Latin,
114 Cyrillic,
115 Devanagari,
116 Cjk,
117 Other,
118}
119
120impl Language {
121 #[must_use]
122 pub const fn slug(self) -> &'static str {
123 match self {
124 Self::English => "en",
125 Self::Russian => "ru",
126 Self::Hindi => "hi",
127 Self::Chinese => "zh",
128 Self::Unknown => "unknown",
129 }
130 }
131}
132
133#[must_use]
142pub fn detect(prompt: &str) -> Language {
143 if let Some(forced) = forced_language::get() {
146 return forced;
147 }
148 let mut latin = 0usize;
149 let mut cyrillic = 0usize;
150 let mut devanagari = 0usize;
151 let mut cjk = 0usize;
152 let mut other_script = 0usize;
153 let mut first_script = None;
154
155 for character in prompt.chars() {
156 let codepoint = u32::from(character);
157 if character.is_ascii_alphabetic() {
158 latin += 1;
159 first_script.get_or_insert(Script::Latin);
160 } else if (0x0400..=0x04FF).contains(&codepoint) {
161 cyrillic += 1;
162 first_script.get_or_insert(Script::Cyrillic);
163 } else if (0x0900..=0x097F).contains(&codepoint) {
164 devanagari += 1;
165 first_script.get_or_insert(Script::Devanagari);
166 } else if (0x4E00..=0x9FFF).contains(&codepoint) {
167 cjk += 1;
168 first_script.get_or_insert(Script::Cjk);
169 } else if character.is_alphabetic() {
170 other_script += 1;
171 first_script.get_or_insert(Script::Other);
172 }
173 }
174
175 let total_script = latin + cyrillic + devanagari + cjk + other_script;
176 if total_script == 0 {
177 return Language::English;
178 }
179
180 if other_script > latin
181 && other_script >= cyrillic
182 && other_script >= devanagari
183 && other_script >= cjk
184 {
185 return Language::Unknown;
186 }
187 if latin > 0 {
188 if let Some(language) = marker_language(prompt, cyrillic, devanagari, cjk) {
189 return language;
190 }
191 match first_script {
192 Some(Script::Cyrillic) if cyrillic >= devanagari.max(cjk) => {
193 return Language::Russian;
194 }
195 Some(Script::Devanagari) if devanagari >= cyrillic.max(cjk) => {
196 return Language::Hindi;
197 }
198 Some(Script::Cjk) if cjk >= cyrillic.max(devanagari) => return Language::Chinese,
199 _ => {}
200 }
201 }
202 if cyrillic >= latin.max(devanagari).max(cjk) && cyrillic > 0 {
203 return Language::Russian;
204 }
205 if devanagari >= latin.max(cyrillic).max(cjk) && devanagari > 0 {
206 return Language::Hindi;
207 }
208 if cjk >= latin.max(cyrillic).max(devanagari) && cjk > 0 {
209 return Language::Chinese;
210 }
211 Language::English
212}
213
214fn marker_language(
215 prompt: &str,
216 cyrillic: usize,
217 devanagari: usize,
218 cjk: usize,
219) -> Option<Language> {
220 let normalized = prompt.to_lowercase();
221 let mut best = None;
222 for (script, count, language) in [
223 (Script::Cyrillic, cyrillic, Language::Russian),
224 (Script::Devanagari, devanagari, Language::Hindi),
225 (Script::Cjk, cjk, Language::Chinese),
226 ] {
227 if count == 0 || !contains_question_marker(&normalized, script) {
228 continue;
229 }
230 match best {
231 Some((best_count, _)) if count <= best_count => {}
232 _ => best = Some((count, language)),
233 }
234 }
235 best.map(|(_, language)| language)
236}
237
238fn contains_question_marker(prompt: &str, script: Script) -> bool {
239 let markers: &[&str] = match script {
240 Script::Cyrillic => &[
241 "\u{0447}\u{0442}\u{043e}",
242 "\u{043a}\u{0430}\u{043a}",
243 "\u{043a}\u{0442}\u{043e}",
244 "\u{0433}\u{0434}\u{0435}",
245 "\u{043a}\u{043e}\u{0433}\u{0434}\u{0430}",
246 "\u{043f}\u{043e}\u{0447}\u{0435}\u{043c}\u{0443}",
247 ],
248 Script::Devanagari => &[
249 "\u{0915}\u{094d}\u{092f}\u{093e}",
250 "\u{0915}\u{094c}\u{0928}",
251 "\u{0915}\u{0939}\u{093e}\u{0901}",
252 "\u{0915}\u{092c}",
253 "\u{0915}\u{0948}\u{0938}\u{0947}",
254 "\u{0915}\u{094d}\u{092f}\u{094b}\u{0902}",
255 ],
256 Script::Cjk => &[
257 "\u{4ec0}\u{4e48}",
258 "\u{5417}",
259 "\u{600e}\u{4e48}",
260 "\u{8c01}",
261 "\u{54ea}",
262 ],
263 Script::Latin | Script::Other => &[],
264 };
265 markers.iter().any(|marker| prompt.contains(marker))
266}