use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
pub use mentra_provider::AnthropicRequestOptions;
pub use mentra_provider::AuthScheme;
pub use mentra_provider::BuiltinProvider;
pub use mentra_provider::CompactionInputItem;
pub use mentra_provider::CompactionRequest;
pub use mentra_provider::CompactionResponse;
pub use mentra_provider::ContentBlock;
pub use mentra_provider::ContentBlockDelta;
pub use mentra_provider::ContentBlockStart;
pub use mentra_provider::EmbeddingData;
pub use mentra_provider::EmbeddingModelInfo;
pub use mentra_provider::EmbeddingProvider;
pub use mentra_provider::EmbeddingRequest;
pub use mentra_provider::EmbeddingResponse;
pub use mentra_provider::EmbeddingUsage;
pub use mentra_provider::GeminiRequestOptions;
pub use mentra_provider::ImageSource;
pub use mentra_provider::MemorySummarizeOutput;
pub use mentra_provider::MemorySummarizeRequest;
pub use mentra_provider::MemorySummarizeResponse;
pub use mentra_provider::Message;
pub use mentra_provider::ModelInfo;
pub use mentra_provider::ModelSelector;
pub use mentra_provider::OpenAIRequestOptions;
pub use mentra_provider::ProviderCapabilities;
pub use mentra_provider::ProviderCredentials;
pub use mentra_provider::ProviderDefinition;
pub use mentra_provider::ProviderDescriptor;
pub use mentra_provider::ProviderError;
pub use mentra_provider::ProviderEvent;
pub use mentra_provider::ProviderEventStream;
pub use mentra_provider::ProviderId;
pub use mentra_provider::ProviderRequestOptions;
pub use mentra_provider::RawMemory;
pub use mentra_provider::RawMemoryMetadata;
pub use mentra_provider::ReasoningEffort;
pub use mentra_provider::ReasoningFormat;
pub use mentra_provider::ReasoningOptions;
pub use mentra_provider::ReasoningProvenance;
pub use mentra_provider::Request;
pub use mentra_provider::Response;
pub use mentra_provider::ResponsesRequestOptions;
pub use mentra_provider::ResponsesStateMode;
pub use mentra_provider::ResponsesTransport;
pub use mentra_provider::Role;
pub use mentra_provider::TokenUsage;
pub use mentra_provider::ToolChoice;
pub use mentra_provider::ToolSearchMode;
pub use mentra_provider::WireApi;
pub use mentra_provider::collect_response_from_stream;
pub use mentra_provider::provider_event_stream_from_response;
#[async_trait]
pub trait Provider: Send + Sync {
fn descriptor(&self) -> ProviderDescriptor;
fn capabilities(&self) -> ProviderCapabilities {
ProviderCapabilities::default()
}
fn fresh_session_scope(&self) -> Result<ProviderSessionScope, ProviderError> {
Err(ProviderError::UnsupportedCapability(
"fresh_session_scope".to_string(),
))
}
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError>;
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError>;
async fn send(&self, request: Request<'_>) -> Result<Response, ProviderError> {
collect_response_from_stream(self.stream(request).await?).await
}
async fn compact(
&self,
_request: CompactionRequest<'_>,
) -> Result<CompactionResponse, ProviderError> {
Err(ProviderError::UnsupportedCapability(
"history_compaction".to_string(),
))
}
async fn summarize_memories(
&self,
_request: MemorySummarizeRequest<'_>,
) -> Result<MemorySummarizeResponse, ProviderError> {
Err(ProviderError::UnsupportedCapability(
"memory_summarization".to_string(),
))
}
}
#[derive(Clone)]
pub struct ProviderSessionScope {
inner: Arc<dyn Provider>,
}
impl ProviderSessionScope {
pub fn new<P>(provider: P) -> Self
where
P: Provider + 'static,
{
Self {
inner: Arc::new(provider),
}
}
}
#[async_trait]
impl Provider for ProviderSessionScope {
fn descriptor(&self) -> ProviderDescriptor {
self.inner.descriptor()
}
fn capabilities(&self) -> ProviderCapabilities {
self.inner.capabilities()
}
fn fresh_session_scope(&self) -> Result<ProviderSessionScope, ProviderError> {
self.inner.fresh_session_scope()
}
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
self.inner.list_models().await
}
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
self.inner.stream(request).await
}
async fn send(&self, request: Request<'_>) -> Result<Response, ProviderError> {
self.inner.send(request).await
}
async fn compact(
&self,
request: CompactionRequest<'_>,
) -> Result<CompactionResponse, ProviderError> {
self.inner.compact(request).await
}
async fn summarize_memories(
&self,
request: MemorySummarizeRequest<'_>,
) -> Result<MemorySummarizeResponse, ProviderError> {
self.inner.summarize_memories(request).await
}
}
#[async_trait]
impl Provider for Arc<dyn Provider> {
fn descriptor(&self) -> ProviderDescriptor {
(**self).descriptor()
}
fn capabilities(&self) -> ProviderCapabilities {
(**self).capabilities()
}
fn fresh_session_scope(&self) -> Result<ProviderSessionScope, ProviderError> {
(**self).fresh_session_scope()
}
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
(**self).list_models().await
}
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
(**self).stream(request).await
}
async fn send(&self, request: Request<'_>) -> Result<Response, ProviderError> {
(**self).send(request).await
}
async fn compact(
&self,
request: CompactionRequest<'_>,
) -> Result<CompactionResponse, ProviderError> {
(**self).compact(request).await
}
async fn summarize_memories(
&self,
request: MemorySummarizeRequest<'_>,
) -> Result<MemorySummarizeResponse, ProviderError> {
(**self).summarize_memories(request).await
}
}
#[derive(Default)]
pub struct ProviderRegistry {
default_provider: Option<ProviderId>,
default_embedding_provider: Option<ProviderId>,
providers: HashMap<ProviderId, Arc<dyn Provider>>,
embedding_providers: HashMap<ProviderId, Arc<dyn EmbeddingProvider>>,
responses_transport: Option<ResponsesTransport>,
}
impl ProviderRegistry {
pub(crate) fn register_builtin_provider(
&mut self,
id: BuiltinProvider,
api_key: impl Into<String>,
) -> Result<(), String> {
let api_key = api_key.into();
let provider: Arc<dyn Provider> = match id {
BuiltinProvider::Anthropic => anthropic::provider(api_key.clone()),
BuiltinProvider::Gemini => gemini::provider(api_key.clone()),
BuiltinProvider::OpenAI => openai::provider(api_key.clone()),
BuiltinProvider::OpenRouter => openrouter::provider(api_key.clone()),
BuiltinProvider::Ollama => ollama::provider(),
BuiltinProvider::LmStudio => lmstudio::provider(),
};
let provider_id: ProviderId = id.into();
if self.default_provider.is_none() {
self.default_provider = Some(provider_id.clone());
}
let ep: Option<Arc<dyn EmbeddingProvider>> = match id {
BuiltinProvider::OpenAI => Some(Arc::new(mentra_provider::responses::openai(api_key))),
BuiltinProvider::OpenRouter => {
Some(Arc::new(mentra_provider::responses::openrouter(api_key)))
}
BuiltinProvider::Ollama => Some(Arc::new(openai_compatible_embedding_provider(
id,
"http://127.0.0.1:11434/",
))),
BuiltinProvider::LmStudio => Some(Arc::new(openai_compatible_embedding_provider(
id,
"http://127.0.0.1:1234/",
))),
_ => None,
};
if let Some(ep) = ep {
if self.default_embedding_provider.is_none() {
self.default_embedding_provider = Some(provider_id.clone());
}
self.embedding_providers.insert(provider_id.clone(), ep);
}
self.providers.insert(provider_id, provider);
Ok(())
}
pub(crate) fn register_provider_instance<P>(&mut self, provider: P)
where
P: Provider + 'static,
{
let descriptor = provider.descriptor();
let id = descriptor.id;
if self.default_provider.is_none() {
self.default_provider = Some(id.clone());
}
self.providers.insert(id, Arc::new(provider));
}
pub(crate) fn register_registered_provider<P>(&mut self, provider: P)
where
P: mentra_provider::Provider + 'static,
{
let descriptor = provider.descriptor();
let id = descriptor.id;
if self.default_provider.is_none() {
self.default_provider = Some(id.clone());
}
self.providers.insert(id, shared_provider(provider));
}
pub(crate) fn register_shared_provider(&mut self, provider: Arc<dyn Provider>) {
let id = provider.descriptor().id;
if self.default_provider.is_none() {
self.default_provider = Some(id.clone());
}
self.providers.insert(id, provider);
}
pub(crate) fn register_ollama(&mut self) {
self.register_shared_provider(ollama::provider());
}
pub(crate) fn register_lmstudio(&mut self) {
self.register_shared_provider(lmstudio::provider());
}
pub(crate) fn get_provider(&self, id: Option<&ProviderId>) -> Option<Arc<dyn Provider>> {
match id {
Some(id) => self.providers.get(id).cloned(),
None => self
.default_provider
.as_ref()
.and_then(|id| self.providers.get(id).cloned()),
}
}
pub fn embedding_provider(&self) -> Option<Arc<dyn EmbeddingProvider>> {
self.default_embedding_provider
.as_ref()
.and_then(|id| self.embedding_providers.get(id))
.map(Arc::clone)
.or_else(|| self.embedding_providers.values().next().map(Arc::clone))
}
pub fn embedding_provider_for(&self, id: &ProviderId) -> Option<Arc<dyn EmbeddingProvider>> {
self.embedding_providers.get(id).map(Arc::clone)
}
pub(crate) fn descriptors(&self) -> Vec<ProviderDescriptor> {
self.providers
.values()
.map(|provider| provider.descriptor())
.collect()
}
pub(crate) fn is_empty(&self) -> bool {
self.providers.is_empty()
}
pub(crate) fn set_responses_transport(&mut self, transport: ResponsesTransport) {
self.responses_transport = Some(transport);
}
pub(crate) fn responses_transport(&self) -> Option<ResponsesTransport> {
self.responses_transport
}
}
pub(crate) fn select_responses_transport(
provider: &dyn Provider,
chosen: Option<ResponsesTransport>,
options: &mut ProviderRequestOptions,
) -> Result<(), crate::error::RuntimeError> {
if let Some(transport) = chosen {
options.responses.transport = transport;
}
if options.responses.transport != ResponsesTransport::WebSocket
|| provider.capabilities().supports_websockets
{
return Ok(());
}
let descriptor = provider.descriptor();
let name = descriptor
.display_name
.unwrap_or_else(|| descriptor.id.as_str().to_string());
Err(crate::error::RuntimeError::OperationDenied(format!(
"provider '{name}' does not serve the Responses websocket transport; \
select ResponsesTransport::HttpSse or register a provider that does \
— answering over HTTP+SSE would return a transport nobody asked for"
)))
}
fn shared_provider<P>(provider: P) -> Arc<dyn Provider>
where
P: mentra_provider::Provider + 'static,
{
Arc::new(SharedProviderProxy { inner: provider })
}
fn openai_compatible_embedding_provider(
builtin: BuiltinProvider,
base_url: &str,
) -> mentra_provider::responses::ResponsesProvider<NoCredentialsSource> {
use mentra_provider::AuthScheme;
use mentra_provider::ProviderCapabilities;
use mentra_provider::WireApi;
use std::collections::HashMap;
let mut definition = ProviderDefinition::new(builtin);
definition.wire_api = WireApi::Responses;
definition.auth_scheme = AuthScheme::None;
definition.capabilities = ProviderCapabilities {
supports_model_listing: true,
supports_streaming: true,
supports_websockets: false,
supports_tool_calls: true,
supports_images: true,
supports_history_compaction: false,
supports_memory_summarization: false,
supports_deferred_tools: false,
supports_hosted_tool_search: false,
supports_hosted_web_search: false,
supports_image_generation: false,
supports_reasoning_effort: false,
reports_reasoning_tokens: false,
reports_thoughts_tokens: false,
supports_structured_tool_results: false,
supports_embeddings: true,
};
definition.base_url = Some(base_url.to_string());
definition.headers = Some(HashMap::new());
mentra_provider::responses::ResponsesProvider::new(definition, NoCredentialsSource)
}
fn chat_completions_provider(
provider: impl Into<mentra_provider::ProviderId>,
display_name: &str,
description: &str,
base_url: &str,
credentials: Option<String>,
) -> Arc<dyn Provider> {
let mut definition = mentra_provider::chat_completions::definition(provider, base_url);
definition.descriptor.display_name = Some(display_name.to_string());
definition.descriptor.description = Some(description.to_string());
match credentials {
Some(api_key) => shared_provider(
mentra_provider::chat_completions::ChatCompletionsProvider::new(
definition,
mentra_provider::StaticCredentialSource::new(api_key),
),
),
None => {
definition.auth_scheme = mentra_provider::AuthScheme::None;
shared_provider(
mentra_provider::chat_completions::ChatCompletionsProvider::new(
definition,
NoCredentialsSource,
),
)
}
}
}
#[derive(Clone)]
struct NoCredentialsSource;
#[async_trait]
impl mentra_provider::CredentialSource for NoCredentialsSource {
async fn credentials(
&self,
) -> Result<mentra_provider::ProviderCredentials, mentra_provider::ProviderError> {
Ok(mentra_provider::ProviderCredentials::default())
}
}
struct SharedProviderProxy<P> {
inner: P,
}
#[async_trait]
impl<P> Provider for SharedProviderProxy<P>
where
P: mentra_provider::Provider + 'static,
{
fn descriptor(&self) -> ProviderDescriptor {
self.inner.descriptor()
}
fn capabilities(&self) -> ProviderCapabilities {
self.inner.definition().capabilities
}
fn fresh_session_scope(&self) -> Result<ProviderSessionScope, ProviderError> {
let scope = mentra_provider::Provider::fresh_session_scope(&self.inner)?;
Ok(ProviderSessionScope::new(SharedProviderProxy {
inner: scope,
}))
}
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
self.inner.list_models().await
}
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
self.inner.stream(request).await
}
async fn compact(
&self,
request: CompactionRequest<'_>,
) -> Result<CompactionResponse, ProviderError> {
self.inner.compact(request).await
}
async fn summarize_memories(
&self,
request: MemorySummarizeRequest<'_>,
) -> Result<MemorySummarizeResponse, ProviderError> {
self.inner.summarize_memories(request).await
}
}
pub mod openai {
use std::sync::Arc;
use async_trait::async_trait;
use super::Provider;
use super::shared_provider;
#[async_trait]
pub trait OpenAICredentialSource: Send + Sync {
async fn api_key(&self) -> Result<String, String>;
}
pub fn provider(api_key: impl Into<String>) -> Arc<dyn Provider> {
shared_provider(mentra_provider::responses::openai(api_key))
}
pub fn with_credential_source(
source: impl OpenAICredentialSource + 'static,
) -> Arc<dyn Provider> {
with_shared_credential_source(Arc::new(source))
}
pub fn with_shared_credential_source(
source: Arc<dyn OpenAICredentialSource>,
) -> Arc<dyn Provider> {
shared_provider(mentra_provider::responses::openai_with_credential_source(
OpenAICredentialAdapter { source },
))
}
#[derive(Clone)]
struct OpenAICredentialAdapter {
source: Arc<dyn OpenAICredentialSource>,
}
#[async_trait]
impl mentra_provider::CredentialSource for OpenAICredentialAdapter {
async fn credentials(
&self,
) -> Result<mentra_provider::ProviderCredentials, mentra_provider::ProviderError> {
let api_key = self
.source
.api_key()
.await
.map_err(mentra_provider::ProviderError::InvalidRequest)?;
Ok(mentra_provider::ProviderCredentials {
bearer_token: Some(api_key),
account_id: None,
headers: Default::default(),
})
}
}
}
pub mod openrouter {
use std::sync::Arc;
use super::Provider;
use super::shared_provider;
pub fn provider(api_key: impl Into<String>) -> Arc<dyn Provider> {
shared_provider(mentra_provider::responses::openrouter(api_key))
}
}
pub mod anthropic {
use std::sync::Arc;
use super::Provider;
use super::shared_provider;
pub fn provider(api_key: impl Into<String>) -> Arc<dyn Provider> {
shared_provider(mentra_provider::anthropic::AnthropicProvider::new(api_key))
}
}
pub mod gemini {
use std::sync::Arc;
use super::Provider;
use super::shared_provider;
pub fn provider(api_key: impl Into<String>) -> Arc<dyn Provider> {
shared_provider(mentra_provider::gemini::GeminiProvider::new(api_key))
}
}
pub mod openai_compatible {
use std::sync::Arc;
use super::Provider;
const DESCRIPTION: &str = "OpenAI-compatible chat/completions provider";
pub fn new(
id: impl Into<mentra_provider::ProviderId>,
base_url: impl AsRef<str>,
api_key: impl Into<String>,
) -> Arc<dyn Provider> {
let id = id.into();
let display_name = id.as_str().to_string();
super::chat_completions_provider(
id,
&display_name,
DESCRIPTION,
base_url.as_ref(),
Some(api_key.into()),
)
}
pub fn without_credentials(
id: impl Into<mentra_provider::ProviderId>,
base_url: impl AsRef<str>,
) -> Arc<dyn Provider> {
let id = id.into();
let display_name = id.as_str().to_string();
super::chat_completions_provider(id, &display_name, DESCRIPTION, base_url.as_ref(), None)
}
}
pub mod ollama {
use std::sync::Arc;
use super::BuiltinProvider;
use super::Provider;
const DEFAULT_BASE_URL: &str = "http://127.0.0.1:11434/";
pub fn provider() -> Arc<dyn Provider> {
with_base_url(DEFAULT_BASE_URL)
}
pub fn with_base_url(base_url: impl AsRef<str>) -> Arc<dyn Provider> {
super::chat_completions_provider(
BuiltinProvider::Ollama,
"Ollama",
"Ollama OpenAI-compatible chat/completions provider",
base_url.as_ref(),
None,
)
}
}
pub mod lmstudio {
use std::sync::Arc;
use super::BuiltinProvider;
use super::Provider;
const DEFAULT_BASE_URL: &str = "http://127.0.0.1:1234/";
pub fn provider() -> Arc<dyn Provider> {
with_base_url(DEFAULT_BASE_URL)
}
pub fn with_base_url(base_url: impl AsRef<str>) -> Arc<dyn Provider> {
super::chat_completions_provider(
BuiltinProvider::LmStudio,
"LM Studio",
"LM Studio OpenAI-compatible chat/completions provider",
base_url.as_ref(),
None,
)
}
}