use std::collections::HashMap;
use std::sync::Arc;
use embacle::config::CliRunnerType;
use embacle::types::{LlmProvider, RunnerError};
use tokio::sync::{Mutex, RwLock};
use crate::runner::factory;
pub type SharedState = Arc<ServerState>;
struct ActiveConfig {
provider: CliRunnerType,
model: Option<String>,
multiplex: Vec<CliRunnerType>,
}
pub struct ServerState {
config: RwLock<ActiveConfig>,
runners: Mutex<HashMap<CliRunnerType, Arc<dyn LlmProvider>>>,
}
impl ServerState {
pub fn new(default_provider: CliRunnerType) -> Self {
Self {
config: RwLock::new(ActiveConfig {
provider: default_provider,
model: None,
multiplex: Vec::new(),
}),
runners: Mutex::new(HashMap::new()),
}
}
pub async fn active_provider(&self) -> CliRunnerType {
self.config.read().await.provider
}
pub async fn set_active_provider(&self, provider: CliRunnerType) {
let mut config = self.config.write().await;
config.provider = provider;
config.model = None;
}
pub async fn active_model(&self) -> Option<String> {
self.config.read().await.model.clone()
}
pub async fn set_active_model(&self, model: Option<String>) {
self.config.write().await.model = model;
}
pub async fn multiplex_providers(&self) -> Vec<CliRunnerType> {
self.config.read().await.multiplex.clone()
}
pub async fn set_multiplex_providers(&self, providers: Vec<CliRunnerType>) {
self.config.write().await.multiplex = providers;
}
pub async fn get_runner(
&self,
provider: CliRunnerType,
) -> Result<Arc<dyn LlmProvider>, RunnerError> {
{
let runners = self.runners.lock().await;
if let Some(runner) = runners.get(&provider) {
return Ok(Arc::clone(runner));
}
}
let runner = factory::create_runner(provider).await?;
let runner: Arc<dyn LlmProvider> = Arc::from(runner);
let runner = self
.runners
.lock()
.await
.entry(provider)
.or_insert_with(|| runner)
.clone();
Ok(runner)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn default_state_uses_provided_provider() {
let state = ServerState::new(CliRunnerType::Copilot);
assert_eq!(state.active_provider().await, CliRunnerType::Copilot);
assert!(state.active_model().await.is_none());
assert!(state.multiplex_providers().await.is_empty());
}
#[tokio::test]
async fn set_provider_resets_model() {
let state = ServerState::new(CliRunnerType::Copilot);
state.set_active_model(Some("gpt-4o".to_owned())).await;
assert_eq!(state.active_model().await, Some("gpt-4o".to_owned()));
state.set_active_provider(CliRunnerType::ClaudeCode).await;
assert_eq!(state.active_provider().await, CliRunnerType::ClaudeCode);
assert!(state.active_model().await.is_none());
}
#[tokio::test]
async fn multiplex_providers_round_trip() {
let state = ServerState::new(CliRunnerType::Copilot);
let providers = vec![CliRunnerType::ClaudeCode, CliRunnerType::OpenCode];
state.set_multiplex_providers(providers.clone()).await;
assert_eq!(state.multiplex_providers().await, providers);
}
}