use std::collections::{BTreeMap, BTreeSet};
use std::sync::Arc;
use af_llm::{CompletionRequest, CompletionResponse, LlmClient, LlmError};
use async_trait::async_trait;
use tokio::sync::mpsc::UnboundedSender;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ModelDescriptor {
pub id: String,
pub provider: String,
#[serde(default)]
pub reasoning_efforts: BTreeSet<String>,
#[serde(default)]
pub input_modalities: BTreeSet<String>,
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum ModelRegistryError {
#[error("model descriptor requires id and provider")]
InvalidDescriptor,
#[error("model is already registered: {0}")]
Duplicate(String),
#[error("model is not registered: {0}")]
Unknown(String),
#[error("preferred model is not allowed: {0}")]
NotAllowed(String),
}
#[derive(Default)]
pub struct ModelRegistry {
models: BTreeMap<String, (ModelDescriptor, Arc<dyn ChatModel>)>,
}
impl ModelRegistry {
pub fn register(
&mut self,
descriptor: ModelDescriptor,
model: Arc<dyn ChatModel>,
) -> Result<(), ModelRegistryError> {
if descriptor.id.trim().is_empty() || descriptor.provider.trim().is_empty() {
return Err(ModelRegistryError::InvalidDescriptor);
}
if self.models.contains_key(&descriptor.id) {
return Err(ModelRegistryError::Duplicate(descriptor.id));
}
self.models
.insert(descriptor.id.clone(), (descriptor, model));
Ok(())
}
pub fn descriptors(&self) -> impl Iterator<Item = &ModelDescriptor> {
self.models.values().map(|(descriptor, _)| descriptor)
}
pub fn resolve(
&self,
preferred: &str,
allowed: &BTreeSet<String>,
) -> Result<Arc<dyn ChatModel>, ModelRegistryError> {
if !allowed.contains(preferred) {
return Err(ModelRegistryError::NotAllowed(preferred.into()));
}
self.models
.get(preferred)
.map(|(_, model)| Arc::clone(model))
.ok_or_else(|| ModelRegistryError::Unknown(preferred.into()))
}
}
#[async_trait]
pub trait ChatModel: Send + Sync {
async fn complete_streaming(
&self,
request: &CompletionRequest,
delta_tx: UnboundedSender<(String, bool)>,
) -> Result<CompletionResponse, LlmError>;
}
#[async_trait]
impl ChatModel for LlmClient {
async fn complete_streaming(
&self,
request: &CompletionRequest,
delta_tx: UnboundedSender<(String, bool)>,
) -> Result<CompletionResponse, LlmError> {
self.complete_stream_single_attempt(request, |content, has_tools| {
let _ = delta_tx.send((content.to_string(), has_tools));
})
.await
}
}
#[async_trait]
impl ChatModel for Arc<dyn ChatModel> {
async fn complete_streaming(
&self,
request: &CompletionRequest,
delta_tx: UnboundedSender<(String, bool)>,
) -> Result<CompletionResponse, LlmError> {
(**self).complete_streaming(request, delta_tx).await
}
}
#[cfg(test)]
mod registry_tests {
use super::*;
use af_llm::{ChatMessage, Choice, CompletionResponse, LlmError};
struct Fixed;
#[async_trait]
impl ChatModel for Fixed {
async fn complete_streaming(
&self,
_: &CompletionRequest,
_: UnboundedSender<(String, bool)>,
) -> Result<CompletionResponse, LlmError> {
Ok(CompletionResponse {
id: "fixed".into(),
choices: vec![Choice {
index: 0,
message: ChatMessage::assistant("ok"),
output_blocks: Vec::new(),
finish_reason: None,
}],
usage: None,
})
}
}
fn descriptor(id: &str) -> ModelDescriptor {
ModelDescriptor {
id: id.into(),
provider: "test".into(),
reasoning_efforts: BTreeSet::from(["low".into()]),
input_modalities: BTreeSet::from(["text".into()]),
}
}
#[test]
fn registry_is_exact_sorted_and_rejects_duplicates() {
let mut registry = ModelRegistry::default();
registry.register(descriptor("z"), Arc::new(Fixed)).unwrap();
registry.register(descriptor("a"), Arc::new(Fixed)).unwrap();
assert_eq!(
registry
.descriptors()
.map(|value| value.id.as_str())
.collect::<Vec<_>>(),
vec!["a", "z"]
);
assert_eq!(
registry
.register(descriptor("a"), Arc::new(Fixed))
.unwrap_err(),
ModelRegistryError::Duplicate("a".into())
);
assert!(registry.resolve("a", &BTreeSet::from(["a".into()])).is_ok());
assert!(matches!(
registry.resolve("z", &BTreeSet::from(["a".into()])),
Err(ModelRegistryError::NotAllowed(value)) if value == "z"
));
}
}