1pub use aria_ffi::{
4 aria_complete, aria_complete_stream, aria_embed, aria_last_error, aria_model_destroy,
5 aria_model_init, aria_transcribe, AriaModelHandle,
6};
7pub use aria_inference::{EngineError, GenerateOpts, Generation, Session, SessionBuilder};
8
9mod download;
10mod setup;
11pub use download::{download_model, download_model_setup, ensure_ffi_lib, DownloadError};
12pub use setup::{
13 apply_setup, fill_setup_urls, SetupConfig, SetupError, SetupUpdates, CN_SITE, CN_UPGRADE,
14};
15
16#[derive(Default, Clone)]
18pub struct OpenOptions {
19 pub token: Option<String>,
21 pub hf_token: Option<String>,
23 pub modelscope_api_token: Option<String>,
25 pub site: Option<String>,
27}
28
29pub struct Engine {
31 session: Option<Session>,
32 cfg: SetupConfig,
33 generic_token: Option<String>,
34}
35
36impl Engine {
37 pub fn new() -> Self {
39 Self {
40 session: None,
41 cfg: SetupConfig::default(),
42 generic_token: None,
43 }
44 }
45
46 pub fn setup(&mut self, updates: &SetupUpdates) -> Result<&mut Self, SetupError> {
48 self.cfg = apply_setup(&self.cfg, updates)?;
49 Ok(self)
50 }
51
52 pub fn setup_status(&self) -> &SetupConfig {
53 &self.cfg
54 }
55
56 pub fn setup_clear(&mut self) -> &mut Self {
58 self.cfg = SetupConfig::default();
59 self
60 }
61
62 pub fn from_bundle(bundle_path: impl AsRef<std::path::Path>) -> Result<Self, EngineError> {
64 let session = SessionBuilder::new().model(bundle_path).build()?;
65 Ok(Self {
66 session: Some(session),
67 cfg: SetupConfig::default(),
68 generic_token: None,
69 })
70 }
71
72 pub fn open(bundle_path: impl AsRef<std::path::Path>) -> Result<Self, EngineError> {
74 Self::from_bundle(bundle_path)
75 }
76
77 fn load_ref(&mut self, model_ref: &str, opts: &OpenOptions) -> Result<(), OpenError> {
78 let _ = ensure_ffi_lib(opts.site.as_deref())?;
79 let path = if model_ref.contains('/')
80 || model_ref.contains('\\')
81 || std::path::Path::new(model_ref).exists()
82 {
83 std::path::PathBuf::from(model_ref)
84 } else {
85 download_model_setup(
86 model_ref,
87 opts.token.as_deref().unwrap_or(""),
88 opts.site.as_deref(),
89 opts.hf_token.as_deref(),
90 opts.modelscope_api_token.as_deref(),
91 )
92 .map_err(OpenError::Download)?
93 };
94 let session = SessionBuilder::new()
95 .model(&path)
96 .build()
97 .map_err(OpenError::Engine)?;
98 self.session = Some(session);
99 Ok(())
100 }
101
102 pub fn open_named(&mut self, model_ref: &str) -> Result<&mut Self, OpenError> {
104 let opts = OpenOptions {
105 token: self.generic_token.clone(),
106 hf_token: if self.cfg.hf_token.is_empty() {
107 None
108 } else {
109 Some(self.cfg.hf_token.clone())
110 },
111 modelscope_api_token: if self.cfg.modelscope_api_token.is_empty() {
112 None
113 } else {
114 Some(self.cfg.modelscope_api_token.clone())
115 },
116 site: if self.cfg.site_url.is_empty() {
117 None
118 } else {
119 Some(self.cfg.site_url.clone())
120 },
121 };
122 self.load_ref(model_ref, &opts)?;
123 Ok(self)
124 }
125
126 pub fn open_model(model_ref: &str, opts: &OpenOptions) -> Result<Self, OpenError> {
131 let mut eng = Self::new();
132 let updates = SetupUpdates {
133 site_url: opts.site.clone(),
134 hf_token: opts.hf_token.clone(),
135 modelscope_api_token: opts.modelscope_api_token.clone(),
136 ..Default::default()
137 };
138 eng.generic_token = opts.token.clone();
139 let _ = eng.setup(&updates);
140 eng.load_ref(model_ref, opts)?;
141 Ok(eng)
142 }
143
144 fn session_mut(&mut self) -> Result<&mut Session, EngineError> {
145 self.session
146 .as_mut()
147 .ok_or_else(|| EngineError::InvalidParam("engine not opened".into()))
148 }
149
150 fn session_ref(&self) -> Result<&Session, EngineError> {
151 self.session
152 .as_ref()
153 .ok_or_else(|| EngineError::InvalidParam("engine not opened".into()))
154 }
155
156 pub fn complete(&mut self, prompt: &str, opts: &GenerateOpts) -> Result<Generation, EngineError> {
157 let turns = [aria_inference::ChatTurn::new("user", prompt)];
158 let session = self.session_mut()?;
159 let tokens = session.encode_chat(&turns);
160 session.generate(&tokens, opts)
161 }
162
163 pub fn embed(&self, text: &str) -> Result<Vec<f32>, EngineError> {
164 self.session_ref()?.embed_text(text)
165 }
166
167 pub fn transcribe(&self, pcm: &[u8]) -> Result<String, EngineError> {
168 self.session_ref()?.transcribe_pcm16le(pcm)
169 }
170}
171
172impl Default for Engine {
173 fn default() -> Self {
174 Self::new()
175 }
176}
177
178#[derive(Debug, thiserror::Error)]
180pub enum OpenError {
181 #[error("download failed: {0}")]
182 Download(#[from] DownloadError),
183 #[error("engine load failed: {0}")]
184 Engine(#[from] EngineError),
185}
186
187#[cfg(test)]
188mod tests {
189 use super::*;
190 use aria_inference::fixture::write_tiny_q4_bundle;
191
192 #[test]
193 fn engine_complete_ok() {
194 let dir = tempfile::tempdir().unwrap();
195 write_tiny_q4_bundle(dir.path()).unwrap();
196 let mut eng = Engine::open(dir.path()).unwrap();
197 let g = eng
198 .complete("hi", &GenerateOpts { max_tokens: 2, temperature: 0.0 })
199 .unwrap();
200 assert!(!g.text.is_empty());
201 assert!(!eng.embed("x").unwrap().is_empty());
202 assert!(!eng.transcribe(&[0, 1, 2, 3]).unwrap().is_empty());
203 }
204
205 #[test]
206 fn open_model_local_path_no_token() {
207 let _guard = crate::download::ENV_LOCK.lock().unwrap();
208 let home = tempfile::tempdir().unwrap();
209 std::env::set_var("ARIA_COMPUTE_HOME", home.path());
210 let libdir = home.path().join("lib");
211 std::fs::create_dir_all(&libdir).unwrap();
212 let name = if cfg!(windows) {
213 "aria-engine_ffi.dll"
214 } else if cfg!(target_os = "macos") {
215 "libaria-engine_ffi.dylib"
216 } else {
217 "libaria-engine_ffi.so"
218 };
219 std::fs::write(libdir.join(name), b"x").unwrap();
220 let dir = tempfile::tempdir().unwrap();
221 write_tiny_q4_bundle(dir.path()).unwrap();
222 let mut eng = Engine::open_model(dir.path().to_str().unwrap(), &OpenOptions::default()).unwrap();
223 let g = eng
224 .complete("hi", &GenerateOpts { max_tokens: 2, temperature: 0.0 })
225 .unwrap();
226 assert!(!g.text.is_empty());
227 std::env::remove_var("ARIA_COMPUTE_HOME");
228 }
229
230 #[test]
231 fn setup_instance_all_fields() {
232 let mut eng = Engine::new();
233 eng.setup(&SetupUpdates {
234 router: Some("http://127.0.0.1:8080".into()),
235 router_api_key: Some("sk-aria_test".into()),
236 site_url: Some(crate::setup::CN_SITE.into()),
237 upgrade_url: Some(crate::setup::CN_UPGRADE.into()),
238 compute: Some("cpu".into()),
239 hf_token: Some("hf_abc".into()),
240 modelscope_api_token: Some("ms_xyz".into()),
241 })
242 .unwrap();
243 let st = eng.setup_status();
244 assert_eq!(st.router, "http://127.0.0.1:8080");
245 assert_eq!(st.router_api_key, "sk-aria_test");
246 assert_eq!(st.compute, "cpu");
247 assert_eq!(st.hf_token, "hf_abc");
248 assert_eq!(st.modelscope_api_token, "ms_xyz");
249 assert_eq!(st.site_url, crate::setup::CN_SITE);
250 }
251
252 #[test]
253 fn setup_accepts_bfvk_router_api_key() {
254 let mut eng = Engine::new();
255 eng.setup(&SetupUpdates {
256 router_api_key: Some("bfvk-cloud".into()),
257 ..Default::default()
258 })
259 .unwrap();
260 assert_eq!(eng.setup_status().router_api_key, "bfvk-cloud");
261 assert!(eng
262 .setup(&SetupUpdates {
263 router_api_key: Some("bad-key".into()),
264 ..Default::default()
265 })
266 .is_err());
267 }
268
269 #[test]
270 fn setup_partial_merge() {
271 let mut eng = Engine::new();
272 eng.setup(&SetupUpdates {
273 hf_token: Some("hf_one".into()),
274 router: Some("http://127.0.0.1:1".into()),
275 ..Default::default()
276 })
277 .unwrap();
278 eng.setup(&SetupUpdates {
279 compute: Some("cuda".into()),
280 ..Default::default()
281 })
282 .unwrap();
283 let st = eng.setup_status();
284 assert_eq!(st.hf_token, "hf_one");
285 assert_eq!(st.router, "http://127.0.0.1:1");
286 assert_eq!(st.compute, "cuda");
287 }
288
289 #[test]
290 fn setup_invalid_enum_leaves_state() {
291 let mut eng = Engine::new();
292 eng.setup(&SetupUpdates {
293 compute: Some("cpu".into()),
294 ..Default::default()
295 })
296 .unwrap();
297 assert!(eng
298 .setup(&SetupUpdates {
299 compute: Some("gpu".into()),
300 ..Default::default()
301 })
302 .is_err());
303 assert_eq!(eng.setup_status().compute, "cpu");
304 }
305
306 #[test]
307 fn setup_clear_resets_instance() {
308 let mut eng = Engine::new();
309 eng.setup(&SetupUpdates {
310 hf_token: Some("hf_x".into()),
311 compute: Some("cpu".into()),
312 ..Default::default()
313 })
314 .unwrap();
315 eng.setup_clear();
316 let st = eng.setup_status();
317 assert_eq!(st.hf_token, "");
318 assert_eq!(st.compute, "auto");
319 }
320
321 #[test]
322 fn setup_fills_urls_from_site_tld() {
323 let mut eng = Engine::new();
324 eng.setup(&SetupUpdates {
325 site_url: Some("https://ariacompute.cn".into()),
326 ..Default::default()
327 })
328 .unwrap();
329 let st = eng.setup_status();
330 assert_eq!(st.upgrade_url, crate::setup::CN_UPGRADE);
331 }
332
333 #[test]
334 fn setup_does_not_write_engine_yml() {
335 let _guard = crate::download::ENV_LOCK.lock().unwrap();
336 let home = tempfile::tempdir().unwrap();
337 std::env::set_var("ARIA_COMPUTE_HOME", home.path());
338 let mut eng = Engine::new();
339 eng.setup(&SetupUpdates {
340 router: Some("http://127.0.0.1:8080".into()),
341 site_url: Some("https://ariacompute.com".into()),
342 hf_token: Some("hf_x".into()),
343 ..Default::default()
344 })
345 .unwrap();
346 assert!(!home.path().join("engine.yml").is_file());
347 assert!(!home.path().join("config.yml").is_file());
348 std::env::remove_var("ARIA_COMPUTE_HOME");
349 }
350}