Skip to main content

embacle_mcp/
state.rs

1// ABOUTME: Shared server state holding active provider, model, and multiplex configuration
2// ABOUTME: Thread-safe via Arc<RwLock> with lazy runner creation on first use
3//
4// SPDX-License-Identifier: Apache-2.0
5// Copyright (c) 2026 dravr.ai
6
7use 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
16/// Type alias for the shared state handle used across the server.
17///
18/// `dravr-tronc` 0.5 shares state as `Arc<S>`, so the mutable provider/model
19/// configuration lives behind a per-field interior `RwLock` rather than an
20/// outer one.
21pub type SharedState = Arc<ServerState>;
22
23/// Interior-mutable provider configuration: the active provider, the optional
24/// model override, and the multiplex fan-out list, co-locked so a provider
25/// switch can atomically reset the model.
26struct ActiveConfig {
27    provider: CliRunnerType,
28    model: Option<String>,
29    multiplex: Vec<CliRunnerType>,
30}
31
32/// Central server state tracking provider configuration and cached runners
33///
34/// Runners are created lazily on first access and cached for reuse.
35/// The active provider and model determine how prompt dispatch behaves.
36pub struct ServerState {
37    config: RwLock<ActiveConfig>,
38    runners: Mutex<HashMap<CliRunnerType, Arc<dyn LlmProvider>>>,
39}
40
41impl ServerState {
42    /// Create server state with the given default provider
43    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    /// Get the currently active provider type
55    pub async fn active_provider(&self) -> CliRunnerType {
56        self.config.read().await.provider
57    }
58
59    /// Switch the active provider (resets the active model)
60    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    /// Get the currently selected model (None means use provider default)
67    pub async fn active_model(&self) -> Option<String> {
68        self.config.read().await.model.clone()
69    }
70
71    /// Set the model to use for subsequent requests
72    pub async fn set_active_model(&self, model: Option<String>) {
73        self.config.write().await.model = model;
74    }
75
76    /// Get the list of providers configured for multiplex dispatch
77    pub async fn multiplex_providers(&self) -> Vec<CliRunnerType> {
78        self.config.read().await.multiplex.clone()
79    }
80
81    /// Set the providers used when multiplexing prompts
82    pub async fn set_multiplex_providers(&self, providers: Vec<CliRunnerType>) {
83        self.config.write().await.multiplex = providers;
84    }
85
86    /// Get or lazily create a runner for the given provider type
87    ///
88    /// Created runners are cached for future calls. The runner cache uses
89    /// interior mutability so callers only need `&self`.
90    pub async fn get_runner(
91        &self,
92        provider: CliRunnerType,
93    ) -> Result<Arc<dyn LlmProvider>, RunnerError> {
94        // Fast path: check cache under lock
95        {
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        // Slow path: create runner without holding the lock
103        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}