use std::sync::Arc;
use mentra::{
BuiltinProvider, ProviderId,
provider_core::{
AuthScheme, Provider, chat_completions,
gateway_ring::{GatewayRing, GatewayRingEvent, GatewayRingPolicy},
},
};
use crate::{error::RunError, provider, runtime::credential::Credential};
use super::{RuntimeBuilder, Wire, provider_settlement::responses_provider};
#[derive(Clone)]
pub struct GatewayMember {
base_url: String,
api_key: Option<String>,
}
impl GatewayMember {
pub fn new(base_url: impl Into<String>) -> Self {
Self {
base_url: base_url.into(),
api_key: None,
}
}
#[must_use]
pub fn with_api_key(self, api_key: impl Into<String>) -> Self {
Self {
api_key: Some(api_key.into()),
..self
}
}
pub fn base_url(&self) -> &str {
&self.base_url
}
}
impl std::fmt::Debug for GatewayMember {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GatewayMember")
.field("base_url", &self.base_url)
.field("api_key", &self.api_key.as_ref().map(|_| "<redacted>"))
.finish()
}
}
type Observer = Arc<dyn Fn(GatewayRingEvent) + Send + Sync>;
#[derive(Clone, Default)]
pub(in crate::runtime) struct GatewayRingSpec {
pub(super) members: Option<Vec<GatewayMember>>,
pub(super) policy: GatewayRingPolicy,
pub(super) observer: Option<Observer>,
}
impl GatewayRingSpec {
pub(in crate::runtime) fn is_stated(&self) -> bool {
self.members.is_some()
}
}
impl std::fmt::Debug for GatewayRingSpec {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GatewayRingSpec")
.field("members", &self.members)
.field("policy", &self.policy)
.field("observer", &self.observer.as_ref().map(|_| "<observer>"))
.finish()
}
}
impl RuntimeBuilder {
#[must_use]
pub fn with_gateway_ring(self, members: impl IntoIterator<Item = GatewayMember>) -> Self {
let spec = GatewayRingSpec {
members: Some(members.into_iter().collect()),
..self.gateway_ring.clone().unwrap_or_default()
};
Self {
gateway_ring: Some(spec),
..self
}
}
#[must_use]
pub fn with_gateway_ring_policy(self, policy: GatewayRingPolicy) -> Self {
let spec = GatewayRingSpec {
policy,
..self.gateway_ring.clone().unwrap_or_default()
};
Self {
gateway_ring: Some(spec),
..self
}
}
#[must_use]
pub fn with_gateway_ring_observer(
self,
observer: impl Fn(GatewayRingEvent) + Send + Sync + 'static,
) -> Self {
let spec = GatewayRingSpec {
observer: Some(Arc::new(observer)),
..self.gateway_ring.clone().unwrap_or_default()
};
Self {
gateway_ring: Some(spec),
..self
}
}
}
pub(super) fn assemble(
spec: GatewayRingSpec,
provider: BuiltinProvider,
wire: Wire,
) -> Result<GatewayRing, RunError> {
let members = spec
.members
.unwrap_or_default()
.iter()
.map(|member| build_member(member, provider, wire))
.collect::<Result<Vec<_>, _>>()?;
let ring = GatewayRing::new(members, spec.policy)
.map_err(|_| provider::ProviderError::EmptyGatewayRing)?;
Ok(match spec.observer {
Some(observer) => ring.with_observer(move |event| observer(event)),
None => ring,
})
}
fn build_member(
member: &GatewayMember,
provider: BuiltinProvider,
wire: Wire,
) -> Result<Arc<dyn Provider>, RunError> {
let base_url = provider::normalize_base_url(&member.base_url)?;
let credential = Credential::new(member.api_key.as_deref());
Ok(match wire {
Wire::Responses => Arc::new(responses_provider(provider, &base_url, credential)),
Wire::ChatCompletions => {
Arc::new(chat_completions_provider(provider, &base_url, credential))
}
})
}
fn chat_completions_provider(
provider: BuiltinProvider,
base_url: &str,
credential: Credential,
) -> chat_completions::ChatCompletionsProvider<Credential> {
let mut definition = chat_completions::definition(ProviderId::from(provider), base_url);
definition.descriptor.display_name = Some(format!("OpenAI-compatible ({base_url})"));
if !credential.is_some() {
definition.auth_scheme = AuthScheme::None;
}
chat_completions::ChatCompletionsProvider::new(definition, credential)
}