308 lines
11 KiB
Rust
308 lines
11 KiB
Rust
pub mod agent_factory;
|
|
pub mod agent_prompt_provider;
|
|
pub mod agent_task_executor;
|
|
pub mod cancel_manager;
|
|
pub mod cli_session;
|
|
pub mod command;
|
|
pub mod compaction;
|
|
pub mod execution;
|
|
pub mod http;
|
|
pub mod memory_maintenance;
|
|
pub mod memory_maintenance_coordinator;
|
|
pub mod message_prepare;
|
|
pub mod outbound_dispatcher;
|
|
pub mod processor;
|
|
pub mod prompt;
|
|
pub mod provider_config_service;
|
|
pub mod runtime;
|
|
pub mod scheduled_agent_task_service;
|
|
pub mod session;
|
|
pub mod session_factory;
|
|
pub mod session_history;
|
|
pub mod session_lifecycle;
|
|
pub mod session_message_sender;
|
|
pub mod session_message_service;
|
|
pub mod session_pool;
|
|
pub mod static_files;
|
|
pub mod tool_registry_factory;
|
|
pub mod todo_prompt_provider;
|
|
pub mod ws;
|
|
|
|
use axum::{Router, routing};
|
|
use std::collections::HashMap;
|
|
use std::sync::Arc;
|
|
use tokio::net::TcpSocket;
|
|
use tokio::sync::Semaphore;
|
|
use tower_http::services::ServeDir;
|
|
|
|
use crate::bus::MessageBus;
|
|
use crate::channels::ChannelManager;
|
|
use crate::config::Config;
|
|
use crate::config::LLMProviderConfig;
|
|
use crate::logging;
|
|
use crate::scheduler::Scheduler;
|
|
use crate::skills::SkillRuntime;
|
|
use crate::tools::task::repository::TaskRepository;
|
|
use agent_task_executor::{AgentTaskExecutor, SchedulerMaintenanceService};
|
|
use cancel_manager::CancelManager;
|
|
use outbound_dispatcher::OutboundDispatcher;
|
|
use processor::InboundProcessor;
|
|
use runtime::build_session_manager_with_sender;
|
|
use session_message_sender::BusSessionMessageSender;
|
|
use session::SessionManager;
|
|
use static_files::static_handler;
|
|
|
|
use tokio::sync::{watch, RwLock};
|
|
|
|
pub struct GatewayState {
|
|
pub config: Arc<RwLock<Config>>,
|
|
pub session_manager: SessionManager,
|
|
pub channel_manager: ChannelManager,
|
|
pub bus: Arc<MessageBus>,
|
|
pub task_repository: Arc<dyn TaskRepository>,
|
|
pub cancel_manager: CancelManager,
|
|
pub restart_tx: watch::Sender<bool>,
|
|
pub mcp_manager: Option<Arc<crate::mcp::client::McpClientManager>>,
|
|
}
|
|
|
|
impl GatewayState {
|
|
pub fn from_config(config: Config, restart_tx: watch::Sender<bool>) -> Result<Self, Box<dyn std::error::Error>> {
|
|
// Get provider config for SessionManager
|
|
let provider_config = config.get_provider_config("default")?;
|
|
let mut provider_configs = HashMap::<String, LLMProviderConfig>::new();
|
|
for agent_name in config.agents.keys() {
|
|
provider_configs.insert(agent_name.clone(), config.get_provider_config(agent_name)?);
|
|
}
|
|
|
|
let agent_prompt_reinject_every = config.gateway.agent_prompt_reinject_every;
|
|
let show_tool_results = config.gateway.show_tool_results;
|
|
|
|
let session_ttl_hours = config.gateway.session_ttl_hours;
|
|
|
|
let skills = Arc::new(SkillRuntime::from_config(config.skills.clone()));
|
|
let channel_manager = ChannelManager::new();
|
|
let bus = channel_manager.bus();
|
|
|
|
let mcp_config = crate::mcp::McpConfig {
|
|
mcp_servers: config.mcp_servers.clone(),
|
|
};
|
|
|
|
let (session_manager, task_repository, mcp_manager) = build_session_manager_with_sender(
|
|
agent_prompt_reinject_every,
|
|
show_tool_results,
|
|
config.time.timezone.clone(),
|
|
provider_config,
|
|
provider_configs,
|
|
skills,
|
|
Arc::new(BusSessionMessageSender::new(bus.clone())),
|
|
std::collections::HashSet::new(),
|
|
config.tools.task.clone(),
|
|
config.subagents.clone(),
|
|
config.memory_maintenance.clone(),
|
|
session_ttl_hours,
|
|
mcp_config,
|
|
Some(bus.clone()),
|
|
)?;
|
|
|
|
let cancel_manager = CancelManager::new();
|
|
|
|
Ok(Self {
|
|
config: Arc::new(RwLock::new(config)),
|
|
session_manager,
|
|
channel_manager,
|
|
bus,
|
|
task_repository,
|
|
cancel_manager,
|
|
restart_tx,
|
|
mcp_manager,
|
|
})
|
|
}
|
|
|
|
/// Start the message processing loops
|
|
pub async fn start_message_processing(&self) {
|
|
let bus_for_outbound = self.bus.clone();
|
|
|
|
// Read config under read lock
|
|
let cfg = self.config.read().await;
|
|
let max_concurrent = cfg.gateway.max_concurrent_requests;
|
|
let provider_config = match cfg.get_provider_config("default") {
|
|
Ok(config) => config,
|
|
Err(e) => {
|
|
tracing::error!(error = %e, "Failed to get provider config");
|
|
return;
|
|
}
|
|
};
|
|
drop(cfg); // release read lock before spawning long-running tasks
|
|
|
|
let semaphore = Arc::new(Semaphore::new(max_concurrent));
|
|
let inbound_processor =
|
|
InboundProcessor::new(self.bus.clone(), self.session_manager.clone(), semaphore, provider_config, self.cancel_manager.clone());
|
|
tokio::spawn(inbound_processor.run());
|
|
|
|
// Spawn outbound dispatcher
|
|
let dispatcher = OutboundDispatcher::new(bus_for_outbound);
|
|
let channel_manager = self.channel_manager.clone();
|
|
|
|
for (name, channel) in channel_manager.channels().await {
|
|
dispatcher.register_channel(&name, channel).await;
|
|
}
|
|
|
|
tokio::spawn(async move {
|
|
tracing::info!("Outbound dispatcher started");
|
|
dispatcher.run().await;
|
|
});
|
|
}
|
|
}
|
|
|
|
pub async fn run(
|
|
host: Option<String>,
|
|
port: Option<u16>,
|
|
) -> Result<bool, Box<dyn std::error::Error>> {
|
|
let config = Config::load_default()?;
|
|
let timezone = config.time.parse_timezone()?;
|
|
|
|
// Initialize logging
|
|
logging::init_logging(timezone);
|
|
tracing::info!("Starting PicoBot Gateway");
|
|
|
|
// Restart signal channel
|
|
let (restart_tx, mut restart_rx) = watch::channel(false);
|
|
|
|
let state = Arc::new(GatewayState::from_config(config, restart_tx)?);
|
|
|
|
// Get provider config for channels
|
|
let cfg = state.config.read().await;
|
|
let provider_config = cfg.get_provider_config("default")?;
|
|
|
|
// Initialize and start channels
|
|
state
|
|
.channel_manager
|
|
.init(&*cfg, provider_config.clone())
|
|
.await?;
|
|
drop(cfg);
|
|
state.channel_manager.start_all().await?;
|
|
|
|
// Start message processing (inbound processor + outbound dispatcher)
|
|
state.start_message_processing().await;
|
|
|
|
let (scheduler_shutdown_tx, scheduler_shutdown_rx) = tokio::sync::watch::channel(false);
|
|
let scheduler_enabled = {
|
|
let cfg = state.config.read().await;
|
|
cfg.scheduler.enabled
|
|
};
|
|
if scheduler_enabled {
|
|
let scheduler_cfg = state.config.read().await.scheduler.clone();
|
|
let scheduler = Scheduler::new(
|
|
state.bus.clone(),
|
|
scheduler_cfg,
|
|
timezone,
|
|
state.session_manager.store(),
|
|
AgentTaskExecutor::new(state.session_manager.clone()),
|
|
SchedulerMaintenanceService::new(state.session_manager.clone()),
|
|
);
|
|
|
|
tokio::spawn(async move {
|
|
scheduler.run(scheduler_shutdown_rx).await;
|
|
});
|
|
}
|
|
|
|
// CLI args override config file values
|
|
let (bind_host, bind_port) = {
|
|
let cfg = state.config.read().await;
|
|
let h = host.unwrap_or_else(|| cfg.gateway.host.clone());
|
|
let p = port.unwrap_or(cfg.gateway.port);
|
|
(h, p)
|
|
};
|
|
|
|
// 使用嵌入的静态文件(编译时打包进二进制)
|
|
// 开发模式下可通过 STATIC_DIR 环境变量使用磁盘文件
|
|
let use_embedded = std::env::var("STATIC_DIR").is_err();
|
|
|
|
let app = if use_embedded {
|
|
Router::new()
|
|
.route("/health", routing::get(http::health))
|
|
.route("/api/config", routing::get(http::get_config).put(http::save_config))
|
|
.route("/api/restart", routing::post(http::restart))
|
|
.route("/api/mcp/status", routing::get(http::mcp_status))
|
|
.route("/ws", routing::get(ws::ws_handler))
|
|
.fallback(static_handler)
|
|
.with_state(state.clone())
|
|
} else {
|
|
let static_dir = std::env::var("STATIC_DIR").unwrap_or_else(|_| "static".to_string());
|
|
Router::new()
|
|
.route("/health", routing::get(http::health))
|
|
.route("/api/config", routing::get(http::get_config).put(http::save_config))
|
|
.route("/api/restart", routing::post(http::restart))
|
|
.route("/api/mcp/status", routing::get(http::mcp_status))
|
|
.route("/ws", routing::get(ws::ws_handler))
|
|
.fallback_service(ServeDir::new(&static_dir))
|
|
.with_state(state.clone())
|
|
};
|
|
|
|
let addr: std::net::SocketAddr = format!("{}:{}", bind_host, bind_port).parse()?;
|
|
let listener = {
|
|
let socket = match addr {
|
|
std::net::SocketAddr::V4(_) => TcpSocket::new_v4()?,
|
|
std::net::SocketAddr::V6(_) => TcpSocket::new_v6()?,
|
|
};
|
|
socket.set_reuseaddr(true)?;
|
|
socket.bind(addr)?;
|
|
socket.listen(1024)?
|
|
};
|
|
tracing::info!(address = %addr, "Gateway listening");
|
|
|
|
// Graceful shutdown / restart signal
|
|
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
|
|
// Side channel to communicate whether this was a restart or shutdown
|
|
let (result_tx, result_rx) = tokio::sync::oneshot::channel::<bool>();
|
|
let channel_manager = state.channel_manager.clone();
|
|
let cancel_manager = state.cancel_manager.clone();
|
|
let mcp_manager = state.mcp_manager.clone();
|
|
|
|
// Spawn ctrl_c / restart handler
|
|
tokio::spawn(async move {
|
|
tokio::select! {
|
|
_ = tokio::signal::ctrl_c() => {
|
|
tracing::info!("Shutdown signal received");
|
|
cancel_manager.cancel_all().await;
|
|
let _ = scheduler_shutdown_tx.send(true);
|
|
if let Some(ref mgr) = mcp_manager {
|
|
tracing::info!("Shutting down MCP servers before shutdown");
|
|
let _ = mgr.shutdown_all().await;
|
|
}
|
|
let _ = channel_manager.stop_all().await;
|
|
let _ = result_tx.send(false);
|
|
let _ = shutdown_tx.send(());
|
|
}
|
|
_ = restart_rx.changed() => {
|
|
if *restart_rx.borrow() {
|
|
tracing::info!("Restart signal received");
|
|
cancel_manager.cancel_all().await;
|
|
let _ = scheduler_shutdown_tx.send(true);
|
|
if let Some(ref mgr) = mcp_manager {
|
|
tracing::info!("Shutting down MCP servers before restart");
|
|
let _ = mgr.shutdown_all().await;
|
|
}
|
|
let _ = channel_manager.stop_all().await;
|
|
let _ = result_tx.send(true);
|
|
let _ = shutdown_tx.send(());
|
|
}
|
|
}
|
|
}
|
|
});
|
|
|
|
// Serve with graceful shutdown
|
|
axum::serve(listener, app)
|
|
.with_graceful_shutdown(async {
|
|
let _ = shutdown_rx.await;
|
|
})
|
|
.await?;
|
|
|
|
// Wait briefly for in-flight requests to complete
|
|
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
|
|
|
|
// Check if this was a restart
|
|
let should_restart = result_rx.await.unwrap_or(false);
|
|
Ok(should_restart)
|
|
}
|