use crate::{ChannelMessage, Role};
use std::sync::atomic::{AtomicBool, Ordering};
static KICKOFF_FIRED: AtomicBool = AtomicBool::new(false);
#[derive(Debug)]
pub enum ProviderInput {
OpenRouterKey(String),
CustomEndpoint { url: String, key: Option<String> },
Invalid,
}
#[must_use]
pub fn parse_provider_input(raw: &str) -> ProviderInput {
let trimmed = raw.trim();
let lines: Vec<&str> = trimmed
.lines()
.map(str::trim)
.filter(|l| !l.is_empty())
.collect();
if lines.len() >= 2 && is_http_url(lines[0]) {
let key = lines[1].to_string();
if crate::config::is_default_endpoint(lines[0]) && key.contains("sk-or-v1-") {
return ProviderInput::OpenRouterKey(key);
}
if crate::config::is_default_endpoint(lines[0]) {
return ProviderInput::Invalid;
}
if !is_clean_url_line(lines[0]) {
return ProviderInput::Invalid;
}
return ProviderInput::CustomEndpoint {
url: lines[0].to_string(),
key: Some(key),
};
}
if lines.len() == 1 && is_http_url(lines[0]) {
if crate::config::is_default_endpoint(lines[0]) {
return ProviderInput::Invalid;
}
if !is_clean_url_line(lines[0]) {
return ProviderInput::Invalid;
}
return ProviderInput::CustomEndpoint {
url: lines[0].to_string(),
key: None,
};
}
if let Some(key_line) = lines.iter().find(|l| l.contains("sk-or-v1-")) {
return ProviderInput::OpenRouterKey(key_line.to_string());
}
ProviderInput::Invalid
}
fn is_http_url(s: &str) -> bool {
let lower = s.to_ascii_lowercase();
lower.starts_with("http://") || lower.starts_with("https://")
}
fn is_clean_url_line(s: &str) -> bool {
s.split_whitespace().count() == 1
}
pub async fn persist_provider_input(input: &ProviderInput) -> anyhow::Result<()> {
match input {
ProviderInput::OpenRouterKey(key) => {
crate::config::persist_settled_string_field("provider_key", key).await?;
}
ProviderInput::CustomEndpoint { url, key } => {
if let Some(k) = key {
crate::config::persist_settled_string_field("provider_endpoint_key", k).await?;
}
crate::config::persist_settled_string_field("provider_endpoint", url).await?;
}
ProviderInput::Invalid => anyhow::bail!("invalid provider input"),
}
Ok(())
}
#[must_use]
pub fn intro_messages() -> [&'static str; 2] {
[
"Welcome to MahBot! Before we start, I need an LLM provider to power your agents.",
"Enter your OpenRouter API key (it starts with `sk-or-v1-`), or paste a custom endpoint. For a custom endpoint, put the URL on the FIRST line and your API key on the SECOND line (the key is optional).",
]
}
#[must_use]
pub fn success_message() -> &'static str {
"Provider configured! Setting up your Support assistant — one moment."
}
#[must_use]
pub fn invalid_message() -> &'static str {
"That doesn't look right. Enter an OpenRouter key (starts with `sk-or-v1-`) or a custom endpoint URL (first line) + optional API key (second line)."
}
pub async fn kickoff_support(user_name: &str) -> anyhow::Result<()> {
use crate::config::OnboardingState;
if crate::config::CONFIG.onboarding_stage() != OnboardingState::Init {
return Ok(());
}
if !crate::config::provider_configured() {
return Ok(());
}
let pool = crate::users::role_pool(user_name).await;
if !pool.contains(&Role::Support) {
return Ok(());
}
if KICKOFF_FIRED.swap(true, Ordering::SeqCst) {
return Ok(());
}
if let Err(e) = crate::users::switch_active_role(user_name, Role::Support).await {
KICKOFF_FIRED.store(false, Ordering::SeqCst);
return Err(e);
}
if let Err(e) = crate::config::persist_settled_string_field(
crate::config::CONFIG_KEY_ONBOARDING_STATE,
OnboardingState::Welcomed.as_str(),
)
.await
{
KICKOFF_FIRED.store(false, Ordering::SeqCst);
return Err(e);
}
let msg = ChannelMessage {
user_name: user_name.to_string(),
reply_target: user_name.to_string(),
content: "hi mah bot".to_string(),
channel: "gui".to_string(),
workspace: format!("personal:{user_name}"),
optimistic_id: None,
callback_query_id: None,
reply_reference: None,
};
if let Some(tx) = crate::GUI_MESSAGE_TX.get()
&& let Err(e) = tx.send(msg)
{
tracing::error!("kickoff: failed to send 'hi mah bot' via GUI_MESSAGE_TX: {e}");
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_openrouter_key() {
match parse_provider_input(" sk-or-v1-abc123def456 ") {
ProviderInput::OpenRouterKey(k) => assert_eq!(k, "sk-or-v1-abc123def456"),
other => panic!("expected OpenRouterKey, got {other:?}"),
}
}
#[test]
fn parses_keyless_custom_endpoint_url() {
match parse_provider_input("https://ollama.local:11434/v1") {
ProviderInput::CustomEndpoint { url, key } => {
assert_eq!(url, "https://ollama.local:11434/v1");
assert_eq!(key, None);
}
other => panic!("expected keyless CustomEndpoint, got {other:?}"),
}
}
#[test]
fn parses_two_line_custom_endpoint_with_key() {
match parse_provider_input("http://localhost:8080/v1\nsk-local-123") {
ProviderInput::CustomEndpoint { url, key } => {
assert_eq!(url, "http://localhost:8080/v1");
assert_eq!(key.as_deref(), Some("sk-local-123"));
}
other => panic!("expected CustomEndpoint with key, got {other:?}"),
}
}
#[test]
fn parses_default_openrouter_url_with_key_as_openrouter() {
match parse_provider_input("https://openrouter.ai/api/v1\nsk-or-v1-abc123def456") {
ProviderInput::OpenRouterKey(k) => assert_eq!(k, "sk-or-v1-abc123def456"),
other => panic!("expected OpenRouterKey, got {other:?}"),
}
}
#[test]
fn parses_reversed_key_then_url_as_openrouter_key() {
match parse_provider_input("sk-or-v1-abc123def456\nhttps://custom.example/v1") {
ProviderInput::OpenRouterKey(k) => assert_eq!(k, "sk-or-v1-abc123def456"),
other => panic!("expected OpenRouterKey, got {other:?}"),
}
}
#[test]
fn rejects_bare_default_openrouter_url() {
assert!(matches!(
parse_provider_input("https://openrouter.ai/api/v1"),
ProviderInput::Invalid
));
}
#[test]
fn rejects_single_line_url_with_embedded_junk() {
assert!(matches!(
parse_provider_input("https://custom.example/v1 sk-local-123"),
ProviderInput::Invalid
));
assert!(matches!(
parse_provider_input("https://custom.example/v1 extra"),
ProviderInput::Invalid
));
}
#[test]
fn rejects_default_openrouter_url_with_unknown_key() {
assert!(matches!(
parse_provider_input("https://openrouter.ai/api/v1\nmy-random-key"),
ProviderInput::Invalid
));
}
#[test]
fn rejects_unknown_input() {
assert!(matches!(
parse_provider_input("just some chat text"),
ProviderInput::Invalid
));
}
}