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_ffi.dll"
214        } else if cfg!(target_os = "macos") {
215            "libaria_ffi.dylib"
216        } else {
217            "libaria_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            site_url: Some(crate::setup::CN_SITE.into()),
236            upgrade_url: Some(crate::setup::CN_UPGRADE.into()),
237            compute: Some("cpu".into()),
238            hf_token: Some("hf_abc".into()),
239            modelscope_api_token: Some("ms_xyz".into()),
240        })
241        .unwrap();
242        let st = eng.setup_status();
243        assert_eq!(st.router, "http://127.0.0.1:8080");
244        assert_eq!(st.compute, "cpu");
245        assert_eq!(st.hf_token, "hf_abc");
246        assert_eq!(st.modelscope_api_token, "ms_xyz");
247        assert_eq!(st.site_url, crate::setup::CN_SITE);
248    }
249
250    #[test]
251    fn setup_partial_merge() {
252        let mut eng = Engine::new();
253        eng.setup(&SetupUpdates {
254            hf_token: Some("hf_one".into()),
255            router: Some("http://127.0.0.1:1".into()),
256            ..Default::default()
257        })
258        .unwrap();
259        eng.setup(&SetupUpdates {
260            compute: Some("cuda".into()),
261            ..Default::default()
262        })
263        .unwrap();
264        let st = eng.setup_status();
265        assert_eq!(st.hf_token, "hf_one");
266        assert_eq!(st.router, "http://127.0.0.1:1");
267        assert_eq!(st.compute, "cuda");
268    }
269
270    #[test]
271    fn setup_invalid_enum_leaves_state() {
272        let mut eng = Engine::new();
273        eng.setup(&SetupUpdates {
274            compute: Some("cpu".into()),
275            ..Default::default()
276        })
277        .unwrap();
278        assert!(eng
279            .setup(&SetupUpdates {
280                compute: Some("gpu".into()),
281                ..Default::default()
282            })
283            .is_err());
284        assert_eq!(eng.setup_status().compute, "cpu");
285    }
286
287    #[test]
288    fn setup_clear_resets_instance() {
289        let mut eng = Engine::new();
290        eng.setup(&SetupUpdates {
291            hf_token: Some("hf_x".into()),
292            compute: Some("cpu".into()),
293            ..Default::default()
294        })
295        .unwrap();
296        eng.setup_clear();
297        let st = eng.setup_status();
298        assert_eq!(st.hf_token, "");
299        assert_eq!(st.compute, "auto");
300    }
301
302    #[test]
303    fn setup_fills_urls_from_site_tld() {
304        let mut eng = Engine::new();
305        eng.setup(&SetupUpdates {
306            site_url: Some("https://ariacompute.cn".into()),
307            ..Default::default()
308        })
309        .unwrap();
310        let st = eng.setup_status();
311        assert_eq!(st.upgrade_url, crate::setup::CN_UPGRADE);
312    }
313
314    #[test]
315    fn setup_does_not_write_engine_yml() {
316        let _guard = crate::download::ENV_LOCK.lock().unwrap();
317        let home = tempfile::tempdir().unwrap();
318        std::env::set_var("ARIA_COMPUTE_HOME", home.path());
319        let mut eng = Engine::new();
320        eng.setup(&SetupUpdates {
321            router: Some("http://127.0.0.1:8080".into()),
322            site_url: Some("https://ariacompute.com".into()),
323            hf_token: Some("hf_x".into()),
324            ..Default::default()
325        })
326        .unwrap();
327        assert!(!home.path().join("engine.yml").is_file());
328        assert!(!home.path().join("config.yml").is_file());
329        std::env::remove_var("ARIA_COMPUTE_HOME");
330    }
331}