use mentra::{
BuiltinProvider, ProviderId,
provider_core::{AuthScheme, responses, responses::ResponsesProvider},
};
use crate::{error::RunError, provider, runtime::credential::Credential};
use super::Wire;
pub(in crate::runtime) struct HostProvider {
pub(super) id: ProviderId,
pub(super) install:
Box<dyn FnOnce(mentra::RuntimeBuilder) -> mentra::RuntimeBuilder + Send + Sync>,
}
impl HostProvider {
pub(in crate::runtime) fn id(&self) -> &ProviderId {
&self.id
}
pub(in crate::runtime) fn registered<P>(provider: P) -> Self
where
P: mentra::provider_core::Provider + 'static,
{
let id = provider.descriptor().id;
Self {
id,
install: Box::new(move |builder| builder.with_registered_provider(provider)),
}
}
}
pub(super) enum ProviderSource {
Host(HostProvider),
Resolved(provider::ProviderChoice),
}
pub(super) fn settle(
host_provider: Option<HostProvider>,
provider: Option<BuiltinProvider>,
base_url: Option<String>,
api_key: Option<String>,
) -> Result<ProviderSource, RunError> {
match host_provider {
Some(host) => {
validate_host_provider_source(
provider.as_ref(),
base_url.as_deref(),
api_key.as_deref(),
)?;
Ok(ProviderSource::Host(host))
}
None => Ok(ProviderSource::Resolved(provider::resolve_with(
provider,
base_url.as_deref(),
api_key.as_deref(),
)?)),
}
}
pub(super) fn validate_host_provider_source(
provider: Option<&BuiltinProvider>,
base_url: Option<&str>,
api_key: Option<&str>,
) -> Result<(), RunError> {
for (also_set, knob) in [
(provider.is_some(), "with_provider"),
(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(())
}
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::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()
}