use std::collections::HashMap;
use std::sync::Arc;
use anyhow::Result;
use clap::Parser;
use tokio::sync::Mutex;
use tracing_subscriber::EnvFilter;
use tuitbot_core::auth::passphrase;
use tuitbot_core::config::Config;
use tuitbot_core::content::ContentGenerator;
use tuitbot_core::context::semantic_index::SemanticIndex;
use tuitbot_core::llm::embedding_factory::create_embedding_provider;
use tuitbot_core::llm::factory::create_provider;
use tuitbot_core::storage;
use tuitbot_core::storage::accounts::DEFAULT_ACCOUNT_ID;
use tuitbot_core::x_api::scraper_health::new_scraper_health;
use tokio_util::sync::CancellationToken;
use tuitbot_core::automation::WatchtowerLoop;
use tuitbot_core::net::local_ip;
use tuitbot_server::auth;
use tuitbot_server::state::AppState;
use tuitbot_server::ws::AccountWsEvent;
#[derive(Parser)]
#[command(name = "tuitbot-server", version, about)]
struct Cli {
#[arg(long, default_value = "3001")]
port: u16,
#[arg(long, default_value = "127.0.0.1")]
host: String,
#[arg(long, default_value = "~/.tuitbot/config.toml")]
config: String,
#[arg(long)]
reset_passphrase: bool,
}
#[tokio::main]
async fn main() -> Result<()> {
tracing_subscriber::fmt()
.with_env_filter(EnvFilter::try_from_default_env().unwrap_or_else(|_| "info".into()))
.init();
let cli = Cli::parse();
let config_path = std::path::PathBuf::from(storage::expand_tilde(&cli.config));
let db_dir = config_path
.parent()
.unwrap_or_else(|| std::path::Path::new("."));
let db_path = db_dir.join("tuitbot.db");
if cli.reset_passphrase {
let new_passphrase = passphrase::reset_passphrase(db_dir)?;
println!("{new_passphrase}");
return Ok(());
}
tracing::info!(
db = %db_path.display(),
host = %cli.host,
port = cli.port,
"starting tuitbot server"
);
let pool = storage::init_db(&db_path.to_string_lossy()).await?;
storage::accounts::ensure_default_account(&pool).await?;
let api_token = auth::ensure_api_token(db_dir)?;
tracing::info!(token_path = %db_dir.join("api_token").display(), "API token ready");
let passphrase_hash = if cli.host == "0.0.0.0" {
match passphrase::ensure_passphrase(db_dir)? {
Some(new_passphrase) => {
println!("\n Web login passphrase: {new_passphrase}");
println!(" (save this — it won't be shown again)\n");
}
None => {
tracing::info!("Passphrase already configured");
}
}
passphrase::load_passphrase_hash(db_dir)?
} else {
match passphrase::load_passphrase_hash(db_dir)? {
Some(hash) => {
tracing::info!("Passphrase loaded from disk");
Some(hash)
}
None => {
tracing::info!("No passphrase configured — awaiting browser claim");
None
}
}
};
let passphrase_hash_mtime = passphrase::passphrase_hash_mtime(db_dir);
let (event_tx, _) = tokio::sync::broadcast::channel::<AccountWsEvent>(256);
let data_dir = db_dir.to_path_buf();
let loaded_config = Config::load(Some(&cli.config)).ok();
let bind_host = if cli.host != "127.0.0.1" {
cli.host.clone()
} else {
loaded_config
.as_ref()
.map(|c| c.server.host.clone())
.unwrap_or_else(|| cli.host.clone())
};
let bind_port = if cli.port != 3001 {
cli.port
} else {
loaded_config
.as_ref()
.map(|c| c.server.port)
.unwrap_or(cli.port)
};
let content_generator = match Config::load(Some(&cli.config)) {
Ok(config) => match create_provider(&config.llm) {
Ok(provider) => {
tracing::info!("LLM provider initialized for AI assist endpoints");
Some(Arc::new(ContentGenerator::new(provider, config.business)))
}
Err(e) => {
tracing::info!(error = %e, "LLM provider not configured — AI assist endpoints disabled");
None
}
},
Err(e) => {
tracing::info!(error = %e, "Config not loaded — AI assist endpoints disabled");
None
}
};
let mut content_generators = HashMap::new();
if let Some(cg) = content_generator {
content_generators.insert(DEFAULT_ACCOUNT_ID.to_string(), cg);
}
let content_sources = loaded_config
.as_ref()
.map(|c| c.content_sources.clone())
.unwrap_or_default();
let connector_config = loaded_config
.as_ref()
.map(|c| c.connectors.clone())
.unwrap_or_default();
let deployment_mode = loaded_config
.as_ref()
.map(|c| c.deployment_mode.clone())
.unwrap_or_default();
let watchtower_cancel = {
let enabled_sources: Vec<_> = content_sources
.sources
.iter()
.filter(|s| {
if !deployment_mode.allows_source_type(&s.source_type) {
tracing::warn!(
source_type = %s.source_type,
deployment_mode = %deployment_mode,
"skipping content source incompatible with deployment mode"
);
return false;
}
s.is_enabled() && (s.path.is_some() || s.folder_id.is_some())
})
.collect();
if !enabled_sources.is_empty() {
let cancel = CancellationToken::new();
let watchtower = WatchtowerLoop::new(
pool.clone(),
content_sources.clone(),
connector_config.clone(),
data_dir.clone(),
);
let cancel_clone = cancel.clone();
tokio::spawn(async move {
watchtower.run(cancel_clone).await;
});
tracing::info!(sources = enabled_sources.len(), "Watchtower started");
Some(cancel)
} else {
None
}
};
let (embedding_provider, semantic_index) = match loaded_config
.as_ref()
.and_then(|c| c.embedding.as_ref())
.filter(|cfg| cfg.enabled)
{
Some(embedding_cfg) => match create_embedding_provider(embedding_cfg) {
Ok(provider) => {
let provider: Arc<dyn tuitbot_core::llm::embedding::EmbeddingProvider> =
Arc::from(provider);
let index = SemanticIndex::new(
provider.dimension(),
provider.model_id().to_string(),
50_000,
);
tracing::info!(
provider = provider.name(),
model = provider.model_id(),
dimension = provider.dimension(),
"Semantic search enabled"
);
(
Some(provider),
Some(Arc::new(tokio::sync::RwLock::new(index))),
)
}
Err(e) => {
tracing::warn!("Embedding provider not available, semantic search disabled: {e}");
(None, None)
}
},
None => (None, None),
};
let state = Arc::new(AppState {
db: pool,
config_path,
data_dir,
event_tx,
api_token,
passphrase_hash: tokio::sync::RwLock::new(passphrase_hash),
passphrase_hash_mtime: tokio::sync::RwLock::new(passphrase_hash_mtime),
bind_host: bind_host.clone(),
bind_port,
login_attempts: Mutex::new(HashMap::new()),
runtimes: Mutex::new(HashMap::new()),
content_generators: Mutex::new(content_generators),
circuit_breaker: None,
scraper_health: loaded_config
.as_ref()
.filter(|c| c.x_api.provider_backend == "scraper")
.map(|_| new_scraper_health()),
watchtower_cancel: tokio::sync::RwLock::new(watchtower_cancel),
content_sources: tokio::sync::RwLock::new(content_sources),
connector_config,
deployment_mode,
pending_oauth: Mutex::new(HashMap::new()),
token_managers: Mutex::new(HashMap::new()),
x_client_id: loaded_config
.as_ref()
.map(|c| c.x_api.client_id.clone())
.unwrap_or_default(),
semantic_index,
embedding_provider,
});
let router = tuitbot_server::build_router(state.clone());
if bind_host == "0.0.0.0" {
tracing::warn!("Binding to 0.0.0.0 — server accessible from LAN");
if let Some(ip) = local_ip() {
println!(" Dashboard: http://{}:{}", ip, bind_port);
}
}
let listener = tokio::net::TcpListener::bind(format!("{}:{}", bind_host, bind_port)).await?;
tracing::info!("listening on http://{}:{}", bind_host, bind_port);
axum::serve(listener, router).await?;
if let Some(cancel) = state.watchtower_cancel.read().await.as_ref() {
cancel.cancel();
}
Ok(())
}