pub mod auth; pub mod http; pub(crate) mod reload; mod router; pub mod uploads; pub mod ws; use axum::{Router, middleware, routing}; use std::net::SocketAddr; use std::sync::Arc; use tokio::net::TcpListener; use crate::bus::{MessageBus, OutboundDispatcher}; use crate::channels::{ChannelManager, CliChatChannel}; use crate::config::{Config, ensure_workspace_dir, expand_path}; use crate::delivery::{ConversationWriteLocks, DeliveryCoordinator, TurnDeliveryService}; use crate::logging; use crate::mcp; use crate::memory::MemoryManager; use crate::scheduler::Scheduler; use crate::session::{SessionManager, SessionManagerServices}; use crate::task_supervisor::TaskSupervisor; pub struct GatewayState { pub config: Config, pub config_path: std::path::PathBuf, pub workspace_dir: std::path::PathBuf, pub session_manager: Arc, pub channel_manager: ChannelManager, pub storage: Arc, pub task_supervisor: TaskSupervisor, pub delivery_coordinator: DeliveryCoordinator, pub connection_shutdown: tokio_util::sync::CancellationToken, pub auth: auth::AuthManager, pub uploads: uploads::UploadRegistry, pub(crate) reload: reload::ReloadHandle, pub(crate) admission: reload::RuntimeAdmission, } impl GatewayState { /// Construct a standalone state. Configuration reload is available only /// when the state is owned by [`run`], which owns the generation loop. pub async fn new() -> Result> { let config_path = crate::config::resolve_default_config_path(); let config = Config::load_from(&config_path)?; Self::from_config( config, config_path, reload::ReloadHandle::unavailable(), true, ) .await } async fn from_config( config: Config, config_path: std::path::PathBuf, reload: reload::ReloadHandle, initialize_process: bool, ) -> Result> { let task_supervisor = TaskSupervisor::new(); let admission = reload::RuntimeAdmission::open(); let delivery_coordinator = DeliveryCoordinator::new(ConversationWriteLocks::default()); let connection_shutdown = tokio_util::sync::CancellationToken::new(); let auth = auth::AuthManager::load( config.gateway.require_pairing, crate::config::get_user_config_dir().join("web_auth.json"), ) .await?; let uploads = uploads::UploadRegistry::new(config.gateway.file_transfer.clone()); // Initialize workspace directory: expand path and ensure it exists let workspace_path = expand_path(&config.workspace_dir); let workspace_path = ensure_workspace_dir(&workspace_path)?; if initialize_process { // Startup is single-threaded. Reload candidates reuse the already // selected workspace and must not mutate process-global cwd. std::env::set_current_dir(&workspace_path).map_err(|e| { format!( "Failed to switch to workspace directory {}: {}", workspace_path.display(), e ) })?; } tracing::info!("Using workspace directory: {}", workspace_path.display()); // Release default AGENTS.md and USER.md to ~/.picobot/ if not exist if initialize_process { ensure_default_config_files(); } // Get provider config for SessionManager let mut provider_config = config.get_provider_config("default")?; // Override workspace_dir with the ensured path provider_config.workspace_dir = workspace_path.clone(); // Initialize Storage let db_path = if let Some(ref path) = config.gateway.session_db_path { std::path::PathBuf::from(path) } else { workspace_path.join("picobot.db") }; let storage = Arc::new( crate::storage::Storage::new(&db_path) .await .map_err(|e| format!("failed to initialize session storage: {}", e))?, ); tracing::info!("Session storage: {}", db_path.display()); // Resolve consolidation provider/model with fallback to main agent config let consolidation_provider = config .memory .resolve_consolidation_provider(&provider_config.name); let consolidation_model = config .memory .resolve_consolidation_model(&provider_config.model_id); let memory_manager = Arc::new(MemoryManager::new( storage.clone(), consolidation_provider, consolidation_model, )); tracing::info!( consolidation_provider = %memory_manager.consolidation_provider, consolidation_model = %memory_manager.consolidation_model, "Memory system initialized" ); // Create MessageBus first (shared by SessionManager and ChannelManager) let bus = MessageBus::new(100); // Channels are resolved by TurnDeliveryService, while Session workers // depend only on that protocol-neutral delivery facade. let cli_chat_channel = Arc::new(CliChatChannel::with_upload_registry(uploads.clone())); let channel_manager = ChannelManager::with_bus(cli_chat_channel, bus.clone()); channel_manager .init(&config, workspace_path.clone()) .await .map_err(|e| format!("Failed to init channels: {}", e))?; let turn_delivery = TurnDeliveryService::new( delivery_coordinator.clone(), channel_manager.clone(), task_supervisor.clone(), ); let browser_config = if config.browser.enabled { Some(config.browser.clone()) } else { None }; // Create SessionManager with bus injection let session_manager = SessionManager::new( provider_config.clone(), storage.clone(), SessionManagerServices::new( bus.clone(), memory_manager, task_supervisor.clone(), turn_delivery, reload.clone(), ) .with_admission(admission.clone()), browser_config, config.gateway.max_concurrent_background_tasks, )?; let session_manager = Arc::new(session_manager); // Register send_message tool with available channel names let available_channels = channel_manager.list_channel_names().await; let valid_channels = available_channels.clone(); session_manager.register_outbound_tool(available_channels); // Register chat_manager tool session_manager .tools() .register(crate::tools::ChatManagerTool::new( storage.clone(), valid_channels.clone(), )); // Initialize scheduler if enabled in config let scheduler_config = config.gateway.scheduler.clone().unwrap_or_default(); if scheduler_config.enabled { session_manager .tools() .register(crate::tools::RoutineMaintenanceTool::new( storage.clone(), config.memory.timeline_retention_days, )); storage .ensure_default_maintenance_job() .await .map_err(|e| format!("failed to seed default maintenance job: {e}"))?; // Register cron tools session_manager .tools() .register(crate::tools::cron::CronAddTool::new( storage.clone(), valid_channels, )); session_manager .tools() .register(crate::tools::cron::CronListTool::new(storage.clone())); session_manager .tools() .register(crate::tools::cron::CronRemoveTool::new(storage.clone())); session_manager .tools() .register(crate::tools::cron::CronEnableTool::new(storage.clone())); session_manager .tools() .register(crate::tools::cron::CronDisableTool::new(storage.clone())); session_manager .tools() .register(crate::tools::cron::CronUpdateTool::new(storage.clone())); tracing::info!("Cron tools registered"); } Ok(Self { config, config_path, workspace_dir: workspace_path, session_manager: session_manager.clone(), channel_manager, storage, task_supervisor, delivery_coordinator, connection_shutdown, auth, uploads, reload, admission, }) } /// Get a reference to the MessageBus pub fn bus(&self) -> Arc { self.channel_manager.bus() } /// Get CLI chat channel for WebSocket handling pub fn cli_chat_channel(&self) -> Arc { self.channel_manager.cli_chat_channel() } /// Start the message processing loops pub async fn start_message_processing(&self) { // MCP connections have external/process-wide side effects. Activate // them only after this generation becomes current, never while it is // merely a reload candidate. let mcp_tools = mcp::connect_all(&self.config.mcp).await; for tool_info in mcp_tools { let wrapper = mcp::McpToolWrapper::new( &tool_info.server_name, tool_info.tool_name, tool_info.description, tool_info.schema, tool_info.connection, ); self.session_manager.tools().register(wrapper); } let bus = self.bus(); let bus_for_outbound = bus.clone(); let session_manager = self.session_manager.clone(); if self.uploads.enabled() { let uploads = self.uploads.clone(); self.task_supervisor .spawn("pending-upload-cleanup", async move { let mut interval = tokio::time::interval(std::time::Duration::from_secs(300)); interval.tick().await; loop { interval.tick().await; let removed = uploads.cleanup_expired().await; if removed > 0 { tracing::debug!(removed, "Expired pending uploads removed"); } } }); } // Relay structured plan changes to WebSocket clients. This remains // separate from chat messages, so task UI updates never pollute history. let mut plan_events = self.session_manager.work_manager().subscribe(); let cli_chat = self.cli_chat_channel(); self.task_supervisor.spawn("task-plan-events", async move { loop { match plan_events.recv().await { Ok(event) => cli_chat.publish_plan_changed(event).await, Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { tracing::warn!(skipped, "Task plan event relay lagged"); } Err(tokio::sync::broadcast::error::RecvError::Closed) => break, } } }); router::spawn_message_routers( bus.clone(), session_manager, self.task_supervisor.clone(), self.admission.clone(), ); // Spawn outbound dispatcher let dispatcher = OutboundDispatcher::new( bus_for_outbound, self.channel_manager.clone(), self.task_supervisor.clone(), self.delivery_coordinator.write_locks(), ); self.task_supervisor .spawn("outbound-dispatcher", async move { tracing::info!("Outbound dispatcher started"); dispatcher.run().await; }); // Spawn scheduler background task if enabled let scheduler_config = self.config.gateway.scheduler.clone().unwrap_or_default(); if scheduler_config.enabled { let sched = Arc::new(Scheduler::with_admission( self.storage.clone(), self.session_manager.clone(), scheduler_config, self.admission.clone(), )); self.task_supervisor.spawn("scheduler", async move { sched.run().await; }); tracing::info!("Scheduler background task spawned"); } } } pub async fn run( host: Option, port: Option, ) -> Result<(), Box> { let config_path = crate::config::resolve_default_config_path(); let startup_process_env = Config::startup_process_env(); let startup_cwd = std::env::current_dir()?; let config = Config::load_from(&config_path)?; // Initialize logging logging::init_logging(); tracing::info!(config_path = %config_path.display(), "Starting PicoBot Gateway"); let mut reload_controller = reload::ReloadController::new(startup_process_env, startup_cwd); let mut state = Arc::new( GatewayState::from_config( config, config_path.clone(), reload_controller.handle.clone(), true, ) .await?, ); // CLI args override config file values let bind_host = host.unwrap_or_else(|| state.config.gateway.host.clone()); let bind_port = port.unwrap_or(state.config.gateway.port); let addr = format!("{}:{}", bind_host, bind_port); let listener = std::net::TcpListener::bind(&addr)?; listener.set_nonblocking(true)?; tracing::info!(address = %addr, "Gateway listening"); let process_signal = wait_for_shutdown_signal(); tokio::pin!(process_signal); let mut current_generation = 1_u64; loop { if let Err(error) = state.channel_manager.start_all().await { reload_controller.set_failed(current_generation, error.to_string()); return Err(error.into()); } state.start_message_processing().await; reload_controller.set_phase(current_generation, reload::ReloadPhase::Active); let app = build_router(state.clone()); let generation_listener = TcpListener::from_std(listener.try_clone()?)?; let generation_shutdown = tokio_util::sync::CancellationToken::new(); let shutdown_wait = generation_shutdown.clone(); let mut serve_task = tokio::spawn(async move { axum::serve( generation_listener, app.into_make_service_with_connect_info::(), ) .with_graceful_shutdown(async move { shutdown_wait.cancelled().await }) .await }); let mut next_state = None; let mut serve_result = None; 'generation: loop { tokio::select! { result = &mut serve_task => { serve_result = Some(result); state.connection_shutdown.cancel(); generation_shutdown.cancel(); break 'generation; } _ = &mut process_signal => { tracing::info!("Shutdown signal received"); state.admission.close(); state.connection_shutdown.cancel(); generation_shutdown.cancel(); break 'generation; } request = reload_controller.receiver.recv() => { let Some(request) = request else { state.admission.close(); state.connection_shutdown.cancel(); generation_shutdown.cancel(); break 'generation; }; let requested_generation = request.generation; reload_controller.set_phase(requested_generation, reload::ReloadPhase::Preparing); let candidate = match reload::load_candidate( &config_path, &reload_controller.startup_process_env, &reload_controller.startup_cwd, &state.config, &state.workspace_dir, ) { Ok(candidate) => candidate, Err(error) => { reload_controller.set_failed(requested_generation, error.to_string()); let _ = request.response.send(Err(error)); continue; } }; let preparation = GatewayState::from_config( candidate, config_path.clone(), reload_controller.handle.clone(), false, ); tokio::pin!(preparation); let prepared = match tokio::select! { result = &mut preparation => Some(result), result = &mut serve_task => { serve_result = Some(result); None } _ = &mut process_signal => None, } { None => { let error = reload::ReloadError::ShuttingDown; reload_controller.set_failed(requested_generation, error.to_string()); let _ = request.response.send(Err(error)); state.admission.close(); state.connection_shutdown.cancel(); generation_shutdown.cancel(); break 'generation; } Some(result) => match result { Ok(prepared) => Arc::new(prepared), Err(error) => { let error = reload::ReloadError::PreparationFailed(format!( "configuration reload failed: {error}" )); reload_controller.set_failed(requested_generation, error.to_string()); let _ = request.response.send(Err(error)); continue; } }}; state.admission.close(); reload_controller.set_phase(requested_generation, reload::ReloadPhase::Draining); let message = "配置校验通过;Gateway 将在当前任务结束后切换到新配置。".to_string(); let _ = request.response.send(Ok(reload::ReloadAccepted { generation: requested_generation, message, })); let drain = async { tokio::join!( state.admission.wait_for_idle(), state.session_manager.wait_until_idle(std::time::Duration::from_secs(60)), ) }; let drain_result = tokio::select! { result = tokio::time::timeout(std::time::Duration::from_secs(60), drain) => { Some(matches!(result, Ok(((), true)))) } result = &mut serve_task => { serve_result = Some(result); None } _ = &mut process_signal => None, }; let Some(drained) = drain_result else { reload_controller.set_failed( requested_generation, "gateway stopped while draining configuration reload", ); state.connection_shutdown.cancel(); generation_shutdown.cancel(); break 'generation; }; if !drained { tracing::warn!("Reload drain period ended before all work became idle"); } tracing::info!(path = %config_path.display(), "Switching to reloaded configuration"); reload_controller.set_phase(requested_generation, reload::ReloadPhase::Activating); state.connection_shutdown.cancel(); generation_shutdown.cancel(); next_state = Some((prepared, requested_generation)); break 'generation; } } } if serve_result.is_none() { serve_result = match tokio::time::timeout(std::time::Duration::from_secs(10), &mut serve_task) .await { Ok(result) => Some(result), Err(_) => { tracing::warn!("Aborting Axum generation after shutdown timeout"); serve_task.abort(); Some(serve_task.await) } }; } if let Err(error) = state.channel_manager.stop_all().await { tracing::error!(error = %error, "Failed to stop channels cleanly"); } state.task_supervisor.cancel(); state .task_supervisor .shutdown(std::time::Duration::from_secs(10)) .await; if let Some(result) = serve_result { match result { Ok(Ok(())) => {} Ok(Err(error)) => return Err(error.into()), Err(error) if error.is_cancelled() => {} Err(error) => return Err(format!("Gateway server task failed: {error}").into()), } } match next_state.take() { Some((prepared, generation)) => { state = prepared; current_generation = generation; } None => break, } } Ok(()) } fn build_router(state: Arc) -> Router { let protected = Router::new() .route("/api/health", routing::get(http::health)) .route( "/api/config", routing::get(http::get_config).put(http::put_config), ) .route("/api/config/reload", routing::post(http::reload_config)) .route( "/api/config/reload/status", routing::get(http::reload_status), ) .route( "/api/profiles/{name}", routing::get(http::get_profile).put(http::put_profile), ) .route("/api/logs", routing::get(http::get_logs)) .route("/api/tasks", routing::get(http::get_tasks)) .route("/api/jobs", routing::get(http::get_jobs)) .route("/api/jobs/{id}/runs", routing::get(http::get_job_runs)) .route("/api/memories", routing::get(http::get_memories)) .route( "/api/chat/{client_id}/uploads", routing::post(http::upload_file).layer(axum::extract::DefaultBodyLimit::disable()), ) .route( "/api/chat/{client_id}/sessions/{session_id}/messages/{message_id}/attachments/{index}", routing::get(http::download_attachment), ) .route("/ws", routing::get(ws::ws_handler)) .route_layer(middleware::from_fn_with_state( state.auth.clone(), auth::require_auth, )); Router::new() .route("/", routing::get(http::webui_index)) .route("/app.js", routing::get(http::webui_script)) .route("/styles.css", routing::get(http::webui_styles)) .route("/fonts/{name}", routing::get(http::webui_font)) .route("/health", routing::get(http::health)) .route("/api/auth/status", routing::get(auth::status)) .route("/api/auth/pair", routing::post(auth::pair)) .route("/api/auth/code", routing::post(auth::issue_code)) .merge(protected) .with_state(state) } async fn wait_for_shutdown_signal() { #[cfg(unix)] { use tokio::signal::unix::{SignalKind, signal}; match signal(SignalKind::terminate()) { Ok(mut terminate) => { tokio::select! { result = tokio::signal::ctrl_c() => { if let Err(error) = result { tracing::error!(error = %error, "Failed to listen for Ctrl-C"); } } _ = terminate.recv() => {} } } Err(error) => { tracing::error!(error = %error, "Failed to listen for SIGTERM"); if let Err(error) = tokio::signal::ctrl_c().await { tracing::error!(error = %error, "Failed to listen for Ctrl-C"); } } } } #[cfg(not(unix))] if let Err(error) = tokio::signal::ctrl_c().await { tracing::error!(error = %error, "Failed to listen for Ctrl-C"); } } /// Release default AGENTS.md and USER.md templates to ~/.picobot/ if not already present. fn ensure_default_config_files() { let picobot_dir = dirs::home_dir().unwrap_or_default().join(".picobot"); if let Err(e) = std::fs::create_dir_all(&picobot_dir) { tracing::warn!(dir = %picobot_dir.display(), error = %e, "Failed to create ~/.picobot directory"); return; } let agents_path = picobot_dir.join("AGENTS.md"); if !agents_path.exists() { let content = include_str!("../../resources/templates/AGENTS.md"); if let Err(e) = std::fs::write(&agents_path, content) { tracing::warn!(path = %agents_path.display(), error = %e, "Failed to write AGENTS.md template"); } else { tracing::info!(path = %agents_path.display(), "Released default AGENTS.md template"); } } let user_path = picobot_dir.join("USER.md"); if !user_path.exists() { let content = include_str!("../../resources/templates/USER.md"); if let Err(e) = std::fs::write(&user_path, content) { tracing::warn!(path = %user_path.display(), error = %e, "Failed to write USER.md template"); } else { tracing::info!(path = %user_path.display(), "Released default USER.md template"); } } let config_example_path = picobot_dir.join("config.example.json"); if !config_example_path.exists() { let content = include_str!("../../resources/templates/config.example.json"); if let Err(e) = std::fs::write(&config_example_path, content) { tracing::warn!(path = %config_example_path.display(), error = %e, "Failed to write config.example.json template"); } else { tracing::info!(path = %config_example_path.display(), "Released config.example.json template"); } } }