pub mod compatible;
pub mod compatible_streaming;
pub mod reasoning_roundtrip;
pub mod reliable;
pub mod transcribe;
use crate::config::{CONFIG, trimmed_or_none};
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();
let image_transcriber = create_transcriber(
Some(&endpoint),
api_key.as_deref(),
Some(CONFIG.image_transcription_model().as_str()),
CONFIG.transcription_provider().as_deref(),
ImageTranscriber::from_inner,
);
let audio_transcriber = create_transcriber(
Some(&endpoint),
api_key.as_deref(),
Some(CONFIG.audio_transcription_model().as_str()),
CONFIG.audio_transcription_provider().as_deref(),
transcribe::AudioTranscriber::from_inner,
);
match warmup_mode {
WarmupMode::Fatal => {
provider.warmup().await?;
}
WarmupMode::NonFatal => {
if let Err(e) = provider.warmup().await {
tracing::warn!(endpoint = %endpoint, "Provider warmup failed (non-fatal): {e}");
}
}
}
*PROVIDER.write().expect("PROVIDER poisoned") = Some(provider);
*IMAGE_TRANSCRIBER
.write()
.expect("IMAGE_TRANSCRIBER poisoned") = image_transcriber;
*AUDIO_TRANSCRIBER
.write()
.expect("AUDIO_TRANSCRIBER poisoned") = audio_transcriber;
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(trimmed_or_none);
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(trimmed_or_none);
let resolved_key = key_owned.as_deref();
let base_url = endpoint
.and_then(trimmed_or_none)
.unwrap_or(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(trimmed_or_none)?;
let model = model.and_then(trimmed_or_none)?;
let route = provider.and_then(trimmed_or_none);
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))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn url_roundtrips() {
struct Case {
name: &'static str,
input: &'static str,
expected_chat: &'static str,
expected_base: &'static str,
}
let cases = [
Case {
name: "already_has_suffix",
input: "https://api.example.com/v1/chat/completions",
expected_chat: "https://api.example.com/v1/chat/completions",
expected_base: "https://api.example.com/v1",
},
Case {
name: "no_suffix",
input: "https://api.example.com/v1",
expected_chat: "https://api.example.com/v1/chat/completions",
expected_base: "https://api.example.com/v1",
},
Case {
name: "trailing_slash",
input: "https://api.example.com/v1/",
expected_chat: "https://api.example.com/v1/chat/completions",
expected_base: "https://api.example.com/v1",
},
Case {
name: "double_trailing_slash",
input: "https://api.example.com/v1//",
expected_chat: "https://api.example.com/v1/chat/completions",
expected_base: "https://api.example.com/v1",
},
Case {
name: "trailing_slash_before_suffix",
input: "https://api.example.com/v1/chat/completions/",
expected_chat: "https://api.example.com/v1/chat/completions",
expected_base: "https://api.example.com/v1",
},
Case {
name: "domain_containing_chat_completions",
input: "https://chat.completions.com/api",
expected_chat: "https://chat.completions.com/api/chat/completions",
expected_base: "https://chat.completions.com/api",
},
];
for c in &cases {
assert_eq!(
ensure_chat_completions_url(c.input),
c.expected_chat,
"case '{}': ensure_chat_completions_url({:?})",
c.name,
c.input,
);
assert_eq!(
ensure_base_url(c.input),
c.expected_base,
"case '{}': ensure_base_url({:?})",
c.name,
c.input,
);
}
let roundtrip_inputs = &[
"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 roundtrip_inputs {
let base = ensure_base_url(url);
let chat = ensure_chat_completions_url(&base);
let roundtripped = ensure_base_url(&chat);
assert_eq!(
roundtripped, base,
"roundtrip(base->chat->base) should be identity for '{url}'",
);
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(chat->base->chat) should be identity for '{url}'",
);
}
}
#[test]
fn provider_routing() {
struct Case {
name: &'static str,
order: &'static str,
allow_fallbacks: bool,
expected: Option<serde_json::Value>,
}
let cases = [
Case {
name: "single_provider",
order: "openai",
allow_fallbacks: false,
expected: Some(serde_json::json!({
"order": ["openai"],
"allow_fallbacks": false,
})),
},
Case {
name: "multiple_providers",
order: "openai, anthropic, google",
allow_fallbacks: true,
expected: Some(serde_json::json!({
"order": ["openai", "anthropic", "google"],
"allow_fallbacks": true,
})),
},
Case {
name: "whitespace_only_yields_none",
order: " , , ",
allow_fallbacks: false,
expected: None,
},
Case {
name: "empty_string_yields_none",
order: "",
allow_fallbacks: true,
expected: None,
},
Case {
name: "leading_trailing_whitespace",
order: " openai ",
allow_fallbacks: false,
expected: Some(serde_json::json!({
"order": ["openai"],
"allow_fallbacks": false,
})),
},
Case {
name: "single_slug_survives_split",
order: "google-gemini",
allow_fallbacks: false,
expected: Some(serde_json::json!({
"order": ["google-gemini"],
"allow_fallbacks": false,
})),
},
];
for c in &cases {
assert_eq!(
provider_routing_json(c.order, c.allow_fallbacks),
c.expected,
"case '{}': provider_routing_json({:?}, {})",
c.name,
c.order,
c.allow_fallbacks,
);
}
}
}