Skip to main content

aria_engine/
lib.rs

1//! Thin Rust SDK: prefer `aria_inference` for native use; re-exports FFI for embedding tests.
2
3pub 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/// Options controlling model auto-download from the regional public hub.
17#[derive(Default, Clone)]
18pub struct OpenOptions {
19    /// Legacy generic hub token. Dashboard `sk-` / `bfvk-` keys are ignored.
20    pub token: Option<String>,
21    /// Hugging Face hub token (`.com`). Same field as `aria-engine setup` `hf_token`.
22    pub hf_token: Option<String>,
23    /// ModelScope hub token (`.cn`). Same field as `aria-engine setup` `modelscope_api_token`.
24    pub modelscope_api_token: Option<String>,
25    /// Site used to pick the regional hub. Defaults to `https://ariacompute.com` (`.com` → HF, `.cn` → ModelScope).
26    pub site: Option<String>,
27}
28
29/// High-level convenience over `Session`.
30pub struct Engine {
31    session: Option<Session>,
32    cfg: SetupConfig,
33    generic_token: Option<String>,
34}
35
36impl Engine {
37    /// Empty construct. Call [`setup`](Self::setup) then [`open`](Self::open) to download/load.
38    pub fn new() -> Self {
39        Self {
40            session: None,
41            cfg: SetupConfig::default(),
42            generic_token: None,
43        }
44    }
45
46    /// Set Config / Run fields on this instance only. Does not write engine.yml.
47    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    /// Reset instance defaults. Does not delete ~/.ariacompute/engine.yml.
57    pub fn setup_clear(&mut self) -> &mut Self {
58        self.cfg = SetupConfig::default();
59        self
60    }
61
62    /// Open a local Aria quant bundle directory (path only; no download).
63    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    /// Open a local Aria quant bundle directory (path only; no download).
73    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    /// Download (if needed) and load a model using instance setup.
103    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    /// Open a model by reference. If `model_ref` looks like a local path
127    /// (contains a separator or exists on disk) it is loaded directly;
128    /// otherwise it is treated as a model name and auto-downloaded from the
129    /// regional public hub into `~/.ariacompute/models/{model}` before loading.
130    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/// Error returned by [`Engine::open_model`].
179#[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_partial_merge() {
254        let mut eng = Engine::new();
255        eng.setup(&SetupUpdates {
256            hf_token: Some("hf_one".into()),
257            router: Some("http://127.0.0.1:1".into()),
258            ..Default::default()
259        })
260        .unwrap();
261        eng.setup(&SetupUpdates {
262            compute: Some("cuda".into()),
263            ..Default::default()
264        })
265        .unwrap();
266        let st = eng.setup_status();
267        assert_eq!(st.hf_token, "hf_one");
268        assert_eq!(st.router, "http://127.0.0.1:1");
269        assert_eq!(st.compute, "cuda");
270    }
271
272    #[test]
273    fn setup_invalid_enum_leaves_state() {
274        let mut eng = Engine::new();
275        eng.setup(&SetupUpdates {
276            compute: Some("cpu".into()),
277            ..Default::default()
278        })
279        .unwrap();
280        assert!(eng
281            .setup(&SetupUpdates {
282                compute: Some("gpu".into()),
283                ..Default::default()
284            })
285            .is_err());
286        assert_eq!(eng.setup_status().compute, "cpu");
287    }
288
289    #[test]
290    fn setup_clear_resets_instance() {
291        let mut eng = Engine::new();
292        eng.setup(&SetupUpdates {
293            hf_token: Some("hf_x".into()),
294            compute: Some("cpu".into()),
295            ..Default::default()
296        })
297        .unwrap();
298        eng.setup_clear();
299        let st = eng.setup_status();
300        assert_eq!(st.hf_token, "");
301        assert_eq!(st.compute, "auto");
302    }
303
304    #[test]
305    fn setup_fills_urls_from_site_tld() {
306        let mut eng = Engine::new();
307        eng.setup(&SetupUpdates {
308            site_url: Some("https://ariacompute.cn".into()),
309            ..Default::default()
310        })
311        .unwrap();
312        let st = eng.setup_status();
313        assert_eq!(st.upgrade_url, crate::setup::CN_UPGRADE);
314    }
315
316    #[test]
317    fn setup_does_not_write_engine_yml() {
318        let _guard = crate::download::ENV_LOCK.lock().unwrap();
319        let home = tempfile::tempdir().unwrap();
320        std::env::set_var("ARIA_COMPUTE_HOME", home.path());
321        let mut eng = Engine::new();
322        eng.setup(&SetupUpdates {
323            router: Some("http://127.0.0.1:8080".into()),
324            site_url: Some("https://ariacompute.com".into()),
325            hf_token: Some("hf_x".into()),
326            ..Default::default()
327        })
328        .unwrap();
329        assert!(!home.path().join("engine.yml").is_file());
330        assert!(!home.path().join("config.yml").is_file());
331        std::env::remove_var("ARIA_COMPUTE_HOME");
332    }
333}