use std::collections::HashMap;
use std::sync::Arc;
use crate::config::{Config, Protocol, ThinkingMode};
use crate::error::{ConfigError, GatewayError};
use crate::queue::EndpointLane;
use crate::upstream::{OpenAiUpstream, Upstream};
pub(crate) struct Endpoint {
pub id: String,
pub upstream: Arc<dyn Upstream>,
pub lane: EndpointLane,
}
impl std::fmt::Debug for Endpoint {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Endpoint")
.field("id", &self.id)
.field("lane", &self.lane)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub(crate) struct Model {
pub name: String,
pub description: String,
pub context: u32,
pub thinking: ThinkingMode,
pub tool_dialect: String,
pub tools_mode: String,
pub upstream_name: String,
pub endpoint: Arc<Endpoint>,
}
#[derive(Debug)]
pub(crate) struct Routing {
by_name: HashMap<String, Arc<Model>>,
models: Vec<Arc<Model>>,
}
impl Routing {
pub(crate) fn new(models: Vec<Arc<Model>>) -> Result<Routing, ConfigError> {
let mut by_name = HashMap::with_capacity(models.len());
for model in &models {
if by_name
.insert(model.name.clone(), Arc::clone(model))
.is_some()
{
return Err(ConfigError::Validation(format!(
"duplicate model name {}",
model.name
)));
}
}
Ok(Routing { by_name, models })
}
#[must_use]
pub(crate) fn models(&self) -> &[Arc<Model>] {
&self.models
}
pub(crate) fn from_config(config: &Config) -> Result<Routing, ConfigError> {
let mut endpoints: HashMap<&str, Arc<Endpoint>> = HashMap::new();
for endpoint in &config.endpoints {
let upstream: Arc<dyn Upstream> = match endpoint.protocol {
Protocol::Openai => Arc::new(OpenAiUpstream::new(
&endpoint.base_url,
endpoint.api_key.clone(),
)),
};
let lane = match config.endpoint_concurrency(endpoint) {
Some(n) => EndpointLane::new(n, &config.queue),
None => EndpointLane::unlimited(),
};
endpoints.insert(
endpoint.id.as_str(),
Arc::new(Endpoint {
id: endpoint.id.clone(),
upstream,
lane,
}),
);
}
let mut models = Vec::with_capacity(config.models.len());
for model in &config.models {
let endpoint_id = model.endpoints.first().ok_or_else(|| {
ConfigError::Validation(format!("model {} has no endpoints", model.name))
})?;
let endpoint = endpoints.get(endpoint_id.as_str()).ok_or_else(|| {
ConfigError::Validation(format!(
"model {} names undefined endpoint {endpoint_id}",
model.name
))
})?;
models.push(Arc::new(Model {
name: model.name.clone(),
description: model.description.clone(),
context: model.context,
thinking: model.thinking,
tool_dialect: "openai".to_owned(),
tools_mode: "native".to_owned(),
upstream_name: model.upstream.clone(),
endpoint: Arc::clone(endpoint),
}));
}
Routing::new(models)
}
pub(crate) fn merge(
mut self,
extras: impl IntoIterator<Item = Arc<Model>>,
) -> Result<Routing, ConfigError> {
for model in extras {
if self.by_name.contains_key(&model.name) {
return Err(ConfigError::Validation(format!(
"duplicate model name {}",
model.name
)));
}
self.by_name.insert(model.name.clone(), Arc::clone(&model));
self.models.push(model);
}
Ok(self)
}
pub(crate) fn model(&self, name: &str) -> Result<Arc<Model>, GatewayError> {
self.by_name
.get(name)
.cloned()
.ok_or_else(|| GatewayError::UnknownModel(name.to_string()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Config;
fn model_named(name: &str) -> Arc<Model> {
let endpoint = Arc::new(Endpoint {
id: "e".to_owned(),
upstream: Arc::new(OpenAiUpstream::new(
"http://127.0.0.1:9",
crate::config::Secret::new(String::new()),
)),
lane: EndpointLane::unlimited(),
});
Arc::new(Model {
name: name.to_owned(),
description: "d".to_owned(),
context: 8192,
thinking: ThinkingMode::Never,
tool_dialect: "openai".to_owned(),
tools_mode: "native".to_owned(),
upstream_name: "u".to_owned(),
endpoint,
})
}
fn routing() -> Routing {
let toml = r#"
[server]
bind = "127.0.0.1:8081"
key = "t"
[[endpoint]]
id = "e"
protocol = "openai"
base_url = "http://127.0.0.1:9"
api_key = ""
[[model]]
name = "known"
description = "a known test model"
context = 8192
upstream = "backend-name"
endpoints = ["e"]
"#;
let config = Config::from_toml_str(toml).unwrap();
Routing::from_config(&config).unwrap()
}
#[test]
fn resolves_a_known_model() {
let r = routing();
let m = r.model("known").unwrap();
assert_eq!(m.upstream_name, "backend-name");
assert_eq!(m.endpoint.id, "e");
}
#[test]
fn unknown_model_errors() {
let r = routing();
assert!(matches!(
r.model("nope"),
Err(GatewayError::UnknownModel(_))
));
}
#[test]
fn new_rejects_duplicate_model_names() {
let dup = Routing::new(vec![model_named("m"), model_named("m")]);
assert!(matches!(dup, Err(ConfigError::Validation(_))));
}
#[test]
fn new_preserves_catalog_order() {
let r = Routing::new(vec![model_named("a"), model_named("b"), model_named("c")])
.expect("distinct names");
let names: Vec<&str> = r.models().iter().map(|m| m.name.as_str()).collect();
assert_eq!(names, ["a", "b", "c"]);
}
#[test]
fn merge_rejects_duplicate_model_names() {
let base = Routing::new(vec![model_named("known")]).expect("distinct");
let merged = base.merge([model_named("known")]);
assert!(matches!(merged, Err(ConfigError::Validation(_))));
}
#[test]
fn merge_appends_after_existing_models() {
let base = Routing::new(vec![model_named("a")]).expect("distinct");
let merged = base.merge([model_named("b")]).expect("distinct extra");
let names: Vec<&str> = merged.models().iter().map(|m| m.name.as_str()).collect();
assert_eq!(names, ["a", "b"]);
}
}