use std::collections::HashMap;
use std::sync::Arc;
use ai_sdk_provider::{EmbeddingModel, ImageModel, LanguageModel, ProviderV3};
pub struct CustomProviderBuilder {
language_models: HashMap<String, Arc<dyn LanguageModel>>,
embedding_models: HashMap<String, Arc<dyn EmbeddingModel<String>>>,
image_models: HashMap<String, Arc<dyn ImageModel>>,
fallback_provider: Option<Arc<dyn ProviderV3>>,
}
impl CustomProviderBuilder {
pub fn new() -> Self {
Self {
language_models: HashMap::new(),
embedding_models: HashMap::new(),
image_models: HashMap::new(),
fallback_provider: None,
}
}
pub fn language_model(
mut self,
model_id: impl Into<String>,
model: Arc<dyn LanguageModel>,
) -> Self {
self.language_models.insert(model_id.into(), model);
self
}
pub fn text_embedding_model(
mut self,
model_id: impl Into<String>,
model: Arc<dyn EmbeddingModel<String>>,
) -> Self {
self.embedding_models.insert(model_id.into(), model);
self
}
pub fn image_model(mut self, model_id: impl Into<String>, model: Arc<dyn ImageModel>) -> Self {
self.image_models.insert(model_id.into(), model);
self
}
pub fn fallback_provider(mut self, provider: Arc<dyn ProviderV3>) -> Self {
self.fallback_provider = Some(provider);
self
}
pub fn build(self) -> Arc<dyn ProviderV3> {
Arc::new(CustomProvider {
language_models: self.language_models,
embedding_models: self.embedding_models,
image_models: self.image_models,
fallback_provider: self.fallback_provider,
})
}
}
impl Default for CustomProviderBuilder {
fn default() -> Self {
Self::new()
}
}
struct CustomProvider {
language_models: HashMap<String, Arc<dyn LanguageModel>>,
embedding_models: HashMap<String, Arc<dyn EmbeddingModel<String>>>,
image_models: HashMap<String, Arc<dyn ImageModel>>,
fallback_provider: Option<Arc<dyn ProviderV3>>,
}
impl ProviderV3 for CustomProvider {
fn specification_version(&self) -> &str {
"v3"
}
fn language_model(&self, model_id: &str) -> Option<Arc<dyn LanguageModel>> {
if let Some(model) = self.language_models.get(model_id) {
return Some(model.clone());
}
if let Some(fallback) = &self.fallback_provider {
return fallback.language_model(model_id);
}
None
}
fn text_embedding_model(&self, model_id: &str) -> Option<Arc<dyn EmbeddingModel<String>>> {
if let Some(model) = self.embedding_models.get(model_id) {
return Some(model.clone());
}
if let Some(fallback) = &self.fallback_provider {
return fallback.text_embedding_model(model_id);
}
None
}
fn image_model(&self, model_id: &str) -> Option<Arc<dyn ImageModel>> {
if let Some(model) = self.image_models.get(model_id) {
return Some(model.clone());
}
if let Some(fallback) = &self.fallback_provider {
return fallback.image_model(model_id);
}
None
}
}
#[cfg(test)]
mod tests {
use super::*;
struct MockProvider;
impl ProviderV3 for MockProvider {
fn language_model(&self, model_id: &str) -> Option<Arc<dyn LanguageModel>> {
if model_id == "fallback-model" {
None
} else {
None
}
}
fn text_embedding_model(&self, _model_id: &str) -> Option<Arc<dyn EmbeddingModel<String>>> {
None
}
fn image_model(&self, _model_id: &str) -> Option<Arc<dyn ImageModel>> {
None
}
}
#[test]
fn test_custom_provider_builder() {
let builder = CustomProviderBuilder::new();
let provider = builder.build();
assert!(provider.language_model("nonexistent").is_none());
}
#[test]
fn test_custom_provider_with_fallback() {
let fallback = Arc::new(MockProvider);
let provider = CustomProviderBuilder::new()
.fallback_provider(fallback)
.build();
let _ = provider.language_model("fallback-model");
}
#[test]
fn test_specification_version() {
let provider = CustomProviderBuilder::new().build();
assert_eq!(provider.specification_version(), "v3");
}
}