diff --git a/.gitignore b/.gitignore index 24cc791..d94f0cb 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,4 @@ -!install/ +install/ # Python __pycache__/ *.py[cod] diff --git a/Dockerfile b/Dockerfile index 9a5e027..5ca09e6 100644 --- a/Dockerfile +++ b/Dockerfile @@ -26,6 +26,9 @@ RUN apt-get update && apt-get install -y --no-install-recommends \ gcc \ g++ \ curl \ + libssl-dev \ + libcrypto++-dev \ + libgmp-dev \ && rm -rf /var/lib/apt/lists/* # 配置 pip 镜像源(加速 Python 包安装) diff --git a/README.md b/README.md index 2fdfa40..aba4f29 100644 --- a/README.md +++ b/README.md @@ -180,19 +180,27 @@ ollama pull qwen3-embedding:8b # Embedding模型,用于向量化 注意: - 确保宿主机上已安装 Docker 和 Docker Compose +- 根据机器架构,下载对应架构的镜像文件后并解压,目前支持 X86 64位和ARM 64架构。 + +``` +X86 64位架构下载 V1.0.0.0.zip,包含镜像 chromadb.tar.gz、soffice.tar.gz、rag-api.tar.gz +ARM 64架构下载 V1.0.0.0-arm64.zip,包含镜像 chromadb-arm64.tar.gz、soffice-arm64.tar.gz、rag-api.tar.gz +``` + +解压后确保: - chromadb.tar.gz、soffice.tar.gz、rag-api.tar.gz 这三个文件必须存在于一个目录下 - .env 文件必须存在于该目录下(按照下面的步骤通过.env.example 生成并修改) - static/ 文件夹必须存在于该目录下 - docker-run.sh 脚本必须存在于该目录下 - nltk/ 文件夹必须存在于该目录下,用于存储 NLTK 数据 - + ```bash # 1. 配置 .env 文件(复制 .env.example 并修改) cp .env.example .env # 编辑 .env 文件,配置 MySQL、Ollama 等连接信息 -# 2. 构建并启动服务 +# 2. 运行docker-run.sh, 构建并启动服务 sudo bash docker-run.sh # 3. 查看服务状态 diff --git a/api/main.py b/api/main.py index abf46ed..0f64775 100644 --- a/api/main.py +++ b/api/main.py @@ -1343,7 +1343,10 @@ async def create_config(config: Dict[str, Any]): # 数据库配置需要:主机、端口、数据库名、表名 if not config.get("mysql_host") or not config.get("mysql_port") or not config.get("database") or not config.get("table_name"): raise HTTPException(status_code=400, detail="数据库配置必须包含主机、端口、数据库名和表名") + # 添加数据库类型到唯一标识符 + db_type = config.get("db_type", "mysql").lower() unique_id_parts.extend([ + db_type, config["mysql_host"].lower(), str(config["mysql_port"]), config["database"].lower(), @@ -1383,8 +1386,9 @@ async def create_config(config: Dict[str, Any]): # Check if it's the same type if existing_config_data.get('type') == config_type: if config_type == 'database': - # For database configs, same source means same host, port, and db name - if (existing_config_data.get('mysql_host') == config.get('mysql_host') and + # For database configs, same source means same db type, host, port, and db name + if (existing_config_data.get('db_type', 'mysql') == config.get('db_type', 'mysql') and + existing_config_data.get('mysql_host') == config.get('mysql_host') and existing_config_data.get('mysql_port') == config.get('mysql_port') and existing_config_data.get('database') == config.get('database')): raise HTTPException( @@ -1443,10 +1447,11 @@ async def create_config(config: Dict[str, Any]): metadata_columns=config.get("metadata_columns"), content_separator=config.get("content_separator"), updated_at_column=config.get("updated_at_column"), - mysql_host=config.get("mysql_host"), - mysql_port=config.get("mysql_port"), - mysql_user=config.get("mysql_user"), - mysql_password=config.get("mysql_password"), + host=config.get("mysql_host"), + port=config.get("mysql_port"), + user=config.get("mysql_user"), + password=config.get("mysql_password"), + db_type=config.get("db_type", "mysql"), file_source_type=config.get("file_source_type"), file_system_base_path=config.get("file_system_base_path"), scp_host=config.get("scp_host"), @@ -1587,12 +1592,13 @@ async def update_folder_config(config_id: str, config: Dict[str, Any]): # 根据不同类型的配置生成有意义的ID if config_type == "database": - # 数据库配置:使用数据库名和表名生成ID + # 数据库配置:使用数据库类型、数据库名和表名生成ID if not config.get("database"): raise HTTPException(status_code=400, detail="Database name is required for database configuration") if not config.get("table_name"): raise HTTPException(status_code=400, detail="Table name is required for database configuration") - new_config_id = f"{config_type}_{config['database'].lower()}_{config['table_name'].lower()}" + db_type = config.get("db_type", "mysql").lower() + new_config_id = f"{config_type}_{db_type}_{config['database'].lower()}_{config['table_name'].lower()}" elif config_type == "folder": # 文件夹配置:根据是否有host字段区分本地和远程 if not config.get("folder_path"): @@ -1648,10 +1654,11 @@ async def update_folder_config(config_id: str, config: Dict[str, Any]): metadata_columns=config.get("metadata_columns"), content_separator=config.get("content_separator"), updated_at_column=config.get("updated_at_column"), - mysql_host=config.get("mysql_host"), - mysql_port=config.get("mysql_port"), - mysql_user=config.get("mysql_user"), - mysql_password=config.get("mysql_password"), + host=config.get("mysql_host"), + port=config.get("mysql_port"), + user=config.get("mysql_user"), + password=config.get("mysql_password"), + db_type=config.get("db_type", "mysql"), file_source_type=config.get("file_source_type"), file_system_base_path=config.get("file_system_base_path"), scp_host=config.get("scp_host"), @@ -1751,6 +1758,7 @@ class DatabaseConnectionParams(BaseModel): port: int = Field(default=3306, description="Database port") username: str = Field(..., description="Database username") password: str = Field(..., description="Database password") + db_type: str = Field(default="mysql", description="Database type: mysql or dameng") class DatabaseParams(BaseModel): """Database parameters""" @@ -1759,6 +1767,7 @@ class DatabaseParams(BaseModel): username: str = Field(..., description="Database username") password: str = Field(..., description="Database password") database: str = Field(..., description="Database name") + db_type: str = Field(default="mysql", description="Database type: mysql or dameng") class TableParams(BaseModel): """Table parameters""" @@ -1768,6 +1777,7 @@ class TableParams(BaseModel): password: str = Field(..., description="Database password") database: str = Field(..., description="Database name") table_name: str = Field(..., description="Table name") + db_type: str = Field(default="mysql", description="Database type: mysql or dameng") @@ -1775,7 +1785,7 @@ class TableParams(BaseModel): @app.post("/database/databases") async def get_databases(params: DatabaseConnectionParams): """ - Get list of databases from MySQL server + Get list of databases from database server Args: params: Database connection parameters @@ -1784,34 +1794,53 @@ async def get_databases(params: DatabaseConnectionParams): List of available databases """ try: - # Direct database connection - connection = mysql.connector.connect( - host=params.host, - port=params.port, - user=params.username, - password=params.password - ) - - cursor = connection.cursor() - cursor.execute("SHOW DATABASES") - - databases = [] - for (database_name,) in cursor: - # Skip system databases - if database_name not in ['information_schema', 'mysql', 'performance_schema', 'sys']: - databases.append(database_name) - - cursor.close() - connection.close() + if params.db_type == "dameng": + # DaMeng database connection + import dmPython + connection = dmPython.connect( + user=params.username, + password=params.password, + server=params.host, + port=params.port + ) + + cursor = connection.cursor() + # Get schemas for DaMeng (similar to databases in MySQL) + cursor.execute("SELECT NAME AS SCHEMA_NAME FROM SYSOBJECTS WHERE TYPE$ = 'SCH';") + + databases = [] + for (schema_name,) in cursor: + databases.append(schema_name) + + cursor.close() + connection.close() + + else: + # MySQL database connection + connection = mysql.connector.connect( + host=params.host, + port=params.port, + user=params.username, + password=params.password + ) + + cursor = connection.cursor() + cursor.execute("SHOW DATABASES") + + databases = [] + for (database_name,) in cursor: + # Skip system databases + if database_name not in ['information_schema', 'mysql', 'performance_schema', 'sys']: + databases.append(database_name) + + cursor.close() + connection.close() return {"databases": databases} - except mysql.connector.Error as e: + except Exception as e: logger.error(f"Database connection error: {e}") raise HTTPException(status_code=500, detail=f"Database connection failed: {str(e)}") - except Exception as e: - logger.error(f"Error getting databases: {e}") - raise HTTPException(status_code=500, detail=f"Error getting databases: {str(e)}") @app.post("/database/tables") @@ -1826,33 +1855,52 @@ async def get_tables(params: DatabaseParams): List of tables in the specified database """ try: - # Direct database connection - connection = mysql.connector.connect( - host=params.host, - port=params.port, - user=params.username, - password=params.password, - database=params.database - ) - - cursor = connection.cursor() - cursor.execute("SHOW TABLES") - - tables = [] - for (table_name,) in cursor: - tables.append(table_name) - - cursor.close() - connection.close() + if params.db_type == "dameng": + # DaMeng database connection + import dmPython + connection = dmPython.connect( + user=params.username, + password=params.password, + server=params.host, + port=params.port + ) + + cursor = connection.cursor() + # Get tables for DaMeng + cursor.execute(f"SELECT TABLE_NAME FROM ALL_TABLES WHERE OWNER = '{params.database.upper()}'") + + tables = [] + for (table_name,) in cursor: + tables.append(table_name) + + cursor.close() + connection.close() + + else: + # MySQL database connection + connection = mysql.connector.connect( + host=params.host, + port=params.port, + user=params.username, + password=params.password, + database=params.database + ) + + cursor = connection.cursor() + cursor.execute("SHOW TABLES") + + tables = [] + for (table_name,) in cursor: + tables.append(table_name) + + cursor.close() + connection.close() return {"tables": tables} - except mysql.connector.Error as e: + except Exception as e: logger.error(f"Database connection error: {e}") raise HTTPException(status_code=500, detail=f"Database connection failed: {str(e)}") - except Exception as e: - logger.error(f"Error getting tables: {e}") - raise HTTPException(status_code=500, detail=f"Error getting tables: {str(e)}") @app.post("/database/table-structure") @@ -1867,40 +1915,66 @@ async def get_table_structure(params: TableParams): Table structure with column details """ try: - # Direct database connection - connection = mysql.connector.connect( - host=params.host, - port=params.port, - user=params.username, - password=params.password, - database=params.database - ) - - cursor = connection.cursor() - cursor.execute(f"DESCRIBE {params.table_name}") - - columns = [] - for (field, type, null, key, default, extra) in cursor: - columns.append({ - "name": field, - "type": type, - "null": null, - "key": key, - "default": default, - "extra": extra - }) - - cursor.close() - connection.close() + if params.db_type == "dameng": + # DaMeng database connection + import dmPython + connection = dmPython.connect( + user=params.username, + password=params.password, + server=params.host, + port=params.port + ) + + cursor = connection.cursor() + # Get table structure for DaMeng + cursor.execute(f"SELECT COLUMN_NAME, DATA_TYPE, NULLABLE, COLUMN_ID FROM ALL_TAB_COLUMNS WHERE OWNER = '{params.database.upper()}' AND TABLE_NAME = '{params.table_name.upper()}' ORDER BY COLUMN_ID") + + columns = [] + for (column_name, data_type, nullable, column_id) in cursor: + columns.append({ + "name": column_name, + "type": data_type, + "null": "YES" if nullable == 'Y' else "NO", + "key": "", # DaMeng doesn't return key info in this query + "default": None, + "extra": "" + }) + + cursor.close() + connection.close() + + else: + # MySQL database connection + connection = mysql.connector.connect( + host=params.host, + port=params.port, + user=params.username, + password=params.password, + database=params.database + ) + + cursor = connection.cursor() + cursor.execute(f"DESCRIBE {params.table_name}") + + columns = [] + for (field, type, null, key, default, extra) in cursor: + columns.append({ + "name": field, + "type": type, + "null": null, + "key": key, + "default": default, + "extra": extra + }) + + cursor.close() + connection.close() return {"columns": columns} - except mysql.connector.Error as e: + except Exception as e: logger.error(f"Database connection error: {e}") raise HTTPException(status_code=500, detail=f"Database connection failed: {str(e)}") - except Exception as e: - logger.error(f"Error getting table structure: {e}") - raise HTTPException(status_code=500, detail=f"Error getting table structure: {str(e)}") diff --git a/config.py b/config.py index 238b17f..7d97553 100644 --- a/config.py +++ b/config.py @@ -24,20 +24,24 @@ class DatabaseDataSourceConfig(BaseDataSourceConfig): 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 列之间的分隔符 + content_separator: str = "\n", # 多个 content 列拼接时的分隔符 updated_at_column: Optional[str] = None, # 用于增量同步的更新时间字段(可选) - # MySQL connection info (optional, will use .env defaults if not provided) - mysql_host: Optional[str] = None, - mysql_port: Optional[int] = None, - mysql_user: Optional[str] = None, - mysql_password: Optional[str] = None, + # 文件源配置 file_source_type: Optional[str] = None, # 可选值: "api", "filesystem", "scp" file_system_base_path: Optional[str] = None, # 文件系统基础路径 @@ -49,8 +53,16 @@ class DatabaseDataSourceConfig(BaseDataSourceConfig): scp_key_path: Optional[str] = None ): super().__init__(name, "database") - self.database = database # MySQL数据库名称 + + # 通用数据库连接信息 + 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: @@ -64,13 +76,7 @@ class DatabaseDataSourceConfig(BaseDataSourceConfig): self.metadata_columns = metadata_columns self.content_separator = content_separator # 多个列之间的分隔符 self.updated_at_column = updated_at_column # 更新时间字段(用于增量同步) - - # MySQL connection info (optional, will fallback to .env defaults) - self.mysql_host = mysql_host - self.mysql_port = mysql_port - self.mysql_user = mysql_user - self.mysql_password = mysql_password - + # 文件源配置 self.file_source_type = file_source_type self.file_system_base_path = file_system_base_path @@ -163,6 +169,7 @@ class Settings(BaseSettings): NLTK_DATA: str = _DEFAULT_NLTK + # 文件下载接口地址,根据实际环境进行修改 FILE_DOWNLOAD_BASE_URL: str = "http://172.20.32.184:8000/api/file/open/downloadByIdentifier" # 添加 SOFFICE 配置 @@ -216,8 +223,15 @@ class Settings(BaseSettings): # 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'), # Default to 'documents' + 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), @@ -225,11 +239,6 @@ class Settings(BaseSettings): 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), - # MySQL connection info - mysql_host=ds_config.get('mysql_host', None), - mysql_port=ds_config.get('mysql_port', None), - mysql_user=ds_config.get('mysql_user', None), - mysql_password=ds_config.get('mysql_password', None), # 文件源配置 file_source_type=ds_config.get('file_source_type', None), file_system_base_path=ds_config.get('file_system_base_path', None), diff --git a/requirements.txt b/requirements.txt index 822e74f..76545b4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -17,6 +17,7 @@ pymysql==1.1.0 sqlalchemy>=2.0.40 cryptography>=46.0.3 mysql-connector-python>=9.5.0 +dmpython==2.5.30 # Utilities python-dotenv==1.0.0 diff --git a/static/config/index.html b/static/config/index.html index 8688ea8..50a2af2 100644 --- a/static/config/index.html +++ b/static/config/index.html @@ -4,8 +4,8 @@ RAG 配置管理 - - + +
@@ -77,6 +77,13 @@
+