use mentra::{
BuiltinProvider, ProviderId,
provider_core::{AuthScheme, responses, responses::ResponsesProvider},
};
use crate::{error::RunError, provider, runtime::credential::Credential};
use super::{GatewayRingSpec, Wire, gateway_ring};
pub(in crate::runtime) struct HostProvider {
pub(super) id: ProviderId,
pub(super) install:
Box<dyn FnOnce(mentra::RuntimeBuilder) -> mentra::RuntimeBuilder + Send + Sync>,
}
pub(super) enum ProviderSource {
Host(HostProvider),
Ring {
spec: GatewayRingSpec,
provider: BuiltinProvider,
},
Resolved(provider::ProviderChoice),
}
pub(super) fn settle(
host_provider: Option<HostProvider>,
gateway_ring: Option<GatewayRingSpec>,
provider: Option<BuiltinProvider>,
base_url: Option<String>,
api_key: Option<String>,
) -> Result<ProviderSource, RunError> {
let gateway_ring = gateway_ring.filter(GatewayRingSpec::is_stated);
match (host_provider, gateway_ring) {
(Some(host), gateway_ring) => {
for (also_set, knob) in [
(provider.is_some(), "with_provider"),
(base_url.is_some(), "with_base_url"),
(api_key.is_some(), "with_api_key"),
(gateway_ring.is_some(), "with_gateway_ring"),
] {
if also_set {
return Err(provider::ProviderError::AmbiguousProviderSource { knob }.into());
}
}
Ok(ProviderSource::Host(host))
}
(None, Some(spec)) => {
for (also_set, knob) in [
(base_url.is_some(), "with_base_url"),
(api_key.is_some(), "with_api_key"),
] {
if also_set {
return Err(provider::ProviderError::AmbiguousProviderSource { knob }.into());
}
}
Ok(ProviderSource::Ring {
spec,
provider: provider.unwrap_or(provider::DEFAULT_COMPATIBLE_PROVIDER),
})
}
(None, None) => Ok(ProviderSource::Resolved(provider::resolve_with(
provider,
base_url.as_deref(),
api_key.as_deref(),
)?)),
}
}
pub(super) fn assemble(
source: ProviderSource,
builder: mentra::RuntimeBuilder,
wire: Wire,
) -> Result<(mentra::Runtime, ProviderId), RunError> {
match source {
ProviderSource::Host(host) => Ok(((host.install)(builder).build()?, host.id)),
ProviderSource::Ring { spec, provider } => {
let ring = gateway_ring::assemble(spec, provider, wire)?;
let id = ProviderId::from(provider);
Ok((builder.with_registered_provider(ring).build()?, id))
}
ProviderSource::Resolved(choice) => {
let assembled = match (&choice.base_url, wire) {
(Some(base_url), Wire::ChatCompletions) => builder.with_openai_compatible(
ProviderId::from(choice.provider),
base_url,
choice.api_key.clone(),
),
(Some(base_url), Wire::Responses) => {
builder.with_registered_provider(responses_provider(
choice.provider,
base_url,
Credential::new(choice.api_key.as_deref()),
))
}
(None, _) => builder
.with_provider(choice.provider, choice.api_key.clone().unwrap_or_default()),
};
Ok((assembled.build()?, ProviderId::from(choice.provider)))
}
}
}
pub(super) fn responses_provider(
provider: BuiltinProvider,
base_url: &str,
credential: Credential,
) -> ResponsesProvider<Credential> {
let mut definition = responses::openai_definition();
definition.base_url = Some(base_url.to_string());
definition.descriptor.id = ProviderId::from(provider);
definition.descriptor.display_name = Some(format!("OpenAI-compatible ({base_url})"));
if !credential.is_some() {
definition.auth_scheme = AuthScheme::None;
}
ResponsesProvider::new(definition, credential).without_hybrid_http_previous_response_id()
}