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::{GenerateOpts, Generation, Session, SessionBuilder};
8
9pub struct Engine {
11 session: Session,
12}
13
14impl Engine {
15 pub fn open(bundle_path: impl AsRef<std::path::Path>) -> Result<Self, aria_inference::EngineError> {
16 let session = SessionBuilder::new().model(bundle_path).build()?;
17 Ok(Self { session })
18 }
19
20 pub fn complete(&mut self, prompt: &str, opts: &GenerateOpts) -> Result<Generation, aria_inference::EngineError> {
21 let turns = [aria_inference::ChatTurn::new("user", prompt)];
22 let tokens = self.session.encode_chat(&turns);
23 self.session.generate(&tokens, opts)
24 }
25
26 pub fn embed(&self, text: &str) -> Result<Vec<f32>, aria_inference::EngineError> {
27 self.session.embed_text(text)
28 }
29
30 pub fn transcribe(&self, pcm: &[u8]) -> Result<String, aria_inference::EngineError> {
31 self.session.transcribe_pcm16le(pcm)
32 }
33}
34
35#[cfg(test)]
36mod tests {
37 use super::*;
38 use aria_inference::fixture::write_tiny_q4_bundle;
39
40 #[test]
41 fn engine_complete_ok() {
42 let dir = tempfile::tempdir().unwrap();
43 write_tiny_q4_bundle(dir.path()).unwrap();
44 let mut eng = Engine::open(dir.path()).unwrap();
45 let g = eng
46 .complete("hi", &GenerateOpts { max_tokens: 2, temperature: 0.0 })
47 .unwrap();
48 assert!(!g.text.is_empty());
49 assert!(!eng.embed("x").unwrap().is_empty());
50 assert!(!eng.transcribe(&[0, 1, 2, 3]).unwrap().is_empty());
51 }
52}