use std::sync::Arc;
use std::time::Instant;
use embacle::config::CliRunnerType;
use embacle::types::{ChatMessage, ChatRequest, RunnerError};
use serde::Serialize;
use crate::state::SharedState;
#[derive(Debug, Serialize)]
pub struct MultiplexResult {
pub responses: Vec<ProviderResponse>,
pub summary: String,
}
#[derive(Debug, Serialize)]
pub struct ProviderResponse {
pub provider: String,
pub content: Option<String>,
pub model: Option<String>,
pub error: Option<String>,
pub duration_ms: u64,
}
pub struct MultiplexEngine {
state: SharedState,
}
impl MultiplexEngine {
pub const fn new(state: SharedState) -> Self {
Self { state }
}
pub async fn execute(
&self,
messages: &[ChatMessage],
providers: &[CliRunnerType],
) -> Result<MultiplexResult, RunnerError> {
let mut handles = Vec::with_capacity(providers.len());
for &provider in providers {
let state = Arc::clone(&self.state);
let messages = messages.to_vec();
handles.push(tokio::spawn(async move {
dispatch_single(state, provider, messages).await
}));
}
let mut responses = Vec::with_capacity(handles.len());
for handle in handles {
match handle.await {
Ok(resp) => responses.push(resp),
Err(e) => responses.push(ProviderResponse {
provider: "unknown".to_owned(),
content: None,
model: None,
error: Some(format!("Task join error: {e}")),
duration_ms: 0,
}),
}
}
let summary = build_summary(&responses);
Ok(MultiplexResult { responses, summary })
}
}
async fn dispatch_single(
state: SharedState,
provider: CliRunnerType,
messages: Vec<ChatMessage>,
) -> ProviderResponse {
let start = Instant::now();
let runner = state.get_runner(provider).await;
let runner = match runner {
Ok(r) => r,
Err(e) => {
return ProviderResponse {
provider: provider.to_string(),
content: None,
model: None,
error: Some(e.to_string()),
duration_ms: elapsed_ms(start),
};
}
};
let request = ChatRequest::new(messages);
match runner.complete(&request).await {
Ok(response) => ProviderResponse {
provider: provider.to_string(),
content: Some(response.content),
model: Some(response.model),
error: None,
duration_ms: elapsed_ms(start),
},
Err(e) => ProviderResponse {
provider: provider.to_string(),
content: None,
model: None,
error: Some(e.to_string()),
duration_ms: elapsed_ms(start),
},
}
}
fn build_summary(responses: &[ProviderResponse]) -> String {
let total = responses.len();
let succeeded = responses.iter().filter(|r| r.content.is_some()).count();
let failed = total - succeeded;
format!("{succeeded} succeeded, {failed} failed out of {total} providers")
}
fn elapsed_ms(start: Instant) -> u64 {
let millis = start.elapsed().as_millis();
u64::try_from(millis).unwrap_or(u64::MAX)
}