use std::sync::Arc;
use axum::http::StatusCode;
use axum::Json;
use ferrox_models::EmbeddingModel;
use crate::{budget, serving, ApiError, Model};
pub(crate) enum Loaded {
Generative(Arc<Model>),
Encoder(Arc<EmbeddingModel>),
}
pub(crate) struct ActiveModel {
pub(crate) id: Option<String>,
pub(crate) loaded: Loaded,
pub(crate) batcher: Option<serving::batch::ContinuousBatcher>,
pub(crate) ceiling: Option<Arc<budget::ContextCeiling>>,
}
fn served_name_override() -> Option<&'static str> {
static NAME: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
NAME.get_or_init(|| {
std::env::var("FERROX_MODEL_NAME")
.ok()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
})
.as_deref()
}
impl ActiveModel {
pub(crate) fn generative(&self) -> Result<&Arc<Model>, ApiError> {
match &self.loaded {
Loaded::Generative(m) => Ok(m),
Loaded::Encoder(e) => Err(not_a_generative_model(e)),
}
}
pub(crate) fn generative_opt(&self) -> Option<&Arc<Model>> {
match &self.loaded {
Loaded::Generative(m) => Some(m),
Loaded::Encoder(_) => None,
}
}
pub(crate) fn encoder(&self) -> Option<&Arc<EmbeddingModel>> {
match &self.loaded {
Loaded::Generative(_) => None,
Loaded::Encoder(e) => Some(e),
}
}
pub(crate) fn encode_any(&self, text: &str) -> Vec<usize> {
match &self.loaded {
Loaded::Generative(m) => m.encode(text),
Loaded::Encoder(e) => e.token_ids(text).into_iter().map(|t| t as usize).collect(),
}
}
pub(crate) fn decode_any(&self, ids: &[usize]) -> String {
match &self.loaded {
Loaded::Generative(m) => m.decode(ids),
Loaded::Encoder(e) => {
let ids: Vec<u32> = ids.iter().map(|&i| i as u32).collect();
e.decode_tokens(&ids)
}
}
}
pub(crate) fn name(&self) -> &str {
if let Some(alias) = served_name_override() {
return alias;
}
match &self.loaded {
Loaded::Generative(m) => m.name(),
Loaded::Encoder(e) => e.name(),
}
}
pub(crate) fn tokenizer_kind(&self) -> &'static str {
match &self.loaded {
Loaded::Generative(m) => m.tokenizer_kind(),
Loaded::Encoder(_) => "gguf-wordpiece",
}
}
pub(crate) fn is_synthetic(&self) -> bool {
self.generative_opt().is_some_and(|m| m.is_synthetic())
}
pub(crate) fn expert_store_stats(&self) -> Option<ferrox_core::expert_store::ExpertStoreStats> {
self.generative_opt().and_then(|m| m.expert_store_stats())
}
}
fn not_a_generative_model(encoder: &EmbeddingModel) -> ApiError {
(
StatusCode::NOT_IMPLEMENTED,
Json(serde_json::json!({"error": {
"message": encoder_refusal(
encoder.name(),
encoder.architecture(),
encoder.n_embd(),
encoder.pooling_type().name(),
),
"type": "unsupported",
"param": "model",
}})),
)
}
fn encoder_refusal(name: &str, arch: &str, n_embd: usize, pooling: &str) -> String {
format!(
"the loaded model '{name}' is an embedding model ({arch} encoder, {n_embd} dims, \
pooling {pooling}). An encoder has no output head, so it cannot generate text at \
all. There is no next token for it to predict. POST /v1/embeddings to use it, or \
load a generative checkpoint."
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_encoder_refusal_names_the_model_and_the_route_that_works() {
let msg = encoder_refusal("bge-small-en-v1.5", "bert", 384, "CLS");
for fact in ["bge-small-en-v1.5", "bert", "384", "CLS", "/v1/embeddings"] {
assert!(msg.contains(fact), "{msg} does not carry {fact}");
}
assert!(msg.contains("embedding model"), "{msg}");
}
}