use std::collections::HashMap;
use std::sync::Arc;
use ai_sdk_provider::ProviderV3;
use super::provider_registry::{DefaultProviderRegistry, ProviderRegistry};
pub fn create_provider_registry(separator: Option<&str>) -> ProviderRegistryBuilder {
ProviderRegistryBuilder {
separator: separator.unwrap_or(":").to_string(),
providers: HashMap::new(),
}
}
pub struct ProviderRegistryBuilder {
separator: String,
providers: HashMap<String, Arc<dyn ProviderV3>>,
}
impl ProviderRegistryBuilder {
pub fn new() -> Self {
Self {
separator: ":".to_string(),
providers: HashMap::new(),
}
}
pub fn with_provider(mut self, id: impl Into<String>, provider: Arc<dyn ProviderV3>) -> Self {
self.providers.insert(id.into(), provider);
self
}
pub fn build(self) -> DefaultProviderRegistry {
let mut registry = DefaultProviderRegistry::new(self.separator);
for (id, provider) in self.providers {
registry.register_provider(id, provider);
}
registry
}
}
impl Default for ProviderRegistryBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use ai_sdk_provider::{EmbeddingModel, ImageModel, LanguageModel, ProviderV3};
struct MockProvider;
impl ProviderV3 for MockProvider {
fn language_model(&self, _model_id: &str) -> Option<Arc<dyn LanguageModel>> {
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_builder_creation() {
let builder = ProviderRegistryBuilder::new();
assert_eq!(builder.separator, ":");
assert_eq!(builder.providers.len(), 0);
}
#[test]
fn test_builder_with_provider() {
let provider = Arc::new(MockProvider);
let builder = ProviderRegistryBuilder::new().with_provider("test", provider);
assert_eq!(builder.providers.len(), 1);
assert!(builder.providers.contains_key("test"));
}
#[test]
fn test_create_provider_registry() {
let provider = Arc::new(MockProvider);
let registry = create_provider_registry(Some(":"))
.with_provider("test", provider)
.build();
let providers = registry.list_providers();
assert_eq!(providers.len(), 1);
assert!(providers.contains(&"test".to_string()));
}
#[test]
fn test_default_separator() {
let registry = create_provider_registry(None).build();
assert_eq!(registry.list_providers().len(), 0);
}
}