use crate::types::{AiLibError, ChatCompletionRequest, ChatCompletionResponse};
use async_trait::async_trait;
use futures::stream::Stream;
#[async_trait]
pub trait ChatProvider: Send + Sync {
fn name(&self) -> &str;
async fn chat(
&self,
request: ChatCompletionRequest,
) -> Result<ChatCompletionResponse, AiLibError>;
async fn stream(
&self,
request: ChatCompletionRequest,
) -> Result<
Box<dyn Stream<Item = Result<ChatCompletionChunk, AiLibError>> + Send + Unpin>,
AiLibError,
>;
async fn list_models(&self) -> Result<Vec<String>, AiLibError>;
async fn get_model_info(&self, model_id: &str) -> Result<ModelInfo, AiLibError>;
async fn batch(
&self,
requests: Vec<ChatCompletionRequest>,
concurrency_limit: Option<usize>,
) -> Result<Vec<Result<ChatCompletionResponse, AiLibError>>, AiLibError> {
batch_utils::process_batch_concurrent(self, requests, concurrency_limit).await
}
}
pub use ChatProvider as ChatApi;
#[derive(Debug, Clone)]
pub struct ChatCompletionChunk {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub choices: Vec<ChoiceDelta>,
}
#[derive(Debug, Clone)]
pub struct ChoiceDelta {
pub index: u32,
pub delta: MessageDelta,
pub finish_reason: Option<String>,
}
#[derive(Debug, Clone)]
pub struct MessageDelta {
pub role: Option<Role>,
pub content: Option<String>,
}
#[derive(Debug, Clone)]
pub struct ModelInfo {
pub id: String,
pub object: String,
pub created: u64,
pub owned_by: String,
pub permission: Vec<ModelPermission>,
}
#[derive(Debug, Clone)]
pub struct ModelPermission {
pub id: String,
pub object: String,
pub created: u64,
pub allow_create_engine: bool,
pub allow_sampling: bool,
pub allow_logprobs: bool,
pub allow_search_indices: bool,
pub allow_view: bool,
pub allow_fine_tuning: bool,
pub organization: String,
pub group: Option<String>,
pub is_blocking: bool,
}
use crate::types::Role;
#[derive(Debug)]
pub struct BatchResult {
pub successful: Vec<ChatCompletionResponse>,
pub failed: Vec<(usize, AiLibError)>,
pub total_requests: usize,
pub total_successful: usize,
pub total_failed: usize,
}
impl BatchResult {
pub fn new(total_requests: usize) -> Self {
Self {
successful: Vec::new(),
failed: Vec::new(),
total_requests,
total_successful: 0,
total_failed: 0,
}
}
pub fn add_success(&mut self, response: ChatCompletionResponse) {
self.successful.push(response);
self.total_successful += 1;
}
pub fn add_failure(&mut self, index: usize, error: AiLibError) {
self.failed.push((index, error));
self.total_failed += 1;
}
pub fn all_successful(&self) -> bool {
self.total_failed == 0
}
pub fn success_rate(&self) -> f64 {
if self.total_requests == 0 {
0.0
} else {
(self.total_successful as f64 / self.total_requests as f64) * 100.0
}
}
}
pub mod batch_utils {
use super::*;
use futures::stream::{self, StreamExt};
use std::sync::Arc;
use tokio::sync::Semaphore;
pub async fn process_batch_concurrent<T: ChatProvider + ?Sized>(
api: &T,
requests: Vec<ChatCompletionRequest>,
concurrency_limit: Option<usize>,
) -> Result<Vec<Result<ChatCompletionResponse, AiLibError>>, AiLibError> {
if requests.is_empty() {
return Ok(Vec::new());
}
let semaphore = concurrency_limit.map(|limit| Arc::new(Semaphore::new(limit)));
let futures = requests.into_iter().enumerate().map(|(index, request)| {
let api_ref = api;
let semaphore_ref = semaphore.clone();
async move {
let _permit = if let Some(sem) = &semaphore_ref {
match sem.acquire().await {
Ok(permit) => Some(permit),
Err(_) => {
return (
index,
Err(AiLibError::ProviderError(
"Failed to acquire semaphore permit".to_string(),
)),
)
}
}
} else {
None
};
let result = api_ref.chat(request).await;
(index, result)
}
});
let results: Vec<_> = stream::iter(futures)
.buffer_unordered(concurrency_limit.unwrap_or(usize::MAX))
.collect()
.await;
let mut sorted_results = Vec::with_capacity(results.len());
sorted_results.resize_with(results.len(), || {
Err(AiLibError::ProviderError("Placeholder".to_string()))
});
for (index, result) in results {
sorted_results[index] = result;
}
Ok(sorted_results)
}
pub async fn process_batch_sequential<T: ChatProvider + ?Sized>(
api: &T,
requests: Vec<ChatCompletionRequest>,
) -> Result<Vec<Result<ChatCompletionResponse, AiLibError>>, AiLibError> {
let mut results = Vec::with_capacity(requests.len());
for request in requests {
let result = api.chat(request).await;
results.push(result);
}
Ok(results)
}
pub async fn process_batch_smart<T: ChatProvider + ?Sized>(
api: &T,
requests: Vec<ChatCompletionRequest>,
concurrency_limit: Option<usize>,
) -> Result<Vec<Result<ChatCompletionResponse, AiLibError>>, AiLibError> {
let request_count = requests.len();
if request_count <= 3 {
return process_batch_sequential(api, requests).await;
}
process_batch_concurrent(api, requests, concurrency_limit).await
}
}