pub mod compatible;
pub mod compatible_streaming;
pub mod error;
pub mod reasoning_roundtrip;
pub mod reliable;
pub mod transcribe;
use crate::config::{CONFIG, trim_non_empty};
pub use crate::{ChatMessage, ChatRequest, ChatResponse, Provider};
use crate::{StreamEvent, StreamResult};
use futures_util::stream;
use std::sync::{Arc, RwLock};
pub use crate::providers::transcribe::ImageTranscriber;
use compatible::OpenAiCompatibleProvider;
use reliable::ReliableProvider;
pub(crate) fn ensure_chat_completions_url(base_url: &str) -> String {
let trimmed = base_url.trim_end_matches('/');
if trimmed.ends_with("/chat/completions") {
trimmed.to_string()
} else {
format!("{trimmed}/chat/completions")
}
}
pub(crate) fn ensure_base_url(endpoint: &str) -> String {
endpoint
.trim_end_matches('/')
.trim_end_matches("/chat/completions")
.to_string()
}
pub(crate) fn provider_routing_json(
order: &str,
allow_fallbacks: bool,
) -> Option<serde_json::Value> {
let providers: Vec<&str> = order
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.collect();
if providers.is_empty() {
return None;
}
Some(serde_json::json!({
"order": providers,
"allow_fallbacks": allow_fallbacks,
}))
}
static PROVIDER: RwLock<Option<Arc<dyn Provider>>> = RwLock::new(None);
static IMAGE_TRANSCRIBER: RwLock<Option<ImageTranscriber>> = RwLock::new(None);
static AUDIO_TRANSCRIBER: RwLock<Option<transcribe::AudioTranscriber>> = RwLock::new(None);
enum WarmupMode {
NonFatal,
Fatal,
}
async fn setup_provider_and_transcribers(warmup_mode: WarmupMode) -> anyhow::Result<()> {
let api_key = CONFIG.provider_key();
let endpoint = CONFIG.provider_endpoint();
let endpoint_opt = if endpoint == crate::config::DEFAULT_PROVIDER_ENDPOINT {
None
} else {
Some(endpoint.as_str())
};
let provider: Arc<dyn Provider> = create_provider(api_key.as_deref(), endpoint_opt)?.into();
match warmup_mode {
WarmupMode::Fatal => {
provider.warmup().await?;
}
WarmupMode::NonFatal => {
if let Err(e) = provider.warmup().await {
tracing::warn!("Provider warmup failed (non-fatal): {e}");
}
}
}
*PROVIDER.write().expect("PROVIDER poisoned") = Some(provider);
create_and_store_transcriber(
&IMAGE_TRANSCRIBER,
"IMAGE_TRANSCRIBER",
Some(&endpoint),
api_key.as_deref(),
Some(CONFIG.image_transcription_model().as_str()),
CONFIG.transcription_provider().as_deref(),
ImageTranscriber::from_inner,
);
create_and_store_transcriber(
&AUDIO_TRANSCRIBER,
"AUDIO_TRANSCRIBER",
Some(&endpoint),
api_key.as_deref(),
Some(CONFIG.audio_transcription_model().as_str()),
CONFIG.audio_transcription_provider().as_deref(),
transcribe::AudioTranscriber::from_inner,
);
Ok(())
}
pub async fn init_global() -> anyhow::Result<()> {
setup_provider_and_transcribers(WarmupMode::NonFatal).await
}
pub async fn warmup_provider_from_config(config: &crate::config::ConfigData) -> anyhow::Result<()> {
let endpoint = config.provider_endpoint.as_deref().and_then(trim_non_empty);
let endpoint_opt = endpoint.filter(|e| e.as_str() != crate::config::DEFAULT_PROVIDER_ENDPOINT);
let provider = create_provider(config.provider_key.as_deref(), endpoint_opt.as_deref())?;
provider.warmup().await?;
Ok(())
}
pub async fn recreate_all() -> anyhow::Result<()> {
setup_provider_and_transcribers(WarmupMode::Fatal).await?;
tracing::info!("Provider and transcriber singletons recreated");
Ok(())
}
#[must_use]
pub fn image_transcriber() -> Option<ImageTranscriber> {
IMAGE_TRANSCRIBER
.read()
.expect("IMAGE_TRANSCRIBER poisoned")
.clone()
}
#[must_use]
pub fn audio_transcriber() -> Option<transcribe::AudioTranscriber> {
AUDIO_TRANSCRIBER
.read()
.expect("AUDIO_TRANSCRIBER poisoned")
.clone()
}
pub async fn chat(request: ChatRequest) -> anyhow::Result<ChatResponse> {
let provider = PROVIDER
.read()
.expect("PROVIDER poisoned")
.clone()
.expect("PROVIDER not initialized");
provider.chat(request).await
}
pub fn stream_chat(request: ChatRequest) -> stream::BoxStream<'static, StreamResult<StreamEvent>> {
let provider = PROVIDER
.read()
.expect("PROVIDER poisoned")
.clone()
.expect("PROVIDER not initialized");
provider.stream_chat(request)
}
pub fn create_provider(
api_key: Option<&str>,
endpoint: Option<&str>,
) -> anyhow::Result<Box<dyn Provider>> {
let key_owned = api_key.and_then(trim_non_empty);
let resolved_key = key_owned.as_deref();
let base_url = endpoint
.and_then(trim_non_empty)
.unwrap_or_else(|| crate::config::DEFAULT_PROVIDER_ENDPOINT.to_string());
let mut extra_headers = std::collections::HashMap::new();
extra_headers.insert("X-Title".to_string(), "MahBot".to_string());
extra_headers.insert(
"HTTP-Referrer".to_string(),
"https://github.com/edezhic".to_string(),
);
let base = OpenAiCompatibleProvider::new("OpenRouter", base_url.as_str(), resolved_key)
.with_extra_headers(extra_headers);
let provider: Box<dyn Provider> = Box::new(base);
let reliable: Box<dyn Provider> = Box::new(ReliableProvider::new(
"openrouter".to_string(),
provider,
10,
500,
));
Ok(reliable)
}
#[must_use]
fn create_transcriber<T>(
api_url: Option<&str>,
api_key: Option<&str>,
model: Option<&str>,
provider: Option<&str>,
wrapper: impl FnOnce(transcribe::MediaTranscriber) -> T,
) -> Option<T> {
let _key = api_key.and_then(trim_non_empty)?;
let model = model.and_then(trim_non_empty)?;
let route = provider.and_then(trim_non_empty);
let base_url = api_url
.unwrap_or(crate::config::DEFAULT_PROVIDER_ENDPOINT)
.to_string();
let inner = transcribe::MediaTranscriber::new(base_url, model, route);
Some(wrapper(inner))
}
fn create_and_store_transcriber<T: Clone + Send + Sync + 'static>(
store: &RwLock<Option<T>>,
name: &str,
api_url: Option<&str>,
api_key: Option<&str>,
model: Option<&str>,
provider: Option<&str>,
wrapper: impl FnOnce(transcribe::MediaTranscriber) -> T,
) {
let transcriber = create_transcriber(api_url, api_key, model, provider, wrapper);
*store.write().unwrap_or_else(|_| panic!("{name} poisoned")) = transcriber;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ensure_chat_completions_url_already_has_suffix() {
assert_eq!(
ensure_chat_completions_url("https://api.example.com/v1/chat/completions"),
"https://api.example.com/v1/chat/completions",
);
}
#[test]
fn ensure_chat_completions_url_no_suffix() {
assert_eq!(
ensure_chat_completions_url("https://api.example.com/v1"),
"https://api.example.com/v1/chat/completions",
);
}
#[test]
fn ensure_chat_completions_url_trailing_slash() {
assert_eq!(
ensure_chat_completions_url("https://api.example.com/v1/"),
"https://api.example.com/v1/chat/completions",
);
}
#[test]
fn ensure_chat_completions_url_double_trailing_slash() {
assert_eq!(
ensure_chat_completions_url("https://api.example.com/v1//"),
"https://api.example.com/v1/chat/completions",
"multiple trailing slashes are collapsed by trim_end_matches('/')",
);
}
#[test]
fn ensure_base_url_already_clean() {
assert_eq!(
ensure_base_url("https://api.example.com/v1"),
"https://api.example.com/v1",
);
}
#[test]
fn ensure_base_url_strips_chat_completions() {
assert_eq!(
ensure_base_url("https://api.example.com/v1/chat/completions"),
"https://api.example.com/v1",
);
}
#[test]
fn ensure_base_url_trailing_slash() {
assert_eq!(
ensure_base_url("https://api.example.com/v1/"),
"https://api.example.com/v1",
);
}
#[test]
fn ensure_base_url_trailing_slash_before_suffix() {
assert_eq!(
ensure_base_url("https://api.example.com/v1/chat/completions/"),
"https://api.example.com/v1",
);
}
#[test]
fn roundtrip_base_to_chat_to_base() {
let cases = &[
"https://api.example.com/v1",
"https://api.example.com/v1/",
"https://api.example.com/v1/chat/completions",
"https://api.example.com/v1/chat/completions/",
];
for &url in cases {
let base = ensure_base_url(url);
let chat = ensure_chat_completions_url(&base);
let roundtripped = ensure_base_url(&chat);
assert_eq!(
roundtripped, base,
"roundtrip(ensure_base_url -> ensure_chat_completions_url -> ensure_base_url) \
should be identity for input '{url}'",
);
}
}
#[test]
fn roundtrip_chat_to_base_to_chat() {
let cases = &[
"https://api.example.com/v1",
"https://api.example.com/v1/",
"https://api.example.com/v1/chat/completions",
"https://api.example.com/v1/chat/completions/",
];
for &url in cases {
let chat = ensure_chat_completions_url(url);
let base = ensure_base_url(&chat);
let roundtripped = ensure_chat_completions_url(&base);
assert_eq!(
roundtripped, chat,
"roundtrip(ensure_chat_completions_url -> ensure_base_url -> \
ensure_chat_completions_url) should be identity for input '{url}'",
);
}
}
#[test]
fn domain_name_containing_slash_chat_completions() {
let url = "https://chat.completions.com/api";
let chat = ensure_chat_completions_url(url);
assert_eq!(chat, "https://chat.completions.com/api/chat/completions");
let base = ensure_base_url(url);
assert_eq!(base, "https://chat.completions.com/api");
}
#[test]
fn routing_single_provider() {
assert_eq!(
provider_routing_json("openai", false),
Some(serde_json::json!({
"order": ["openai"],
"allow_fallbacks": false,
})),
);
}
#[test]
fn routing_multiple_providers() {
assert_eq!(
provider_routing_json("openai, anthropic, google", true),
Some(serde_json::json!({
"order": ["openai", "anthropic", "google"],
"allow_fallbacks": true,
})),
);
}
#[test]
fn routing_whitespace_only_yields_none() {
assert_eq!(provider_routing_json(" , , ", false), None);
}
#[test]
fn routing_empty_string_yields_none() {
assert_eq!(provider_routing_json("", true), None);
}
#[test]
fn routing_leading_trailing_whitespace() {
assert_eq!(
provider_routing_json(" openai ", false),
Some(serde_json::json!({
"order": ["openai"],
"allow_fallbacks": false,
})),
);
}
#[test]
fn routing_single_slug_survives_split() {
assert_eq!(
provider_routing_json("google-gemini", false),
Some(serde_json::json!({
"order": ["google-gemini"],
"allow_fallbacks": false,
})),
);
}
}