use crate::{config::Config, error::{Error, Result}};
use axum::Router;
use std::net::SocketAddr;
use tokio::net::TcpListener;
#[cfg(feature = "database")]
use crate::database::Database;
#[cfg(feature = "vector")]
use crate::database::VectorDatabase;
#[derive(Debug)]
pub struct WebServer {
app: Router,
config: Config,
#[cfg(feature = "database")]
database: Option<Database>,
#[cfg(feature = "vector")]
vector_db: Option<VectorDatabase>,
}
impl WebServer {
pub fn new(app: Router, config: Config) -> Self {
Self {
app,
config,
#[cfg(feature = "database")]
database: None,
#[cfg(feature = "vector")]
vector_db: None,
}
}
#[cfg(feature = "database")]
pub fn with_database(mut self, database: Database) -> Self {
self.database = Some(database);
self
}
#[cfg(feature = "vector")]
pub fn with_vector_db(mut self, vector_db: VectorDatabase) -> Self {
self.vector_db = Some(vector_db);
self
}
#[cfg(feature = "database")]
pub fn database(&self) -> Option<&Database> {
self.database.as_ref()
}
#[cfg(feature = "vector")]
pub fn vector_db(&self) -> Option<&VectorDatabase> {
self.vector_db.as_ref()
}
pub async fn run(self, addr: Option<&str>) -> Result<()> {
let default_addr = self.config.server_address();
let bind_addr = addr.unwrap_or(&default_addr);
tracing::info!("🚀 启动 HwhKit Web 服务器");
tracing::info!("📡 监听地址: {}", bind_addr);
tracing::info!("🏗️ 架构模式: {:?}", self.config.server.architecture);
self.log_middleware_status();
let socket_addr: SocketAddr = bind_addr.parse().map_err(|e| {
Error::ServerStart(format!("无效的地址格式 '{}': {}", bind_addr, e))
})?;
let listener = TcpListener::bind(socket_addr).await.map_err(|e| {
Error::ServerStart(format!("无法绑定到地址 '{}': {}", bind_addr, e))
})?;
tracing::info!("✅ 服务器启动成功,等待连接...");
axum::serve(listener, self.app).await.map_err(|e| {
Error::ServerStart(format!("服务器运行时错误: {}", e))
})?;
Ok(())
}
pub async fn serve(self) -> Result<()> {
self.run(None).await
}
pub fn config(&self) -> &Config {
&self.config
}
pub fn app(&self) -> &Router {
&self.app
}
fn log_middleware_status(&self) {
tracing::info!("🔧 中间件状态:");
if self.config.middleware.cors.enabled {
tracing::info!(" ✅ CORS: 已启用");
tracing::info!(" 📋 允许的源: {:?}", self.config.middleware.cors.origins);
} else {
tracing::info!(" ❌ CORS: 已禁用");
}
if self.config.middleware.static_files.enabled {
tracing::info!(" ✅ 静态文件: 已启用");
tracing::info!(" 📁 目录: {}", self.config.middleware.static_files.dir);
tracing::info!(" 🔗 前缀: {}", self.config.middleware.static_files.prefix);
} else {
tracing::info!(" ❌ 静态文件: 已禁用");
}
if self.config.middleware.templates.enabled {
tracing::info!(" ✅ 模板引擎: 已启用");
tracing::info!(" 📁 目录: {}", self.config.middleware.templates.dir);
} else {
tracing::info!(" ❌ 模板引擎: 已禁用");
}
if self.config.middleware.jwt.enabled {
tracing::info!(" ✅ JWT 认证: 已启用");
tracing::info!(" ⏰ 过期时间: {} 秒", self.config.middleware.jwt.expires_in);
} else {
tracing::info!(" ❌ JWT 认证: 已禁用");
}
if self.config.middleware.logging.requests {
tracing::info!(" ✅ 请求日志: 已启用");
tracing::info!(" 📊 级别: {}", self.config.middleware.logging.level);
} else {
tracing::info!(" ❌ 请求日志: 已禁用");
}
#[cfg(feature = "database")]
if self.config.database.enabled {
tracing::info!(" ✅ 关系数据库: 已启用");
tracing::info!(" 🗄️ 类型: {:?}", self.config.database.db_type);
tracing::info!(" 🔗 最大连接数: {}", self.config.database.max_connections);
} else {
tracing::info!(" ❌ 关系数据库: 已禁用");
}
#[cfg(feature = "vector")]
if self.config.qdrant.enabled {
tracing::info!(" ✅ 向量数据库: 已启用");
tracing::info!(" 🔍 URL: {}", self.config.qdrant.url);
tracing::info!(" 📦 默认集合: {}", self.config.qdrant.default_collection);
} else {
tracing::info!(" ❌ 向量数据库: 已禁用");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{ArchitectureType, Config};
use axum::{routing::get, Router};
async fn test_handler() -> &'static str {
"Hello, World!"
}
#[test]
fn test_web_server_creation() {
let app = Router::new().route("/", get(test_handler));
let config = Config::default();
let server = WebServer::new(app, config);
assert_eq!(server.config().server.port, 3000);
assert_eq!(server.config().server.architecture, ArchitectureType::Api);
}
#[test]
fn test_server_address_parsing() {
let app = Router::new();
let mut config = Config::default();
config.server.host = "127.0.0.1".to_string();
config.server.port = 8080;
let server = WebServer::new(app, config);
assert_eq!(server.config().server_address(), "127.0.0.1:8080");
}
}