RAG/config.py

383 lines
18 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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: Optional[str] = None,
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, # 最后同步时间
git_repositories: Optional[List[Dict[str, str]]] = None, # Git服务器模式下的仓库列表每个仓库包含repository和branch字段
git_server_root_path: Optional[str] = None, # Git服务器模式下的仓库根目录地址
):
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 # 最后同步时间
self.git_repositories = git_repositories # Git服务器模式下的仓库列表
self.git_server_root_path = git_server_root_path # Git服务器模式下的仓库根目录地址
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
# RAG Query Classification Settings
CODE_RELATED_THRESHOLD_LOW: float = 0.3 # 低于此值认为是非代码问题
CODE_RELATED_THRESHOLD_HIGH: float = 0.7 # 高于此值认为是代码相关问题
FILTER_METADATA_FIELDS: str = "class_name,func_name,file_path" # 从query中提取的metadata字段
# 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
git_config = GitDataSourceConfig(
name=name, # 使用数据库表中的name列
git_url=ds_config.get('git_url'),
branch=ds_config.get('branch', ds_config.get('git_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')
)
# 添加Git服务器模式的字段
git_mode = ds_config.get('git_mode', 'single')
if git_mode == 'server':
git_config.git_mode = git_mode
git_config.git_server_host = ds_config.get('git_server_host')
git_config.git_server_port = ds_config.get('git_server_port', 22)
git_config.git_server_username = ds_config.get('git_server_username')
git_config.git_server_password = ds_config.get('git_server_password')
git_config.git_repository = ds_config.get('git_repository')
git_config.git_branch = ds_config.get('git_branch', 'main')
git_config.git_repositories = ds_config.get('git_repositories', [])
configs.append(git_config)
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()