use std::sync::Arc;
use std::time::Instant;
use embacle::config::CliRunnerType;
use embacle::types::{ChatMessage, ChatRequest, ResponseFormat, RunnerError};
use crate::state::SharedState;
#[derive(Debug, Clone, Default)]
pub struct MultiplexParams {
pub temperature: Option<f32>,
pub max_tokens: Option<u32>,
pub top_p: Option<f32>,
pub stop: Option<Vec<String>>,
pub response_format: Option<ResponseFormat>,
}
#[derive(Debug)]
pub struct MultiplexResult {
pub responses: Vec<ProviderResponse>,
pub summary: String,
}
#[derive(Debug)]
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 fn new(state: &SharedState) -> Self {
Self {
state: Arc::clone(state),
}
}
pub async fn execute(
&self,
messages: &[ChatMessage],
providers: &[CliRunnerType],
params: &MultiplexParams,
) -> 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();
let params = params.clone();
handles.push(tokio::spawn(async move {
dispatch_single(state, provider, messages, ¶ms).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>,
params: &MultiplexParams,
) -> 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 mut request = ChatRequest::new(messages);
request.temperature = params.temperature;
request.max_tokens = params.max_tokens;
request.top_p = params.top_p;
request.stop.clone_from(¶ms.stop);
request.response_format.clone_from(¶ms.response_format);
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)
}