use foundation_db::traits::DocumentStore;
use crate::agentic::{AgentSession, AgentSessionBuilder, MemoryStore};
use crate::types::{
ModelId, ModelProvider, ProviderRouter, RoutableProvider, RoutableProviderBox, RoutingRule,
SessionId,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Role {
Primary,
Memory,
Fallback,
}
#[derive(Default)]
pub struct RouterMix {
providers: Vec<Box<dyn RoutableProvider>>,
rules: Vec<RoutingRule>,
primary: Option<ModelId>,
memory: Option<ModelId>,
fallbacks: Vec<ModelId>,
}
impl RouterMix {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn primary<P>(self, provider: P, model: ModelId) -> Self
where
P: ModelProvider + Send + Sync + 'static,
P::Model: Send + Sync,
{
self.add_role(Role::Primary, provider, model)
}
#[must_use]
pub fn memory<P>(self, provider: P, model: ModelId) -> Self
where
P: ModelProvider + Send + Sync + 'static,
P::Model: Send + Sync,
{
self.add_role(Role::Memory, provider, model)
}
#[must_use]
pub fn fallback<P>(self, provider: P, model: ModelId) -> Self
where
P: ModelProvider + Send + Sync + 'static,
P::Model: Send + Sync,
{
self.add_role(Role::Fallback, provider, model)
}
fn add_role<P>(mut self, role: Role, provider: P, model: ModelId) -> Self
where
P: ModelProvider + Send + Sync + 'static,
P::Model: Send + Sync,
{
let provider_id = provider
.describe()
.expect(
"provider must return a descriptor with a provider_id to be \
registered in a RouterMix — a provider without identity cannot be routed",
)
.provider;
let name = model.name().to_string();
let boxed: Box<dyn RoutableProvider> =
Box::new(RoutableProviderBox::with_identity(provider, name.clone(), provider_id));
self.providers.push(boxed);
self.rules.push(RoutingRule {
model: model.clone(),
provider_name: name,
});
match role {
Role::Primary => self.primary = Some(model),
Role::Memory => self.memory = Some(model),
Role::Fallback => self.fallbacks.push(model),
}
self
}
#[must_use]
pub fn build(self) -> RouterPreset {
let mut builder = ProviderRouter::builder();
for provider in self.providers {
builder = builder.add_provider(provider);
}
for rule in self.rules {
builder = builder.rule(rule);
}
RouterPreset {
router: builder.build(),
primary_model: self
.primary
.unwrap_or_else(|| ModelId::Name(String::new(), None)),
memory_model: self.memory,
fallback_models: self.fallbacks,
}
}
}
pub struct RouterPreset {
pub router: ProviderRouter,
pub primary_model: ModelId,
pub memory_model: Option<ModelId>,
pub fallback_models: Vec<ModelId>,
}
impl RouterPreset {
#[must_use]
pub fn into_agent_builder<D, M>(self, session_id: SessionId) -> AgentSessionBuilder<D, M>
where
D: DocumentStore + 'static,
M: MemoryStore + 'static,
{
let mut builder =
AgentSession::<D, M>::builder(session_id, self.router).with_model(self.primary_model);
if let Some(memory) = self.memory_model {
builder = builder.with_memory_model(memory);
}
if !self.fallback_models.is_empty() {
builder = builder.with_fallback_models(self.fallback_models);
}
builder
}
}