diff --git a/.gitignore b/.gitignore index f602ffa0..1fa62661 100644 --- a/.gitignore +++ b/.gitignore @@ -31,14 +31,9 @@ cache *.pdf *.yaml src/data -neo4j* */package-lock.json web/package-lock.json saves notebooks -local_neo4j/data -local_neo4j/logs -local_neo4j/import -local_neo4j/plugins -local_neo4j/conf -graphrag \ No newline at end of file +graphrag +docker/volumes \ No newline at end of file diff --git a/docker-compose.yml b/docker-compose.yml deleted file mode 100644 index 3c82054a..00000000 --- a/docker-compose.yml +++ /dev/null @@ -1,59 +0,0 @@ -version: '3.9' - -services: - # 后端服务 - backend: - build: . - container_name: backend - working_dir: /app - volumes: - - ./src:/app/src # 映射源代码 - - ./saves:/app/saves - - /hdd/zwj/models:/hdd/zwj/models - ports: - - "5000:5000" # 曝露后端端口 - depends_on: - - neo4j # 确保 Neo4j 启动 - networks: - - app-network - environment: - - NEO4J_URI=bolt://neo4j:7687 # 使用 neo4j 容器名称而非 localhost - - NEO4J_USERNAME=neo4j - - NEO4J_PASSWORD=0123456789 - - # 前端服务 - frontend: - build: - context: ./web # 设置构建上下文为 web 文件夹 - dockerfile: Dockerfile # 使用 web 文件夹内的 Dockerfile - container_name: frontend - ports: - - "80:80" - depends_on: - - backend - networks: - - app-network - - # Neo4j 图数据库服务 - neo4j: - image: neo4j:latest - container_name: neo4j - volumes: - - ./local_neo4j/conf:/var/lib/neo4j/conf - - ./local_neo4j/import:/var/lib/neo4j/import - - ./local_neo4j/plugins:/plugins - - ./local_neo4j/data:/data - - ./local_neo4j/logs:/var/lib/neo4j/logs - restart: always - ports: - - "7474:7474" - - "7687:7687" - environment: - - NEO4J_AUTH=neo4j/0123456789 - networks: - - app-network - -# 定义网络 -networks: - app-network: - driver: bridge diff --git a/Dockerfile b/docker/api.Dockerfile similarity index 78% rename from Dockerfile rename to docker/api.Dockerfile index 098c1a14..e852f621 100644 --- a/Dockerfile +++ b/docker/api.Dockerfile @@ -5,13 +5,11 @@ FROM pytorch/pytorch:2.4.1-cuda11.8-cudnn9-runtime WORKDIR /app # 复制 requirements.txt 文件(这一步如果文件没变,Docker 会使用缓存) -COPY requirements.txt /app/requirements.txt +COPY ../requirements.txt /app/requirements.txt # 安装依赖(Docker 会缓存这一步,除非 requirements.txt 发生变化) RUN pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple # 复制代码到容器中 -COPY ./src /app/src +COPY ../src /app/src -# 运行应用 -CMD ["python", "-m", "src.api"] diff --git a/docker/docker-compose.dev.yml b/docker/docker-compose.dev.yml new file mode 100644 index 00000000..5191d33f --- /dev/null +++ b/docker/docker-compose.dev.yml @@ -0,0 +1,128 @@ +services: + api: + build: + context: .. + dockerfile: docker/api.Dockerfile + container_name: api-dev + working_dir: /app + volumes: + - ../src:/app/src + - ../saves:/app/saves + - /hdd/zwj/models:/hdd/zwj/models + ports: + - "5000:5000" + depends_on: + - graph + - milvus + networks: + - app-network + environment: + - NEO4J_URI=bolt://graph:7687 + - NEO4J_USERNAME=neo4j + - NEO4J_PASSWORD=0123456789 + - MILVUS_URI=http://milvus:19530 + command: uvicorn src.main:app --host 0.0.0.0 --port 5000 --reload + + web: + build: + context: .. + dockerfile: docker/web.Dockerfile + target: development + container_name: web-dev + volumes: + - ../web:/app + - /app/node_modules + ports: + - "5173:5173" + depends_on: + - api + networks: + - app-network + environment: + - NODE_ENV=development + - VITE_API_URL=http://api:5000 # 添加这行 + command: npm run server + + graph: + image: neo4j:latest + container_name: graph-dev + ports: + - "7474:7474" + - "7687:7687" + volumes: + - ./volumes/neo4j/data:/data + - ./volumes/neo4j/logs:/var/lib/neo4j/logs + environment: + - NEO4J_AUTH=neo4j/0123456789 + - NEO4J_server_bolt_listen__address=0.0.0.0:7687 + - NEO4J_server_http_listen__address=0.0.0.0:7474 + networks: + - app-network + + etcd: + container_name: milvus-etcd-dev + image: quay.io/coreos/etcd:v3.5.5 + environment: + - ETCD_AUTO_COMPACTION_MODE=revision + - ETCD_AUTO_COMPACTION_RETENTION=1000 + - ETCD_QUOTA_BACKEND_BYTES=4294967296 + - ETCD_SNAPSHOT_COUNT=50000 + volumes: + - ./volumes/milvus/etcd:/etcd + command: etcd -advertise-client-urls=http://127.0.0.1:2379 -listen-client-urls http://0.0.0.0:2379 --data-dir /etcd + healthcheck: + test: ["CMD", "etcdctl", "endpoint", "health"] + interval: 30s + timeout: 20s + retries: 3 + networks: + - app-network + + minio: + container_name: milvus-minio-dev + image: minio/minio:RELEASE.2023-03-20T20-16-18Z + environment: + MINIO_ACCESS_KEY: minioadmin + MINIO_SECRET_KEY: minioadmin + volumes: + - ./volumes/milvus/minio:/minio_data + command: minio server /minio_data + healthcheck: + test: ["CMD", "curl", "-f", "http://localhost:9000/minio/health/live"] + interval: 30s + timeout: 20s + retries: 3 + networks: + - app-network + + milvus: + image: milvusdb/milvus:latest + container_name: milvus-standalone-dev + command: ["milvus", "run", "standalone"] + security_opt: + - seccomp:unconfined + environment: + ETCD_ENDPOINTS: etcd:2379 + MINIO_ADDRESS: minio:9000 + MILVUS_LOG_LEVEL: error # Add this line to reduce log output + volumes: + - ./volumes/milvus/milvus:/var/lib/milvus + - ./volumes/milvus/logs:/var/lib/milvus/logs + healthcheck: + test: ["CMD", "curl", "-f", "http://localhost:9091/healthz"] + interval: 30s + start_period: 90s + timeout: 20s + retries: 3 + ports: + - "19530:19530" + - "9091:9091" + depends_on: + - "etcd" + - "minio" + networks: + - app-network + +networks: + app-network: + driver: bridge \ No newline at end of file diff --git a/docker/docker-compose.yml b/docker/docker-compose.yml new file mode 100644 index 00000000..bd8b757e --- /dev/null +++ b/docker/docker-compose.yml @@ -0,0 +1,126 @@ +services: + api: + build: + context: .. + dockerfile: docker/api.Dockerfile + container_name: api-prod + working_dir: /app + volumes: + - ../src:/app/src + - ../saves:/app/saves + - /hdd/zwj/models:/hdd/zwj/models + ports: + - "8000:8000" + depends_on: + - graph + - milvus + networks: + - app-network + environment: + - NEO4J_URI=bolt://graph:7687 + - NEO4J_USERNAME=neo4j + - NEO4J_PASSWORD=0123456789 + - MILVUS_URI=http://milvus:19530 + command: uvicorn src.main:app --host 0.0.0.0 --port 8000 + + # 前端服务 + web: + build: + context: .. + dockerfile: docker/web.Dockerfile + target: production + container_name: web-prod + ports: + - "80:80" + depends_on: + - api + networks: + - app-network + environment: + - NODE_ENV=production + - VITE_API_URL=http://api:8000 # 添加这行 + + # Neo4j 图数据库服务 + graph: + image: neo4j:latest + container_name: graph-prod + ports: + - "7474:7474" + - "7687:7687" + volumes: + - ./volumes/neo4j/data:/data + - ./volumes/neo4j/logs:/var/lib/neo4j/logs + environment: + - NEO4J_AUTH=neo4j/0123456789 + networks: + - app-network + + etcd: + container_name: milvus-etcd-prod + image: quay.io/coreos/etcd:v3.5.5 + environment: + - ETCD_AUTO_COMPACTION_MODE=revision + - ETCD_AUTO_COMPACTION_RETENTION=1000 + - ETCD_QUOTA_BACKEND_BYTES=4294967296 + - ETCD_SNAPSHOT_COUNT=50000 + volumes: + - ./volumes/milvus/etcd:/etcd + command: etcd -advertise-client-urls=http://127.0.0.1:2379 -listen-client-urls http://0.0.0.0:2379 --data-dir /etcd + healthcheck: + test: ["CMD", "etcdctl", "endpoint", "health"] + interval: 30s + timeout: 20s + retries: 3 + networks: + - app-network + + minio: + container_name: milvus-minio-prod + image: minio/minio:RELEASE.2023-03-20T20-16-18Z + environment: + MINIO_ACCESS_KEY: minioadmin + MINIO_SECRET_KEY: minioadmin + volumes: + - ./volumes/milvus/minio:/minio_data + command: minio server /minio_data + healthcheck: + test: ["CMD", "curl", "-f", "http://localhost:9000/minio/health/live"] + interval: 30s + timeout: 20s + retries: 3 + networks: + - app-network + + + # Milvus 服务 + milvus: + image: milvusdb/milvus:latest + container_name: milvus-standalone-prod + command: ["milvus", "run", "standalone"] + security_opt: + - seccomp:unconfined + environment: + ETCD_ENDPOINTS: etcd:2379 + MINIO_ADDRESS: minio:9000 + volumes: + - ./volumes/milvus/milvus:/var/lib/milvus + healthcheck: + test: ["CMD", "curl", "-f", "http://localhost:9091/healthz"] + interval: 30s + start_period: 90s + timeout: 20s + retries: 3 + ports: + - "19530:19530" + - "9091:9091" + depends_on: + - "etcd" + - "minio" + networks: + - app-network + + +# 定义网络 +networks: + app-network: + driver: bridge diff --git a/docker/nginx/default.conf b/docker/nginx/default.conf new file mode 100644 index 00000000..39b7aa57 --- /dev/null +++ b/docker/nginx/default.conf @@ -0,0 +1,25 @@ +server { + listen 80; + server_name localhost; + + # 增加客户端请求体大小限制 + client_max_body_size 20M; + + location / { + root /usr/share/nginx/html; + try_files $uri /index.html; + } + + location /api/ { + proxy_pass http://api:8000/; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + + # 对于上传请求,增加超时时间 + proxy_read_timeout 600; + proxy_connect_timeout 600; + proxy_send_timeout 600; + } +} \ No newline at end of file diff --git a/docker/nginx/nginx.conf b/docker/nginx/nginx.conf new file mode 100644 index 00000000..f7c3d999 --- /dev/null +++ b/docker/nginx/nginx.conf @@ -0,0 +1,25 @@ +user nginx; +worker_processes auto; + +error_log /var/log/nginx/error.log notice; +pid /var/run/nginx.pid; + +events { + worker_connections 1024; +} + +http { + include /etc/nginx/mime.types; + default_type application/octet-stream; + + log_format main '$remote_addr - $remote_user [$time_local] "$request" ' + '$status $body_bytes_sent "$http_referer" ' + '"$http_user_agent" "$http_x_forwarded_for"'; + + access_log /var/log/nginx/access.log main; + + sendfile on; + keepalive_timeout 65; + + include /etc/nginx/conf.d/*.conf; +} \ No newline at end of file diff --git a/local_neo4j/test_neo4j.py b/docker/test/test_neo4j.py old mode 100644 new mode 100755 similarity index 100% rename from local_neo4j/test_neo4j.py rename to docker/test/test_neo4j.py diff --git a/docker/web.Dockerfile b/docker/web.Dockerfile new file mode 100644 index 00000000..17d551c4 --- /dev/null +++ b/docker/web.Dockerfile @@ -0,0 +1,35 @@ +# 开发阶段 +FROM node:latest AS development +WORKDIR /app + +# 复制 package.json 和 package-lock.json(如果存在) +COPY ./web/package*.json ./ + +# 安装依赖 +RUN npm install --registry https://registry.npmmirror.com + +# 复制源代码 +COPY ./web . + +# 暴露端口 +EXPOSE 5173 + +# 启动开发服务器的命令在 docker-compose 文件中定义 + +# 生产阶段 +FROM node:latest AS build-stage +WORKDIR /app + +COPY ./web/package*.json ./ +RUN npm install --registry https://registry.npmmirror.com + +COPY ./web . +RUN npm run build + +# 生产环境运行阶段 +FROM nginx:alpine AS production +COPY --from=build-stage /app/dist /usr/share/nginx/html +COPY ./docker/nginx/nginx.conf /etc/nginx/nginx.conf +COPY ./docker/nginx/default.conf /etc/nginx/conf.d/default.conf +EXPOSE 80 +CMD ["nginx", "-g", "daemon off;"] \ No newline at end of file diff --git a/local_neo4j/docker-compose.yml b/local_neo4j/docker-compose.yml deleted file mode 100644 index 943bd6f4..00000000 --- a/local_neo4j/docker-compose.yml +++ /dev/null @@ -1,16 +0,0 @@ -version: '3.9' -services: - - neo4j: - image: neo4j:latest - volumes: - - ./conf:/var/lib/neo4j/conf - - ./import:/var/lib/neo4j/import - - ./plugins:/plugins - - ./data:/data - - ./logs:/var/lib/neo4j/logs - ports: - - 7474:7474 - - 7687:7687 - environment: - - NEO4J_AUTH=neo4j/0123456789 diff --git a/requirements.txt b/requirements.txt index 2def7042..ed96727a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -21,4 +21,7 @@ neo4j>=5.22.0 sentencepiece==0.1.99 llama-index-readers-file opencv-python-headless -docx2txt \ No newline at end of file +docx2txt +uvicorn[standard] +fastapi +python-multipart \ No newline at end of file diff --git a/run.sh b/run.sh index a05c6e52..60967a79 100644 --- a/run.sh +++ b/run.sh @@ -4,7 +4,7 @@ stop_services() { echo "Stopping services..." pkill -f "npm run server" - pkill -f "python -m src.api" + pkill -f "uvicorn src.main:app --host 0.0.0.0 --port 5000 --reload" exit } @@ -12,7 +12,7 @@ stop_services() { trap stop_services SIGINT SIGTERM # Start the server -python -m src.api & +uvicorn src.main:app --host 0.0.0.0 --port 5000 --reload & # Start the frontend service cd web diff --git a/scripts/milvus/standalone_embed.sh b/scripts/milvus/standalone_embed.sh deleted file mode 100644 index 610f26f8..00000000 --- a/scripts/milvus/standalone_embed.sh +++ /dev/null @@ -1,143 +0,0 @@ -#!/usr/bin/env bash - -# Licensed to the LF AI & Data foundation under one -# or more contributor license agreements. See the NOTICE file -# distributed with this work for additional information -# regarding copyright ownership. The ASF licenses this file -# to you under the Apache License, Version 2.0 (the -# "License"); you may not use this file except in compliance -# with the License. You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -run_embed() { - cat << EOF > embedEtcd.yaml -listen-client-urls: http://0.0.0.0:2379 -advertise-client-urls: http://0.0.0.0:2379 -quota-backend-bytes: 4294967296 -auto-compaction-mode: revision -auto-compaction-retention: '1000' -EOF - - cat << EOF > user.yaml -# Extra config to override default milvus.yaml -EOF - - sudo docker run -d \ - --name milvus-standalone \ - --security-opt seccomp:unconfined \ - -e ETCD_USE_EMBED=true \ - -e ETCD_DATA_DIR=/var/lib/milvus/etcd \ - -e ETCD_CONFIG_PATH=/milvus/configs/embedEtcd.yaml \ - -e COMMON_STORAGETYPE=local \ - -v $(pwd)/volumes/milvus:/var/lib/milvus \ - -v $(pwd)/embedEtcd.yaml:/milvus/configs/embedEtcd.yaml \ - -v $(pwd)/user.yaml:/milvus/configs/user.yaml \ - -p 19530:19530 \ - -p 9091:9091 \ - -p 2379:2379 \ - --health-cmd="curl -f http://localhost:9091/healthz" \ - --health-interval=30s \ - --health-start-period=90s \ - --health-timeout=20s \ - --health-retries=3 \ - milvusdb/milvus:v2.4.5 \ - milvus run standalone 1> /dev/null -} - -wait_for_milvus_running() { - echo "Wait for Milvus Starting..." - while true - do - res=`sudo docker ps|grep milvus-standalone|grep healthy|wc -l` - if [ $res -eq 1 ] - then - echo "Start successfully." - echo "To change the default Milvus configuration, add your settings to the user.yaml file and then restart the service." - break - fi - sleep 1 - done -} - -start() { - res=`sudo docker ps|grep milvus-standalone|grep healthy|wc -l` - if [ $res -eq 1 ] - then - echo "Milvus is running." - exit 0 - fi - - res=`sudo docker ps -a|grep milvus-standalone|wc -l` - if [ $res -eq 1 ] - then - sudo docker start milvus-standalone 1> /dev/null - else - run_embed - fi - - if [ $? -ne 0 ] - then - echo "Start failed." - exit 1 - fi - - wait_for_milvus_running -} - -stop() { - sudo docker stop milvus-standalone 1> /dev/null - - if [ $? -ne 0 ] - then - echo "Stop failed." - exit 1 - fi - echo "Stop successfully." - -} - -delete() { - res=`sudo docker ps|grep milvus-standalone|wc -l` - if [ $res -eq 1 ] - then - echo "Please stop Milvus service before delete." - exit 1 - fi - sudo docker rm milvus-standalone 1> /dev/null - if [ $? -ne 0 ] - then - echo "Delete failed." - exit 1 - fi - sudo rm -rf $(pwd)/volumes - sudo rm -rf $(pwd)/embedEtcd.yaml - sudo rm -rf $(pwd)/user.yaml - echo "Delete successfully." -} - - -case $1 in - restart) - stop - start - ;; - start) - start - ;; - stop) - stop - ;; - delete) - delete - ;; - *) - echo "please use bash standalone_embed.sh restart|start|stop|delete" - ;; -esac \ No newline at end of file diff --git a/scripts/vllm/main.py b/scripts/vllm/main.py new file mode 100644 index 00000000..40efaeaa --- /dev/null +++ b/scripts/vllm/main.py @@ -0,0 +1,20 @@ +from vllm import LLM, SamplingParams + +llm = LLM(model="/hdd/zwj/models/meta-llama/Meta-Llama-3-8B-Instruct") + + +prompts = [ + "Hello, my name is", + "The president of the United States is", + "The capital of France is", + "The future of AI is", +] +sampling_params = SamplingParams(temperature=0.8, top_p=0.95) + +outputs = llm.generate(prompts, sampling_params) + +# Print the outputs. +for output in outputs: + prompt = output.prompt + generated_text = output.outputs[0].text + print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}") \ No newline at end of file diff --git a/vllm/run.sh b/scripts/vllm/run.sh similarity index 100% rename from vllm/run.sh rename to scripts/vllm/run.sh diff --git a/vllm/test_vllm.py b/scripts/vllm/test_vllm.py similarity index 100% rename from vllm/test_vllm.py rename to scripts/vllm/test_vllm.py diff --git a/src/.env.template b/src/.env.template deleted file mode 100644 index 3dfedca2..00000000 --- a/src/.env.template +++ /dev/null @@ -1,2 +0,0 @@ -ZHIPUAI_API_KEY= -CUDA_VISIBLE_DEVICES=0 \ No newline at end of file diff --git a/src/api.py b/src/api.py deleted file mode 100644 index 68c78193..00000000 --- a/src/api.py +++ /dev/null @@ -1,15 +0,0 @@ -from dotenv import load_dotenv -load_dotenv() - -import os -from src.views import create_app -from src.utils import setup_logger - -logger = setup_logger("Server") - -app = create_app() - -if __name__ == '__main__': - logger.info("Starting server") - app.secret_key = os.urandom(24) - app.run(host='0.0.0.0', port=5000, debug=False, threaded=True) diff --git a/src/cli.py b/src/cli.py deleted file mode 100644 index 539eb10c..00000000 --- a/src/cli.py +++ /dev/null @@ -1,41 +0,0 @@ -import os -from dotenv import load_dotenv -from src.core import HistoryManager -from src.core import Retriever -from src.config import Config -from src.models import select_model - -load_dotenv() - - -if __name__ == "__main__": - config = Config("config/base.yaml") - model = select_model(config) - retriever = Retriever(config) - - print(f"[{config.model_provider}:{config.get('model_name', 'default')}] Type 'exit' to quit") - - history_manager = HistoryManager() - while True: - query = input("\nUser: ") - if query == "exit": - break - - # 检索结果 - query, refs = retriever(query) - - messages = history_manager.add_user(query) - response = model.predict(messages, stream=config.stream) - - if config.stream: - content = "" - print(f"AI: ", end='', flush=True) - for chunk in response: - content += chunk.content - print(f"{chunk.content}", end='', flush=True) - print() - else: - content = response.content - print(f"AI: {content}") - - history_manager.add_ai(content) diff --git a/src/config/__init__.py b/src/config/__init__.py index 7f7fdc62..55824f26 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -80,13 +80,20 @@ class Config(SimpleConfig): def handle_self(self): self.model_names = MODEL_NAMES + model_provider_info = self.model_names.get(self.model_provider, {}) - if self.model_name not in self.model_names[self.model_provider]["models"]: - logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name") - self.model_name = self.model_names[self.model_provider]["default"] + if self.model_provider != "custom": + if self.model_name not in model_provider_info["models"]: + logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name") + self.model_name = model_provider_info["default"] - default_model_name = self.model_names[self.model_provider]["default"] - self.model_name = self.get("model_name") or default_model_name + default_model_name = model_provider_info["default"] + self.model_name = self.get("model_name") or default_model_name + else: + self.model_name = self.get("model_name") + if self.model_name not in [item["name"] for item in self.custom_models]: + logger.warning(f"Model name {self.model_name} not in custom models, using default model name") + self.model_name = self.custom_models[0]["name"] self.model_provider_status = {} for provider in self.model_names: @@ -200,15 +207,6 @@ MODEL_NAMES = { ] }, - "vllm": { - "name": "VLLM", - "default": "vllm", - "env": ["VLLM_API_KEY", "VLLM_API_BASE"], - "models": [ - "vllm", - ] - }, - # https://bailian.console.aliyun.com/?switchAgent=10226727&productCode=p_efm#/model-market "dashscope": { "name": "阿里百炼 (DashScope)", @@ -241,7 +239,7 @@ MODEL_NAMES = { "meta-llama/Meta-Llama-3.1-70B-Instruct", "meta-llama/Meta-Llama-3.1-405B-Instruct", ] - } + }, } diff --git a/src/core/database.py b/src/core/database.py index ac7ec9b4..c1715953 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -242,7 +242,7 @@ class DataBaseLite: self.dimension = dimension self.db_id = kwargs.get("db_id", hashstr(name)) self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}") - self.metadata = kwargs.get("metaname", {}) + self.metadata = kwargs.get("metadata", {}) self.files = kwargs.get("files", []) self.embed_model = kwargs.get("embed_model", None) diff --git a/src/core/graphbase.py b/src/core/graphbase.py index b4ed57b0..9ab9d66a 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -16,70 +16,6 @@ logger = setup_logger("server-graphbase") warnings.filterwarnings("ignore", category=UserWarning) -""" -from neo4j import GraphDatabase -import random - -class KnowledgeGraph: - def __init__(self, uri, user, password): - self._driver = GraphDatabase.driver(uri, auth=(user, password)) - - def close(self): - self._driver.close() - - def use_database(self, kgdb_name): - with self._driver.session() as session: - session.run(f"USE {kgdb_name}") - - def get_sample_nodes(self, kgdb_name='neo4j', num=50): - self.use_database(kgdb_name) - selected_nodes = set() - nodes_to_expand = set() - result_nodes = [] - - with self._driver.session() as session: - while len(selected_nodes) < num: - # 如果需要扩展的节点为空,随机选择一个新节点 - if not nodes_to_expand: - result = session.run("MATCH (n) RETURN n, rand() as r ORDER BY r LIMIT 1") - for record in result: - nodes_to_expand.add(record['n'].id) - result_nodes.append({'n': record['n'], 'r': None, 'm': None}) - - # 从需要扩展的节点中随机选择一个节点 - current_node_id = random.choice(list(nodes_to_expand)) - nodes_to_expand.remove(current_node_id) - - # 获取当前节点的邻居节点,最多5个 - result = session.run( - f"MATCH (n)-[r]-(m) WHERE id(n) = {current_node_id} RETURN n, r, m LIMIT 5" - ) - - for record in result: - neighbor_node_id = record['m'].id - if neighbor_node_id not in selected_nodes: - selected_nodes.add(neighbor_node_id) - nodes_to_expand.add(neighbor_node_id) - result_nodes.append({'n': record['n'], 'r': record['r'], 'm': record['m']}) - - # 如果已经达到最大值,停止扩展 - if len(selected_nodes) >= num: - break - - return result_nodes[:num] - -# 示例用法 -uri = "bolt://localhost:7687" -user = "neo4j" -password = "password" -kg = KnowledgeGraph(uri, user, password) -sample_nodes = kg.get_sample_nodes(num=50) -for node in sample_nodes: - print(f"Node: {node['n'].id}, Relationship: {node['r']}, Neighbor: {node['m'].id}") -kg.close() - -""" - UIE_MODEL = None class GraphDatabase: @@ -97,8 +33,14 @@ class GraphDatabase: username = os.environ.get("NEO4J_USERNAME", "neo4j") password = os.environ.get("NEO4J_PASSWORD", "0123456789") logger.info(f"Connecting to Neo4j at {uri}/{self.kgdb_name}") - self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password)) - self.status = "open" + try: + self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password)) + self.status = "open" + logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_database_info()}") + except Exception as e: + logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}") + self.config.enable_knowledge_graph = False + def close(self): """关闭数据库连接""" @@ -132,14 +74,19 @@ class GraphDatabase: """获取指定数据库的信息""" self.use_database(db_name) def query(tx): - entity_count = tx.run("MATCH (n:Entity) RETURN count(n) AS count").single()["count"] + entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"] relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"] triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"] + + # 获取所有标签 + labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"] + return { "database_name": db_name, "entity_count": entity_count, "relationship_count": relationship_count, "triples_count": triples_count, + "labels": labels, "status": self.status } diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 873f925a..65c24c64 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -1,7 +1,7 @@ import os from src.models.embedding import EmbeddingModel -from pymilvus import MilvusClient +from pymilvus import MilvusClient, MilvusException from src.utils import setup_logger, hashstr logger = setup_logger("KnowledgeBase") @@ -9,17 +9,29 @@ logger = setup_logger("KnowledgeBase") class KnowledgeBase: def __init__(self, config=None, embed_model=None) -> None: - self.config = config - self._init_config(config) - + self.config = config or {} assert embed_model, "embed_model=None" self.embed_model = embed_model - self.client = MilvusClient(self.milvus_path) + self.client = None + if not self.connect_to_milvus(): + raise ConnectionError("Failed to connect to Milvus") - def _init_config(self, config): - self.milvus_path = os.path.join(config.save_dir, "data/vector_base/milvus.db") - os.makedirs(os.path.dirname(self.milvus_path), exist_ok=True) + def connect_to_milvus(self): + """ + 连接到 Milvus 服务。 + 使用配置中的 URI,如果没有配置,则使用默认值。 + """ + try: + uri = os.getenv('MILVUS_URI', self.config.get('milvus_uri', "http://milvus:19530")) + self.client = MilvusClient(uri=uri) + # 可以添加一个简单的测试来确保连接成功 + self.client.list_collections() + logger.info(f"Successfully connected to Milvus at {uri}") + return True + except MilvusException as e: + logger.error(f"Failed to connect to Milvus: {e}") + return False def get_collection_names(self): return self.client.list_collections() diff --git a/src/models/__init__.py b/src/models/__init__.py index f941fbb5..46e3999b 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -20,10 +20,6 @@ def select_model(config): from src.models.chat_model import Qianfan return Qianfan(model_name) - elif model_provider == "vllm": - from src.models.chat_model import VLLM - return VLLM(model_name) - elif model_provider == "dashscope": from src.models.chat_model import DashScope return DashScope(model_name) @@ -36,6 +32,14 @@ def select_model(config): from src.models.chat_model import SiliconFlow return SiliconFlow(model_name) + elif model_provider == "custom": + model_info = next((x for x in config.custom_models if x["name"] == model_name), None) + if model_info is None: + raise ValueError(f"Model {model_name} not found in custom models") + + from src.models.chat_model import CustomModel + return CustomModel(model_info) + elif model_provider is None: raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`") else: diff --git a/src/models/chat_model.py b/src/models/chat_model.py index b387d8bb..8f1924e7 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -62,13 +62,6 @@ class Zhipu(OpenAIBase): base_url = "https://open.bigmodel.cn/api/paas/v4/" super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) -class VLLM(OpenAIBase): - def __init__(self, model_name=None): - model_name = model_name or "vllm" - api_key = os.getenv("VLLM_API_KEY") - base_url = os.getenv("VLLM_API_BASE") - super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) - class SiliconFlow(OpenAIBase): def __init__(self, model_name=None): model_name = model_name or "meta-llama/Meta-Llama-3.1-8B-Instruct" @@ -76,6 +69,13 @@ class SiliconFlow(OpenAIBase): base_url = "https://api.siliconflow.cn/v1" super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) +class CustomModel(OpenAIBase): + def __init__(self, model_info): + model_name = model_info["name"] + api_key = model_info["api_key"] + base_url = model_info["api_base"] + super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) + class GeneralResponse: def __init__(self, content): diff --git a/src/routers/data_router.py b/src/routers/data_router.py index 98776c7b..ed9794ab 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -23,7 +23,7 @@ async def create_database( database_name: str = Body(...), description: str = Body(...), db_type: str = Body(...), - dimension: int = Body(None) + dimension: Optional[int] = Body(None) ): logger.debug(f"Create database {database_name}") database_info = startup.dbm.create_database( @@ -52,7 +52,7 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod msg = startup.dbm.add_files(db_id, files) return msg -@data.get("/database-info") +@data.get("/info") async def get_database_info(db_id: str): logger.debug(f"Get database {db_id} info") database = startup.dbm.get_database_info(db_id) @@ -69,7 +69,13 @@ async def delete_document(db_id: str = Body(...), file_id: str = Body(...)): @data.get("/document") async def get_document_info(db_id: str, file_id: str): logger.debug(f"GET document {file_id} info in {db_id}") - info = startup.dbm.get_file_info(db_id, file_id) + + try: + info = startup.dbm.get_file_info(db_id, file_id) + except Exception as e: + logger.error(f"Failed to get file info, {e}, {db_id=}, {file_id=}") + info = {"message": "Failed to get file info", "status": "failed"}, 500 + return info @data.post("/upload") diff --git a/src/views/__init__.py b/src/views/__init__.py deleted file mode 100644 index 5fa4b936..00000000 --- a/src/views/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -from flask import Flask -from flask_cors import CORS -from src.views.common_view import common -from src.views.database_view import db -from src.views.tools_view import tools - - -def create_app(): - app = Flask(__name__) - CORS(app, resources=r'/*') - - app.register_blueprint(common) - app.register_blueprint(db) - app.register_blueprint(tools) - - return app diff --git a/src/views/common_view.py b/src/views/common_view.py deleted file mode 100644 index 34a1b42b..00000000 --- a/src/views/common_view.py +++ /dev/null @@ -1,101 +0,0 @@ -import itertools -import json -from flask import Blueprint, jsonify, request, Response - -from src.core import HistoryManager -from src.utils.logging_config import setup_logger -from src.core.startup import startup -from collections import deque - -common = Blueprint('common', __name__) -logger = setup_logger("server-common") - -@common.route('/', methods=["GET"]) -def route_index(): - return jsonify({"message": "You Got It!"}) - -@common.errorhandler(404) -def page_not_found(e): - return jsonify({"message": "DEBUG: " + str(e)}), 404 - - -@common.errorhandler(403) -def page_not_found(e): - return jsonify({"message": str(e)}), 403 -@common.route('/', methods=['GET']) -def chat_get(): - return "Chat Get!" - -@common.route('/chat', methods=['POST']) -def chat(): - request_data = json.loads(request.data) - query = request_data['query'] - meta = request_data.get('meta') - history_manager = HistoryManager(request_data['history']) - - new_query, refs = startup.retriever(query, history_manager.messages, meta) - - messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round')) - history_manager.add_user(query) - logger.debug(f"Web history: {history_manager.messages}") - - def generate_response(): - content = "" - for delta in startup.model.predict(messages, stream=True): - if not delta.content: - continue - - if hasattr(delta, 'is_full') and delta.is_full: - content = delta.content - else: - content += delta.content - - response_chunk = json.dumps({ - "history": history_manager.update_ai(content), - "response": content, - "refs": refs # TODO: 优化 refs,不需要每次都返回 - }, ensure_ascii=False).encode('utf8') + b'\n' - yield response_chunk - - return Response(generate_response(), content_type='application/json', status=200) - -@common.route('/call', methods=['POST']) -def call(): - request_data = json.loads(request.data) - query = request_data['query'] - response = startup.model.predict(query) - logger.debug({"query": query, "response": response.content}) - - return jsonify({ - "response": response.content, - }) - -@common.route('/config', methods=['get']) -def get_config(): - return jsonify(startup.config) - -@common.route('/config', methods=['post']) -def update_config(): - request_data = json.loads(request.data) - startup.config.update(request_data) - startup.config.save() - return jsonify(startup.config) - -@common.route('/restart', methods=['POST']) -def restart(): - startup.restart() - return jsonify({"message": "Restarted!"}) - - -@common.route('/log', methods=['GET']) -def get_log(): - from src.utils.logging_config import LOG_FILE - from collections import deque - - with open(LOG_FILE, 'r') as f: - # 使用deque保存最后1000行 - last_lines = deque(f, maxlen=1000) - - log = ''.join(last_lines) - - return jsonify({"log": log}) \ No newline at end of file diff --git a/src/views/database_view.py b/src/views/database_view.py deleted file mode 100644 index a904c294..00000000 --- a/src/views/database_view.py +++ /dev/null @@ -1,163 +0,0 @@ -import os -import json -import threading -from functools import wraps -from flask import Blueprint, jsonify, request, Response - -from src.utils import setup_logger, hashstr -from src.core.startup import startup - -db = Blueprint('database', __name__, url_prefix="/database") - -logger = setup_logger("server-database") - -progress = {} # 只针对单个用户的进度 - -def handle_exceptions(f): - @wraps(f) - def decorated_function(*args, **kwargs): - try: - logger.debug(f"Entering {f.__name__}") - result = f(*args, **kwargs) - logger.debug(f"Exiting {f.__name__}") - return result - except Exception as e: - logger.error(f"Error in {f.__name__}: {str(e)}") - return jsonify({"message": str(e), "error": "处理请求时发生错误"}), 500 - return decorated_function - -@db.route('/', methods=['GET']) -def get_databases(): - try: - database = startup.dbm.get_databases() - except Exception as e: - return jsonify({"message": "获取数据库列表失败", "databases": []}) - - return jsonify(database) - -@db.route('/', methods=['POST']) -def create_database(): - data = json.loads(request.data) - database_name = data.get('database_name') - description = data.get('description') - db_type = data.get('db_type') - dimension = data.get('dimension') - logger.debug(f"Create database {database_name}") - database = startup.dbm.create_database(database_name, description, db_type, dimension=dimension) - return jsonify(database) - -@db.route('/', methods=['DELETE']) -def delete_database(): - data = json.loads(request.data) - db_id = data.get('db_id') - logger.debug(f"Delete database {db_id}") - startup.dbm.delete_database(db_id) - return jsonify({"message": "删除成功"}) - -@db.route('/query-test', methods=['POST']) -def query_test(): - data = json.loads(request.data) - query = data.get('query') - meta = data.get('meta') - logger.debug(f"Query test in {meta}: {query}") - - result = startup.retriever.query_knowledgebase(query, history=None, refs={"meta": meta}) - return jsonify(result) - -@db.route('/add_by_file', methods=['POST']) -def create_document_by_file(): - data = json.loads(request.data) - db_id = data.get('db_id') - files = data.get('files') - logger.debug(f"Add document in {db_id} by file: {files}") - msg = startup.dbm.add_files(db_id, files) - return jsonify(msg) - - -@db.route('/info', methods=['GET']) -def get_database_info(): - db_id = request.args.get('db_id') - if not db_id: - return jsonify({"message": "db_id is required"}), 400 - - logger.debug(f"Get database {db_id} info") - database = startup.dbm.get_database_info(db_id) - - if database is None: - return jsonify({"message": "database not found"}), 404 - - return jsonify(database) - - -@db.route('/document', methods=['DELETE']) -def delete_document(): - data = json.loads(request.data) - db_id = data.get('db_id') - file_id = data.get('file_id') - logger.debug(f"DELETE document {file_id} info in {db_id}") - startup.dbm.delete_file(db_id, file_id) - return jsonify({"message": "删除成功"}) - -@db.route('/document', methods=['GET']) -def get_document_info(): - db_id = request.args.get('db_id') - file_id = request.args.get('file_id') - logger.debug(f"GET document {file_id} info in {db_id}") - info = startup.dbm.get_file_info(db_id, file_id) - return jsonify(info) - -@db.route('/upload', methods=['POST']) -def upload_file(): - if 'file' not in request.files: - return jsonify({'message': 'No file part in the request'}), 400 - file = request.files['file'] - if file.filename == '': - return jsonify({'message': 'No selected file'}), 400 - # elif file.filename.split('.')[-1] not in ['pdf', 'txt', 'md']: - # return jsonify({'message': 'Unsupported file type'}), 400 - if file: - upload_dir = os.path.join(startup.config.save_dir, "data/uploads") - os.makedirs(upload_dir, exist_ok=True) - filename = f"{hashstr(file.filename, 4, with_salt=True)}_{file.filename}".lower() - file_path = os.path.join(upload_dir, filename) - file.save(file_path) - return jsonify({'message': 'File successfully uploaded', 'file_path': file_path}), 200 - -@db.route('/graph', methods=['GET']) -def get_graph_info(): - graph_info = startup.dbm.get_graph() - return jsonify(graph_info) - -@db.route('/graph/node', methods=['GET']) -@handle_exceptions -def get_graph_node(): - assert request.args.get("entity_name"), "entity_name is required" - logger.debug(f"Get graph node {request.args.get('entity_name')} with {request.args}") - result = startup.dbm.graph_base.query_node(**request.args) - return jsonify({'result': startup.retriever.format_query_results(result), 'message': 'success'}), 200 - -@db.route('/graph/nodes', methods=['GET']) -@handle_exceptions -def get_graph_nodes(): - kgdb_name = request.args.get('kgdb_name') - num = request.args.get('num') - assert kgdb_name, "kgdb_name is required" - assert startup.config.enable_knowledge_graph, "Knowledge graph is not enabled" - - logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes") - result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num) - return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200 - -@db.route('/graph/add', methods=['POST']) -@handle_exceptions -def add_graph_entity(): - data = json.loads(request.data) - kgdb_name = data.get('kgdb_name') - file_path = data.get('file_path') - assert file_path.endswith('.jsonl'), "file_path must be a jsonl file" - assert startup.config.enable_knowledge_graph, "Knowledge graph is not enabled" - - startup.dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name) - - return jsonify({'message': 'Entity successfully added'}), 200 - diff --git a/src/views/tools_view.py b/src/views/tools_view.py deleted file mode 100644 index 14384838..00000000 --- a/src/views/tools_view.py +++ /dev/null @@ -1,48 +0,0 @@ -import os -import json -import threading -from flask import Blueprint, jsonify, request, Response - -from src.utils import setup_logger, hashstr -from src.core.startup import startup - -tools = Blueprint('tools', __name__, url_prefix="/tools") - -logger = setup_logger("server-tools") - - -@tools.route("/", methods=["GET"]) -def route_index(): - tools = [ - { - "name": "text_chunking", - "title": "文本分块", - "description": "将文本分块以更好地理解。可以输入文本或者上传文件。", - "url": "/tools/text_chunking", - "method": "POST", - }, - { - "name": "pdf2txt", - "title": "PDF转文本", - "description": "将PDF文件转换为文本文件。", - "url": "/tools/pdf2txt", - "method": "POST", - } - ] - - return jsonify(tools) - - -@tools.route("/text_chunking", methods=["POST"]) -def text_chunking(): - from src.core.indexing import chunk - text = request.json.get("text") - nodes = chunk(text, params=request.json) - return jsonify({"nodes": [node.to_dict() for node in nodes]}) - -@tools.route("/pdf2txt", methods=["POST"]) -def handle_pdf2txt(): - from src.plugins import pdf2txt - file = request.json.get("file") - text = pdf2txt(file, return_text=True) - return jsonify({"text": text}) diff --git a/web/Dockerfile b/web/Dockerfile deleted file mode 100644 index edfe9245..00000000 --- a/web/Dockerfile +++ /dev/null @@ -1,23 +0,0 @@ -# 使用 Node.js 作为基础镜像,用于构建前端 -FROM node:latest AS build-stage -WORKDIR /app - -# 将 package.json 和 package-lock.json 复制到工作目录 -COPY ./package*.json ./ - -# 安装依赖 -RUN npm install - -# 复制前端源代码并运行构建 -COPY . . -RUN npm run build - -# 使用 Nginx 作为生产镜像 -FROM nginx:alpine AS production-stage -COPY --from=build-stage /app/dist /usr/share/nginx/html - -# 复制 Nginx 配置文件 -COPY ../nginx.conf /etc/nginx/nginx.conf - -EXPOSE 80 -CMD ["nginx", "-g", "daemon off;"] diff --git a/web/nginx.conf b/web/nginx.conf deleted file mode 100644 index 270924f4..00000000 --- a/web/nginx.conf +++ /dev/null @@ -1,34 +0,0 @@ -worker_processes 1; - -events { - worker_connections 1024; -} - -http { - include /etc/nginx/mime.types; - default_type application/octet-stream; - - sendfile on; - keepalive_timeout 65; - - # This is where the server block should be placed - server { - listen 80; - - server_name localhost; - - location / { - root /usr/share/nginx/html; - try_files $uri /index.html; - } - - # Proxy to backend API service - location /api/ { - proxy_pass http://backend:5000/; - proxy_set_header Host $host; - proxy_set_header X-Real-IP $remote_addr; - proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; - proxy_set_header X-Forwarded-Proto $scheme; - } - } -} diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 0129ed25..e369f89d 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -1,4 +1,3 @@ -