360 lines
16 KiB
Python
360 lines
16 KiB
Python
"""
|
||
Configuration management for RAG system
|
||
"""
|
||
import json
|
||
import os
|
||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||
from typing import Optional, List, Dict, Any
|
||
from pydantic import Field
|
||
from loguru import logger
|
||
|
||
|
||
|
||
|
||
|
||
class BaseDataSourceConfig:
|
||
"""Base configuration class for all data sources"""
|
||
def __init__(self, name: str, type: str):
|
||
self.name = name # 数据源标识名称
|
||
self.type = type # 数据源类型: "database", "local_folder", "remote_folder"
|
||
|
||
|
||
class DatabaseDataSourceConfig(BaseDataSourceConfig):
|
||
"""Database data source configuration"""
|
||
def __init__(
|
||
self,
|
||
name: str,
|
||
# 通用数据库连接信息
|
||
database: str,
|
||
table_name: str = "documents",
|
||
host: Optional[str] = None,
|
||
port: Optional[int] = None,
|
||
user: Optional[str] = None,
|
||
password: Optional[str] = None,
|
||
|
||
db_type: Optional[str] = "mysql", # 新增:数据库类型,支持 "mysql", "dameng" 等
|
||
|
||
id_column: str = "id",
|
||
content_column: str = "content", # 可以是单个列名,或逗号分隔的多个列名
|
||
file_column: Optional[str] = None, # 单列名,指向表中的文件标识符字段
|
||
title_column: Optional[str] = "title",
|
||
metadata_columns: Optional[str] = None,
|
||
content_separator: str = "\n", # 多个 content 列拼接时的分隔符
|
||
updated_at_column: Optional[str] = None, # 用于增量同步的更新时间字段(可选)
|
||
|
||
# 文件源配置
|
||
file_source_type: Optional[str] = None, # 可选值: "api", "filesystem", "scp"
|
||
file_system_base_path: Optional[str] = None, # 文件系统基础路径
|
||
# SCP配置(可选)
|
||
scp_host: Optional[str] = None,
|
||
scp_port: Optional[int] = 22,
|
||
scp_username: Optional[str] = None,
|
||
scp_password: Optional[str] = None,
|
||
scp_key_path: Optional[str] = None
|
||
):
|
||
super().__init__(name, "database")
|
||
|
||
# 通用数据库连接信息
|
||
self.db_type = db_type # 数据库类型,mysql or dameng
|
||
self.database = database # 数据库名称
|
||
self.table_name = table_name
|
||
self.host = host
|
||
self.port = port
|
||
self.user = user
|
||
self.password = password
|
||
|
||
self.id_column = id_column
|
||
# 支持多个 content_column(逗号分隔)
|
||
if content_column:
|
||
self.content_columns = [col.strip() for col in content_column.split(",")]
|
||
else:
|
||
self.content_columns = []
|
||
# 保留原始 content_column 用于向后兼容
|
||
self.content_column = content_column
|
||
self.file_column = file_column
|
||
self.title_column = title_column
|
||
self.metadata_columns = metadata_columns
|
||
self.content_separator = content_separator # 多个列之间的分隔符
|
||
self.updated_at_column = updated_at_column # 更新时间字段(用于增量同步)
|
||
|
||
# 文件源配置
|
||
self.file_source_type = file_source_type
|
||
self.file_system_base_path = file_system_base_path
|
||
|
||
# SCP配置
|
||
self.scp_host = scp_host
|
||
self.scp_port = scp_port
|
||
self.scp_username = scp_username
|
||
self.scp_password = scp_password
|
||
|
||
|
||
class FolderDataSourceConfig(BaseDataSourceConfig):
|
||
"""Folder data source configuration (local or remote via SSH/SFTP)"""
|
||
def __init__(
|
||
self,
|
||
name: str,
|
||
folder_path: str,
|
||
host: str = "localhost",
|
||
port: int = 22,
|
||
username: Optional[str] = None,
|
||
password: Optional[str] = None,
|
||
recursive: bool = True,
|
||
ignore_patterns: Optional[List[str]] = None
|
||
):
|
||
super().__init__(name, "folder")
|
||
self.folder_path = folder_path # 文件夹路径
|
||
self.host = host # 主机地址
|
||
self.port = port # 主机端口
|
||
self.username = username # 主机用户名
|
||
self.password = password # 主机密码(可选)
|
||
self.recursive = recursive # 是否递归遍历子文件夹
|
||
self.ignore_patterns = ignore_patterns # 忽略的文件模式列表
|
||
|
||
|
||
class GitDataSourceConfig(BaseDataSourceConfig):
|
||
"""Git data source configuration"""
|
||
def __init__(
|
||
self,
|
||
name: str,
|
||
git_url: str,
|
||
branch: str = "main",
|
||
protocol: str = "https", # https 或 ssh
|
||
ssh_key: Optional[str] = None, # SSH私钥
|
||
https_token: Optional[str] = None, # HTTPS令牌
|
||
local_repo_path: Optional[str] = None, # 本地存储路径
|
||
poll_interval: int = 300, # 轮询间隔(秒)
|
||
support_lang: Optional[List[str]] = None, # 支持的编程语言
|
||
latest_commit_id: Optional[str] = None, # 最新commit ID
|
||
last_sync_time: Optional[str] = None # 最后同步时间
|
||
):
|
||
super().__init__(name, "git")
|
||
self.git_url = git_url # Git仓库地址
|
||
self.branch = branch # 分支名称
|
||
self.protocol = protocol # 协议类型
|
||
self.ssh_key = ssh_key # SSH私钥(加密存储)
|
||
self.https_token = https_token # HTTPS令牌(加密存储)
|
||
self.local_repo_path = local_repo_path # 本地存储路径
|
||
self.poll_interval = poll_interval # 轮询间隔
|
||
self.support_lang = support_lang # 支持的编程语言
|
||
self.latest_commit_id = latest_commit_id # 最新commit ID
|
||
self.last_sync_time = last_sync_time # 最后同步时间
|
||
|
||
|
||
class Settings(BaseSettings):
|
||
"""
|
||
Application settings
|
||
|
||
All configuration values can be overridden via .env file.
|
||
Default values are provided as fallback when .env is not present or values are missing.
|
||
See .env.example for all available configuration options.
|
||
"""
|
||
|
||
# API Settings
|
||
# These can be overridden in .env file
|
||
API_HOST: str = "0.0.0.0"
|
||
API_PORT: int = 8001 # Default changed to 8001 to avoid conflict with ChromaDB
|
||
API_TITLE: str = "RAG API"
|
||
API_VERSION: str = "1.0.0"
|
||
|
||
# File Upload Settings
|
||
MAX_UPLOAD_SIZE_MB: int = 5 # Maximum file upload size in MB (default: 5MB)
|
||
|
||
# ChromaDB Settings
|
||
# Use HttpClient mode if CHROMA_SERVER_HOST is set, otherwise use PersistentClient
|
||
# For Docker deployment, set CHROMA_SERVER_HOST=localhost in .env
|
||
CHROMA_SERVER_HOST: Optional[str] = None # e.g., "localhost" or "chromadb" (for Docker)
|
||
CHROMA_SERVER_PORT: int = 8000 # ChromaDB server port
|
||
CHROMA_DB_PATH: str = "./chroma_db" # Only used for PersistentClient mode
|
||
CHROMA_COLLECTION_NAME: str = "rag_collection"
|
||
|
||
# Ollama Settings
|
||
# Configure OLLAMA_BASE_URL in .env file based on your deployment
|
||
# - Local: http://localhost:11434
|
||
# - Remote: http://192.168.1.100:11434
|
||
OLLAMA_BASE_URL: str = "http://localhost:11434"
|
||
OLLAMA_MODEL: str = "qwen3:1.7b" # LLM model for text generation
|
||
OLLAMA_EMBEDDING_MODEL: str = "qwen3-embedding:0.6b" # Embedding model for vectorization
|
||
|
||
# RAG Settings
|
||
EMBEDDING_DIMENSION: int = 768
|
||
CHUNK_SIZE: int = 4000
|
||
CHUNK_OVERLAP: int = 200
|
||
TOP_K: int = 5 # Number of documents to retrieve
|
||
|
||
# Sync Settings
|
||
SYNC_INTERVAL: int = 300 # Sync interval in seconds
|
||
AUTO_SYNC: bool = True
|
||
|
||
# NLTK Settings
|
||
# Prefer a local checked-in NLTK data directory (useful for offline environments).
|
||
# If a local copy exists under the repository (example: ./https:/gitee.com/gislite/nltk_data/raw/gh-pages/)
|
||
# use that; otherwise fall back to the official raw GitHub mirror.
|
||
_PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
|
||
_NLTK_CANDIDATES = [
|
||
os.path.join(_PROJECT_ROOT, './nltk_data'),
|
||
]
|
||
_DEFAULT_NLTK = next((p for p in _NLTK_CANDIDATES if os.path.isdir(p)),
|
||
"https://raw.githubusercontent.com/nltk/nltk_data/gh-pages/")
|
||
|
||
NLTK_DATA: str = _DEFAULT_NLTK
|
||
|
||
# 文件下载接口地址,根据实际环境进行修改
|
||
FILE_DOWNLOAD_BASE_URL: str = "http://172.20.32.184:8000/api/file/open/downloadByIdentifier"
|
||
|
||
# 添加 SOFFICE 配置
|
||
SOFFICE_HOST: str = "127.0.0.1"
|
||
SOFFICE_PORT: int = 8003
|
||
|
||
# Git 相关配置
|
||
GIT_LOCAL_STORAGE_ROOT: str = "./git_repos" # Git仓库本地存储根目录
|
||
GIT_DEFAULT_BRANCH: str = "main" # 默认分支
|
||
GIT_POLL_INTERVAL: int = 300 # 默认轮询间隔(秒)
|
||
GIT_MAX_REPO_SIZE_MB: int = 500 # 最大仓库大小(MB)
|
||
|
||
# Pydantic v2 configuration
|
||
model_config = SettingsConfigDict(
|
||
env_file=".env",
|
||
env_file_encoding="utf-8",
|
||
case_sensitive=True,
|
||
extra="ignore" # Ignore extra fields in .env file that are not defined in Settings
|
||
)
|
||
|
||
def get_data_sources(self) -> List[BaseDataSourceConfig]:
|
||
"""
|
||
Get list of data source configurations
|
||
|
||
Returns:
|
||
List of BaseDataSourceConfig objects. Empty list if no configurations found.
|
||
"""
|
||
configs = []
|
||
|
||
try:
|
||
# Load configs from SQLite database only
|
||
import sqlite3
|
||
from pathlib import Path
|
||
|
||
DATA_DIR = Path(__file__).parent / "data"
|
||
DATA_DIR.mkdir(parents=True, exist_ok=True)
|
||
DB_PATH = DATA_DIR / "sessions.db"
|
||
|
||
try:
|
||
conn = sqlite3.connect(DB_PATH)
|
||
cursor = conn.cursor()
|
||
|
||
# 查询所有数据源配置
|
||
try:
|
||
cursor.execute('SELECT name, config FROM data_sources')
|
||
rows = cursor.fetchall()
|
||
|
||
for row in rows:
|
||
name, config_json = row
|
||
try:
|
||
ds_config = json.loads(config_json)
|
||
|
||
# Create data source objects based on type
|
||
source_type = ds_config.get('type', 'database')
|
||
|
||
if source_type == 'database':
|
||
# Create database data source
|
||
configs.append(DatabaseDataSourceConfig(
|
||
name=name, # 使用数据库表中的name列
|
||
# 通用数据库连接信息
|
||
db_type= ds_config.get('db_type', 'mysql'), # Default to 'mysql'
|
||
database=ds_config['database'], # Required field
|
||
table_name=ds_config.get('table_name', 'documents'),
|
||
host=ds_config.get('host', None),
|
||
port=ds_config.get('port', None),
|
||
user=ds_config.get('user', None),
|
||
password=ds_config.get('password', None),
|
||
|
||
id_column=ds_config.get('id_column', 'id'), # Default to 'id'
|
||
content_column=ds_config.get('content_column', 'content'), # Default to 'content'
|
||
file_column=ds_config.get('file_column', None),
|
||
title_column=ds_config.get('title_column', 'title'), # Default to 'title'
|
||
metadata_columns=ds_config.get('metadata_columns', None),
|
||
content_separator=ds_config.get('content_separator', '\n'),
|
||
updated_at_column=ds_config.get('updated_at_column', None),
|
||
# 文件源配置
|
||
file_source_type=ds_config.get('file_source_type', None),
|
||
file_system_base_path=ds_config.get('file_system_base_path', None),
|
||
# SCP配置
|
||
scp_host=ds_config.get('scp_host', None),
|
||
scp_port=ds_config.get('scp_port', 22),
|
||
scp_username=ds_config.get('scp_username', None),
|
||
scp_password=ds_config.get('scp_password', None)
|
||
))
|
||
elif source_type == 'folder':
|
||
# Create folder data source
|
||
configs.append(FolderDataSourceConfig(
|
||
name=name, # 使用数据库表中的name列
|
||
folder_path=ds_config.get('folder_path', '.'),
|
||
host=ds_config.get('host', 'localhost'),
|
||
port=ds_config.get('port', 22),
|
||
username=ds_config.get('username', None),
|
||
password=ds_config.get('password', None),
|
||
recursive=ds_config.get('recursive', True),
|
||
ignore_patterns=ds_config.get('ignore_patterns', None)
|
||
))
|
||
elif source_type == 'git':
|
||
# Create git data source
|
||
configs.append(GitDataSourceConfig(
|
||
name=name, # 使用数据库表中的name列
|
||
git_url=ds_config.get('git_url'),
|
||
branch=ds_config.get('branch', 'main'),
|
||
protocol=ds_config.get('protocol', 'https'),
|
||
ssh_key=ds_config.get('ssh_key'),
|
||
https_token=ds_config.get('https_token'),
|
||
local_repo_path=ds_config.get('local_repo_path'),
|
||
poll_interval=ds_config.get('poll_interval', 300),
|
||
support_lang=ds_config.get('support_lang'),
|
||
latest_commit_id=ds_config.get('latest_commit_id'),
|
||
last_sync_time=ds_config.get('last_sync_time')
|
||
))
|
||
else:
|
||
from loguru import logger
|
||
logger.warning(f"Unknown data source type: {source_type}, skipping")
|
||
|
||
from loguru import logger
|
||
logger.info(f"Loaded config from database: {name}")
|
||
except Exception as e:
|
||
from loguru import logger
|
||
logger.error(f"Error parsing config from database: {name}, error: {e}")
|
||
except sqlite3.OperationalError as e:
|
||
# 表不存在的情况,返回空列表
|
||
from loguru import logger
|
||
logger.warning(f"SQLite table error: {e}. Returning empty config list.")
|
||
except Exception as e:
|
||
# 其他数据库错误,返回空列表
|
||
from loguru import logger
|
||
logger.error(f"Error querying data sources from database: {e}")
|
||
|
||
# 关闭数据库连接
|
||
conn.close()
|
||
|
||
except Exception as e:
|
||
# 数据库连接失败,返回空列表
|
||
from loguru import logger
|
||
logger.error(f"Error connecting to SQLite database: {e}")
|
||
except Exception as e:
|
||
# 任何其他错误,返回空列表
|
||
from loguru import logger
|
||
logger.error(f"Error in get_data_sources: {e}")
|
||
|
||
# 返回配置列表,即使为空
|
||
return configs
|
||
|
||
def get_database_configs(self) -> List[DatabaseDataSourceConfig]:
|
||
"""
|
||
Get list of database configurations (backward compatibility)
|
||
|
||
Returns:
|
||
List of DatabaseDataSourceConfig objects
|
||
"""
|
||
all_sources = self.get_data_sources()
|
||
# Filter only database sources
|
||
return [source for source in all_sources if isinstance(source, DatabaseDataSourceConfig)]
|
||
|
||
|
||
settings = Settings()
|
||
|