Skip to main content

chatty_rs/backend/
manager.rs

1#[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>, /* Alias - Backend */
14    models: HashMap<String, String>,          /* Model ID - Alias  */
15}
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}