pub(crate) mod compatible;
pub(crate) mod reasoning;
pub(crate) mod reasoning_roundtrip;
pub(crate) mod reliable;
pub(crate) mod transcribe;
pub(crate) use reasoning::plaintext_for_display;
use crate::config::{CONFIG, resolve_or, trimmed_or_none};
use crate::util::UnwrapPoison;
pub(crate) use crate::{ChatRequest, ChatResponse, Provider};
#[cfg(test)]
use crate::ChatMessage;
#[cfg(test)]
pub(crate) fn test_request(
messages: Vec<ChatMessage>,
tools: Option<Vec<crate::ToolSpec>>,
) -> ChatRequest {
ChatRequest {
messages,
tools,
model: "test".to_string(),
allow_image_parts: false,
max_tokens: None,
reasoning_effort: None,
provider_order: None,
meta: None,
}
}
use std::sync::{Arc, RwLock};
use std::time::Instant;
pub(crate) use crate::providers::transcribe::{MediaTranscriber, transcribe_video_file};
use crate::retry::{FailureClass, RetryFailureRecord};
use compatible::OpenAiCompatibleProvider;
#[derive(Debug)]
pub(crate) struct ScopedCallError {
pub inner: anyhow::Error,
pub record: RetryFailureRecord,
pub class: FailureClass,
}
impl ScopedCallError {
#[must_use]
pub(crate) fn new(
inner: anyhow::Error,
record: RetryFailureRecord,
class: FailureClass,
) -> Self {
Self {
inner,
record,
class,
}
}
}
impl std::fmt::Display for ScopedCallError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.inner)
}
}
impl std::error::Error for ScopedCallError {}
#[must_use]
pub(crate) fn failure_class(class: reliable::ErrorClass, truncated: bool) -> FailureClass {
match class {
reliable::ErrorClass::NonRetryable => FailureClass::NonRetryable,
reliable::ErrorClass::Retryable if truncated => FailureClass::TruncatedEnvelope,
reliable::ErrorClass::Retryable => FailureClass::Transport,
}
}
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 {
let trimmed = endpoint.trim_end_matches('/');
trimmed
.strip_suffix("/chat/completions")
.unwrap_or(trimmed)
.to_string()
}
pub(crate) fn provider_routing_json(order: &str) -> 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": false,
}))
}
static PROVIDER: RwLock<Option<Arc<dyn Provider>>> = RwLock::new(None);
static MEDIA_TRANSCRIBER: RwLock<Option<MediaTranscriber>> = RwLock::new(None);
fn build_provider_and_transcriber(
config: &crate::config::ConfigData,
) -> (Arc<dyn Provider>, Option<MediaTranscriber>) {
let provider: Arc<dyn Provider> = create_provider(
config.provider_key.as_deref(),
Some(crate::config::DEFAULT_PROVIDER_ENDPOINT),
)
.into();
let media_transcriber = build_media_transcriber(config);
(provider, media_transcriber)
}
pub fn init_global() -> anyhow::Result<()> {
let config = CONFIG.snapshot();
let (provider, media_transcriber) = build_provider_and_transcriber(&config);
*PROVIDER.write().unwrap_poison() = Some(provider.clone());
*MEDIA_TRANSCRIBER.write().unwrap_poison() = media_transcriber;
let endpoint_str = crate::config::DEFAULT_PROVIDER_ENDPOINT;
tokio::spawn(async move {
if let Err(e) = provider.warmup().await {
tracing::warn!(endpoint = %endpoint_str, "Provider warmup failed (non-fatal): {e}");
}
});
Ok(())
}
pub(crate) async fn warmup_provider_from_config(
config: &crate::config::ConfigData,
) -> anyhow::Result<()> {
let provider = create_provider(
config.provider_key.as_deref(),
Some(crate::config::DEFAULT_PROVIDER_ENDPOINT),
);
provider.warmup().await?;
Ok(())
}
pub(crate) async fn recreate_all(config: &crate::config::ConfigData) -> anyhow::Result<()> {
let (provider, media_transcriber) = build_provider_and_transcriber(config);
provider.warmup().await?;
*PROVIDER.write().unwrap_poison() = Some(provider);
*MEDIA_TRANSCRIBER.write().unwrap_poison() = media_transcriber;
tracing::info!("Provider and transcriber singletons recreated");
let use_local = config.audio_transcription_use_local.as_deref() != Some("false");
if use_local && !crate::audio::local_transcriber::is_loaded() {
if crate::audio::local_transcriber::try_init_from_cache().await {
tracing::info!("Local Qwen3-ASR transcriber loaded from cache after config reload");
} else {
tracing::info!(
"Local Qwen3-ASR transcriber will be downloaded in background after config reload"
);
}
}
Ok(())
}
pub(crate) fn recreate_media_transcriber() {
let config = CONFIG.snapshot();
let transcriber = build_media_transcriber(&config);
*MEDIA_TRANSCRIBER.write().unwrap_poison() = transcriber;
tracing::info!("Media transcriber recreated from updated config");
}
#[must_use]
pub(crate) fn media_transcriber() -> Option<MediaTranscriber> {
MEDIA_TRANSCRIBER.read().unwrap_poison().clone()
}
pub(crate) async fn chat_scoped(
request: ChatRequest,
idle_timeout: std::time::Duration,
deadline: Instant,
) -> Result<ChatResponse, ScopedCallError> {
let provider = PROVIDER
.read()
.unwrap_poison()
.clone()
.expect("PROVIDER not initialized");
provider.chat_scoped(request, idle_timeout, deadline).await
}
#[cfg(test)]
pub(crate) fn swap_provider_for_test(provider: Arc<dyn Provider>) -> Option<Arc<dyn Provider>> {
let mut guard = PROVIDER.write().unwrap_poison();
let previous = guard.clone();
*guard = Some(provider);
previous
}
#[cfg(test)]
pub(crate) fn restore_provider_for_test(previous: Option<Arc<dyn Provider>>) {
*PROVIDER.write().unwrap_poison() = previous;
}
pub(crate) fn create_provider(api_key: Option<&str>, endpoint: Option<&str>) -> 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 headers = std::collections::HashMap::new();
headers.insert("X-Title".to_string(), "MahBot".to_string());
headers.insert(
"HTTP-Referrer".to_string(),
"https://github.com/edezhic/mahbot".to_string(),
);
let base = OpenAiCompatibleProvider::new("OpenRouter", base_url.as_str(), resolved_key)
.with_extra_headers(headers);
Box::new(base)
}
#[must_use]
fn create_transcriber(
api_url: Option<&str>,
api_key: Option<&str>,
model: Option<&str>,
provider: Option<&str>,
) -> Option<MediaTranscriber> {
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();
Some(MediaTranscriber::new(base_url, model, route))
}
#[must_use]
fn build_media_transcriber(config: &crate::config::ConfigData) -> Option<MediaTranscriber> {
let model = resolve_or(
config.multimodal_model.clone(),
crate::config::DEFAULT_MULTIMODAL_MODEL,
);
let route = config
.model_routings
.iter()
.find(|mr| mr.model == model)
.and_then(|mr| mr.provider_order.clone());
create_transcriber(
Some(crate::config::DEFAULT_PROVIDER_ENDPOINT),
config.provider_key.as_deref(),
Some(model.as_str()),
route.as_deref(),
)
}
#[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",
},
Case {
name: "repeated_suffix",
input: "https://api.example.com/v1/chat/completions/chat/completions",
expected_chat: "https://api.example.com/v1/chat/completions/chat/completions",
expected_base: "https://api.example.com/v1/chat/completions",
},
];
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,
expected: Option<serde_json::Value>,
}
let cases = [
Case {
name: "single_provider",
order: "openai",
expected: Some(serde_json::json!({
"order": ["openai"],
"allow_fallbacks": false,
})),
},
Case {
name: "multiple_providers",
order: "openai, anthropic, google",
expected: Some(serde_json::json!({
"order": ["openai", "anthropic", "google"],
"allow_fallbacks": false,
})),
},
Case {
name: "whitespace_only_yields_none",
order: " , , ",
expected: None,
},
Case {
name: "empty_string_yields_none",
order: "",
expected: None,
},
Case {
name: "leading_trailing_whitespace",
order: " openai ",
expected: Some(serde_json::json!({
"order": ["openai"],
"allow_fallbacks": false,
})),
},
Case {
name: "single_slug_survives_split",
order: "google-gemini",
expected: Some(serde_json::json!({
"order": ["google-gemini"],
"allow_fallbacks": false,
})),
},
];
for c in &cases {
assert_eq!(
provider_routing_json(c.order),
c.expected,
"case '{}': provider_routing_json({:?})",
c.name,
c.order,
);
}
}
}