1use std::collections::HashMap;
8use std::sync::Arc;
9
10use embacle::config::CliRunnerType;
11use embacle::types::{LlmProvider, RunnerError};
12use tokio::sync::{Mutex, RwLock};
13
14use crate::runner::factory;
15
16pub type SharedState = Arc<ServerState>;
22
23struct ActiveConfig {
27 provider: CliRunnerType,
28 model: Option<String>,
29 multiplex: Vec<CliRunnerType>,
30}
31
32pub struct ServerState {
37 config: RwLock<ActiveConfig>,
38 runners: Mutex<HashMap<CliRunnerType, Arc<dyn LlmProvider>>>,
39}
40
41impl ServerState {
42 pub fn new(default_provider: CliRunnerType) -> Self {
44 Self {
45 config: RwLock::new(ActiveConfig {
46 provider: default_provider,
47 model: None,
48 multiplex: Vec::new(),
49 }),
50 runners: Mutex::new(HashMap::new()),
51 }
52 }
53
54 pub async fn active_provider(&self) -> CliRunnerType {
56 self.config.read().await.provider
57 }
58
59 pub async fn set_active_provider(&self, provider: CliRunnerType) {
61 let mut config = self.config.write().await;
62 config.provider = provider;
63 config.model = None;
64 }
65
66 pub async fn active_model(&self) -> Option<String> {
68 self.config.read().await.model.clone()
69 }
70
71 pub async fn set_active_model(&self, model: Option<String>) {
73 self.config.write().await.model = model;
74 }
75
76 pub async fn multiplex_providers(&self) -> Vec<CliRunnerType> {
78 self.config.read().await.multiplex.clone()
79 }
80
81 pub async fn set_multiplex_providers(&self, providers: Vec<CliRunnerType>) {
83 self.config.write().await.multiplex = providers;
84 }
85
86 pub async fn get_runner(
91 &self,
92 provider: CliRunnerType,
93 ) -> Result<Arc<dyn LlmProvider>, RunnerError> {
94 {
96 let runners = self.runners.lock().await;
97 if let Some(runner) = runners.get(&provider) {
98 return Ok(Arc::clone(runner));
99 }
100 }
101
102 let runner = factory::create_runner(provider).await?;
104 let runner: Arc<dyn LlmProvider> = Arc::from(runner);
105
106 let runner = self
107 .runners
108 .lock()
109 .await
110 .entry(provider)
111 .or_insert_with(|| runner)
112 .clone();
113 Ok(runner)
114 }
115}
116
117#[cfg(test)]
118mod tests {
119 use super::*;
120
121 #[tokio::test]
122 async fn default_state_uses_provided_provider() {
123 let state = ServerState::new(CliRunnerType::Copilot);
124 assert_eq!(state.active_provider().await, CliRunnerType::Copilot);
125 assert!(state.active_model().await.is_none());
126 assert!(state.multiplex_providers().await.is_empty());
127 }
128
129 #[tokio::test]
130 async fn set_provider_resets_model() {
131 let state = ServerState::new(CliRunnerType::Copilot);
132 state.set_active_model(Some("gpt-4o".to_owned())).await;
133 assert_eq!(state.active_model().await, Some("gpt-4o".to_owned()));
134
135 state.set_active_provider(CliRunnerType::ClaudeCode).await;
136 assert_eq!(state.active_provider().await, CliRunnerType::ClaudeCode);
137 assert!(state.active_model().await.is_none());
138 }
139
140 #[tokio::test]
141 async fn multiplex_providers_round_trip() {
142 let state = ServerState::new(CliRunnerType::Copilot);
143 let providers = vec![CliRunnerType::ClaudeCode, CliRunnerType::OpenCode];
144 state.set_multiplex_providers(providers.clone()).await;
145 assert_eq!(state.multiplex_providers().await, providers);
146 }
147}