chatty_rs/backend/
manager.rs1#[cfg(test)]
2#[path = "manager_test.rs"]
3mod tests;
4
5use crate::backend::{ArcBackend, Backend};
6use crate::models::{ArcEventTx, BackendPrompt, Model};
7use async_trait::async_trait;
8use eyre::{Context, Result, bail};
9use std::collections::HashMap;
10
11#[derive(Default)]
12pub struct Manager {
13 connections: HashMap<String, ArcBackend>, models: HashMap<String, String>, }
16
17impl Manager {
18 pub fn len(&self) -> usize {
19 self.connections.len()
20 }
21
22 pub fn is_empty(&self) -> bool {
23 self.connections.is_empty()
24 }
25
26 pub async fn add_connection(&mut self, connection: ArcBackend) -> eyre::Result<()> {
27 let alias = connection.name().to_string();
28
29 if self.connections.contains_key(&alias) {
30 bail!(format!("connection {} already exists", alias))
31 }
32
33 connection
34 .list_models()
35 .await
36 .wrap_err(format!("listing models backend {}", alias))?
37 .into_iter()
38 .for_each(|m| {
39 self.models.insert(m.id().to_string(), alias.clone());
40 });
41
42 self.connections.insert(alias, connection);
43 Ok(())
44 }
45
46 pub fn get_connection(&self, model: &str) -> Option<&ArcBackend> {
47 self.connections.get(self.models.get(model)?)
48 }
49}
50
51#[async_trait]
52impl Backend for Manager {
53 fn name(&self) -> &str {
54 "Manager"
55 }
56
57 async fn list_models(&self) -> Result<Vec<Model>> {
58 Ok(self
59 .models
60 .iter()
61 .map(|(id, alias)| Model::new(id).with_provider(alias))
62 .collect())
63 }
64
65 async fn get_completion(&self, prompt: BackendPrompt, event_tx: ArcEventTx) -> Result<()> {
66 let connection = match self.get_connection(prompt.model()) {
67 Some(connection) => connection,
68 None => {
69 return Err(eyre::eyre!("model is not available"));
70 }
71 };
72 connection
73 .get_completion(prompt, event_tx)
74 .await
75 .wrap_err(format!("get completion from backend {}", connection.name()))?;
76 Ok(())
77 }
78}