use crate::{
config::{ArchitectureType, Config},
error::{Error, Result},
middleware::MiddlewareManager,
server::WebServer,
};
use axum::Router;
use std::path::Path;
#[derive(Debug)]
pub struct WebServerBuilder {
config: Config,
router: Option<Router>,
custom_middleware: Vec<Box<dyn MiddlewareFactory>>,
}
pub trait MiddlewareFactory: Send + Sync + std::fmt::Debug {
fn create_layer(&self) -> tower::layer::util::Identity;
fn name(&self) -> &str;
}
impl Default for WebServerBuilder {
fn default() -> Self {
Self::new()
}
}
impl WebServerBuilder {
pub fn new() -> Self {
Self {
config: Config::default(),
router: None,
custom_middleware: Vec::new(),
}
}
pub fn config_from_file<P: AsRef<Path>>(mut self, path: P) -> Self {
match Config::from_file(path) {
Ok(config) => {
self.config = config;
}
Err(e) => {
eprintln!("警告: 无法加载配置文件: {},使用默认配置", e);
}
}
self
}
pub fn config(mut self, config: Config) -> Self {
self.config = config;
self
}
pub fn listen(mut self, host: &str, port: u16) -> Self {
self.config.server.host = host.to_string();
self.config.server.port = port;
self
}
pub fn architecture(mut self, arch_type: ArchitectureType) -> Self {
self.config.server.architecture = arch_type;
self
}
pub fn cors(mut self, origins: Vec<String>) -> Self {
self.config.middleware.cors.enabled = true;
self.config.middleware.cors.origins = origins;
self
}
pub fn static_files<P: AsRef<Path>>(mut self, dir: P, prefix: &str) -> Self {
self.config.middleware.static_files.enabled = true;
self.config.middleware.static_files.dir = dir.as_ref().to_string_lossy().to_string();
self.config.middleware.static_files.prefix = prefix.to_string();
self
}
pub fn templates<P: AsRef<Path>>(mut self, dir: P, extension: &str) -> Self {
if self.config.server.architecture == ArchitectureType::Full {
self.config.middleware.templates.enabled = true;
self.config.middleware.templates.dir = dir.as_ref().to_string_lossy().to_string();
self.config.middleware.templates.extension = extension.to_string();
}
self
}
pub fn jwt_auth(mut self, secret: &str, expires_in: u64) -> Self {
self.config.middleware.jwt.enabled = true;
self.config.middleware.jwt.secret = secret.to_string();
self.config.middleware.jwt.expires_in = expires_in;
self
}
pub fn log_level(mut self, level: &str) -> Self {
self.config.middleware.logging.level = level.to_string();
self
}
pub fn routes(mut self, router: Router) -> Self {
self.router = Some(router);
self
}
pub fn middleware<M: MiddlewareFactory + 'static>(mut self, middleware: M) -> Self {
self.custom_middleware.push(Box::new(middleware));
self
}
pub fn custom_config<T: serde::Serialize>(mut self, key: &str, value: T) -> Self {
if let Ok(json_value) = serde_json::to_value(value) {
self.config.middleware.custom.insert(key.to_string(), json_value);
}
self
}
pub async fn build(self) -> Result<WebServer> {
self.config.validate()?;
self.init_logging()?;
#[cfg(feature = "database")]
let database = if self.config.database.enabled {
use crate::database::Database;
match Database::from_config(self.config.database.clone()).await {
Ok(db) => {
if let Err(e) = db.ping().await {
tracing::warn!("数据库连接测试失败: {}", e);
}
Some(db)
}
Err(e) => {
tracing::error!("数据库初始化失败: {}", e);
return Err(e);
}
}
} else {
None
};
#[cfg(feature = "vector")]
let vector_db = if self.config.qdrant.enabled {
use crate::database::VectorDatabase;
match VectorDatabase::from_config(self.config.qdrant.clone()).await {
Ok(db) => {
if let Err(e) = db.ping().await {
tracing::warn!("向量数据库连接测试失败: {}", e);
}
Some(db)
}
Err(e) => {
tracing::error!("向量数据库初始化失败: {}", e);
return Err(e);
}
}
} else {
None
};
let mut middleware_manager = MiddlewareManager::new(self.config.clone());
for middleware in self.custom_middleware {
middleware_manager.add_custom_middleware(middleware);
}
let base_router = self.router.unwrap_or_default();
let app = middleware_manager.apply_middleware(base_router).await?;
let server = WebServer::new(app, self.config);
#[cfg(feature = "database")]
let server = if let Some(db) = database {
server.with_database(db)
} else {
server
};
#[cfg(feature = "vector")]
let server = if let Some(vdb) = vector_db {
server.with_vector_db(vdb)
} else {
server
};
Ok(server)
}
fn init_logging(&self) -> Result<()> {
use tracing_subscriber::{EnvFilter, fmt, prelude::*};
let filter = EnvFilter::try_from_default_env()
.or_else(|_| EnvFilter::try_new(&self.config.middleware.logging.level))
.map_err(|e| Error::Config(format!("无效的日志级别: {}", e)))?;
let fmt_layer = fmt::layer()
.with_target(false)
.with_thread_ids(false)
.with_file(false)
.with_line_number(false);
let _ = tracing_subscriber::registry()
.with(filter)
.with(fmt_layer)
.try_init();
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ArchitectureType;
#[test]
fn test_builder_creation() {
let builder = WebServerBuilder::new();
assert_eq!(builder.config.server.port, 3000);
assert_eq!(builder.config.server.architecture, ArchitectureType::Api);
}
#[test]
fn test_builder_configuration() {
let builder = WebServerBuilder::new()
.listen("127.0.0.1", 8080)
.architecture(ArchitectureType::Full)
.cors(vec!["http://localhost:3000".to_string()])
.log_level("debug");
assert_eq!(builder.config.server.host, "127.0.0.1");
assert_eq!(builder.config.server.port, 8080);
assert_eq!(builder.config.server.architecture, ArchitectureType::Full);
assert!(builder.config.middleware.cors.enabled);
assert_eq!(builder.config.middleware.logging.level, "debug");
}
}