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<RwLock<ServerState>>;
pub struct ServerState {
active_provider: CliRunnerType,
active_model: Option<String>,
multiplex_providers: Vec<CliRunnerType>,
runners: Mutex<HashMap<CliRunnerType, Arc<dyn LlmProvider>>>,
}
impl ServerState {
pub fn new(default_provider: CliRunnerType) -> Self {
Self {
active_provider: default_provider,
active_model: None,
multiplex_providers: Vec::new(),
runners: Mutex::new(HashMap::new()),
}
}
pub const fn active_provider(&self) -> CliRunnerType {
self.active_provider
}
pub fn set_active_provider(&mut self, provider: CliRunnerType) {
self.active_provider = provider;
self.active_model = None;
}
pub fn active_model(&self) -> Option<&str> {
self.active_model.as_deref()
}
pub fn set_active_model(&mut self, model: Option<String>) {
self.active_model = model;
}
pub fn multiplex_providers(&self) -> &[CliRunnerType] {
&self.multiplex_providers
}
pub fn set_multiplex_providers(&mut self, providers: Vec<CliRunnerType>) {
self.multiplex_providers = 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::*;
#[test]
fn default_state_uses_provided_provider() {
let state = ServerState::new(CliRunnerType::Copilot);
assert_eq!(state.active_provider(), CliRunnerType::Copilot);
assert!(state.active_model().is_none());
assert!(state.multiplex_providers().is_empty());
}
#[test]
fn set_provider_resets_model() {
let mut state = ServerState::new(CliRunnerType::Copilot);
state.set_active_model(Some("gpt-4o".to_owned()));
assert_eq!(state.active_model(), Some("gpt-4o"));
state.set_active_provider(CliRunnerType::ClaudeCode);
assert_eq!(state.active_provider(), CliRunnerType::ClaudeCode);
assert!(state.active_model().is_none());
}
#[test]
fn multiplex_providers_round_trip() {
let mut state = ServerState::new(CliRunnerType::Copilot);
let providers = vec![CliRunnerType::ClaudeCode, CliRunnerType::OpenCode];
state.set_multiplex_providers(providers.clone());
assert_eq!(state.multiplex_providers(), &providers);
}
}