This commit is contained in:
Wenjie Zhang 2024-10-08 22:16:17 +08:00
parent 02686db89f
commit 39ff8a07a9
40 changed files with 694 additions and 849 deletions

9
.gitignore vendored
View File

@ -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
graphrag
docker/volumes

View File

@ -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

View File

@ -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"]

View File

@ -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

126
docker/docker-compose.yml Normal file
View File

@ -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

25
docker/nginx/default.conf Normal file
View File

@ -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;
}
}

25
docker/nginx/nginx.conf Normal file
View File

@ -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;
}

View File

35
docker/web.Dockerfile Normal file
View File

@ -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;"]

View File

@ -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

View File

@ -21,4 +21,7 @@ neo4j>=5.22.0
sentencepiece==0.1.99
llama-index-readers-file
opencv-python-headless
docx2txt
docx2txt
uvicorn[standard]
fastapi
python-multipart

4
run.sh
View File

@ -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

View File

@ -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

20
scripts/vllm/main.py Normal file
View File

@ -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}")

View File

@ -1,2 +0,0 @@
ZHIPUAI_API_KEY=
CUDA_VISIBLE_DEVICES=0

View File

@ -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)

View File

@ -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)

View File

@ -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",
]
}
},
}

View File

@ -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)

View File

@ -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
}

View File

@ -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()

View File

@ -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:

View File

@ -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):

View File

@ -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")

View File

@ -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

View File

@ -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})

View File

@ -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

View File

@ -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})

View File

@ -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;"]

View File

@ -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;
}
}
}

View File

@ -1,4 +1,3 @@
<!-- ChatComponent.vue -->
<template>
<div class="chat" ref="chatContainer">
<div class="header">
@ -128,7 +127,7 @@
placeholder="输入问题……"
:auto-size="{ minRows: 1, maxRows: 10 }"
/>
<a-button size="large" @click="sendMessage" :disabled="(!conv.inputText && !isStreaming)">
<a-button size="large" @click="sendMessage" :disabled="(!conv.inputText && !isStreaming)" type="link">
<template #icon> <SendOutlined v-if="!isStreaming" /> <LoadingOutlined v-else/> </template>
</a-button>
</div>
@ -373,7 +372,7 @@ const loadDatabases = () => {
// fetch
const fetchChatResponse = (user_input, cur_res_id) => {
fetch('/api/chat', {
fetch('/api/chat/', {
method: 'POST',
body: JSON.stringify({
query: user_input,
@ -408,7 +407,7 @@ const fetchChatResponse = (user_input, cur_res_id) => {
updateMessage(data.response, cur_res_id, data.refs, "loading");
conv.value.history = data.history;
} catch (e) {
console.error('JSON 解析错误:', e);
// console.error('JSON :', e);
}
});
buffer = ''; //

View File

@ -51,7 +51,7 @@ const getRemoteDatabase = () => {
onMounted(() => {
getRemoteDatabase()
getRemoteConfig()
configStore.refreshConfig()
})
// 使 vue3 setup composition API

View File

@ -30,5 +30,13 @@ export const useConfigStore = defineStore('config', () => {
})
}
return { config, setConfig, setConfigValue }
function refreshConfig() {
fetch('/api/config')
.then(response => response.json())
.then(data => {
setConfig(data)
})
}
return { config, setConfig, setConfigValue, refreshConfig }
})

View File

@ -1,7 +1,7 @@
<template>
<div>
<HeaderComponent
:title="database.name"
:title="database.name || '数据库信息'"
>
<template #description>
<div class="database-info">
@ -87,7 +87,7 @@
width="50%"
v-model:open="state.drawer"
class="custom-class"
:title="selectedFile?.filename"
:title="selectedFile?.filename || '文件详情'"
placement="right"
@after-open-change="afterOpenChange"
>
@ -302,6 +302,9 @@ const onQuery = () => {
meta.db_name = database.value.metaname
fetch('/api/data/query-test', {
method: "POST",
headers: {
"Content-Type": "application/json" // Content-Type
},
body: JSON.stringify({
query: queryText.value.trim(),
meta: meta
@ -361,6 +364,9 @@ const deleteDatabse = () => {
state.lock = true
fetch('/api/data/', {
method: "DELETE",
headers: {
"Content-Type": "application/json" // Content-Type
},
body: JSON.stringify({
db_id: databaseId.value
}),
@ -451,6 +457,9 @@ const deleteFile = (fileId) => {
state.lock = true
fetch('/api/data/document', {
method: "DELETE",
headers: {
"Content-Type": "application/json" // Content-Type
},
body: JSON.stringify({
db_id: databaseId.value,
file_id: fileId
@ -478,8 +487,11 @@ const addDocumentByFile = () => {
state.refreshInterval = setInterval(() => {
getDatabaseInfo();
}, 1000);
fetch('/api/data/add_by_file', {
fetch('/api/data/add-by-file', {
method: "POST",
headers: {
"Content-Type": "application/json" // Content-Type
},
body: JSON.stringify({
db_id: databaseId.value,
files: files

View File

@ -132,11 +132,14 @@ const createDatabase = () => {
}
fetch('/api/data/', {
method: "POST",
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
database_name: newDatabase.name,
description: newDatabase.description,
db_type: "knowledge",
dimension: newDatabase.dimension,
dimension: newDatabase.dimension ? parseInt(newDatabase.dimension) : null,
})
})
.then(response => response.json())
@ -161,24 +164,6 @@ const navigateToGraph = () => {
router.push({ path: `/database/graph` });
};
// const loadGraph = () => {
// graphloading.value = true
// fetch('/api/data/graph', {
// method: "GET",
// })
// .then(response => response.json())
// .then(data => {
// console.log(data)
// graph.value = data.graph
// graphloading.value = false
// })
// .catch(error => {
// console.error(error)
// message.error(error.message)
// graphloading.value = false
// })
// }
watch(() => route.path, (newPath, oldPath) => {
if (newPath === '/database') {
loadDatabases();

View File

@ -149,6 +149,9 @@ const addDocumentByFile = () => {
const files = fileList.value.filter(file => file.status === 'done').map(file => file.response.file_path)
fetch('/api/data/graph/add', {
method: 'POST',
headers: {
"Content-Type": "application/json" // Content-Type
},
body: JSON.stringify({
file_path: files[0]
}),

View File

@ -83,7 +83,69 @@
<div class="setting" v-if="state.section == 'model'">
<h3>模型配置</h3>
<p>请在 <code>src/.env</code> 文件中配置对应的 APIKEY</p>
<div class="model-provider-card" v-for="(item, key) in modelKeys" :key="key" :style="{backgroundColor: modelProvider == item ? colorMap[item] : white }">
<div class="model-provider-card">
<div class="card-header">
<h3>自定义模型</h3>
</div>
<div class="card-body">
<div
:class="{'model_selected': modelProvider == 'custom' && configStore.config.model_name == item.name, 'card-models': true, 'custom-model': true}"
v-for="(item, key) in configStore.config.custom_models" :key="key"
@click="handleChange('model_provider', 'custom'); handleChange('model_name', item.name)"
>
<div class="card-models__header">
<div class="name">{{ item.name }}</div>
<div class="action">
<!-- 添加确认删除 -->
<a-popconfirm
title="确认删除该模型?"
@confirm="handleDeleteCustomModel(item.name)"
okText="确认删除"
cancelText="取消"
ok-type="danger"
:disabled="configStore.config.model_name == item.name"
>
<a-button type="text" :disabled="configStore.config.model_name == item.name" @click.stop><DeleteOutlined /></a-button>
</a-popconfirm>
<a-button type="text" @click.stop="handleEditCustomModel(item)"><EditOutlined /></a-button>
</div>
</div>
<div class="api_base">{{ item.api_base }}</div>
<!-- <div class="select-btn"></div> -->
</div>
<div class="card-models custom-model" @click="customModel.visible=true">
<div class="card-models__header">
<div class="name"> + 添加模型</div>
</div>
<div class="api_base">添加兼容 OpenAI 的模型</div>
<a-modal
class="custom-model-modal"
v-model:open="customModel.visible"
:title="customModel.modelTitle"
@ok="handleAddCustomModel"
@cancel="handleCancelCustomModel"
:okText="'确认'"
:cancelText="'取消'"
:okButtonProps="{disabled: !customModel.name || !customModel.api_base}"
:ok-type="'primary'"
>
<p>添加的模型是兼容 OpenAI 的模型比如 vllmOllama</p>
<a-form :model="customModel" layout="vertical" >
<a-form-item label="模型名称" name="name" :rules="[{ required: true, message: '请输入模型名称' }]">
<a-input v-model:value="customModel.name" />
</a-form-item>
<a-form-item label="API Base" name="api_base" :rules="[{ required: true, message: '请输入API Base' }]">
<a-input v-model:value="customModel.api_base" type="password"/>
</a-form-item>
<a-form-item label="API KEY" name="api_key">
<a-input v-model:value="customModel.api_key" autocomplete="off"/>
</a-form-item>
</a-form>
</a-modal>
</div>
</div>
</div>
<div class="model-provider-card" v-for="(item, key) in modelKeys" :key="key">
<div class="card-header">
<h3>{{ modelNames[item].name }}</h3>
<a :href="modelNames[item].url" target="_blank">详情</a>
@ -96,7 +158,7 @@
@click="handleChange('model_provider', item); handleChange('model_name', model)"
>
<div class="model_name">{{ model }}</div>
<div class="select-btn"></div>
<!-- <div class="select-btn"></div> -->
</div>
</div>
</div>
@ -116,7 +178,7 @@
<script setup>
import { message } from 'ant-design-vue';
import { computed, reactive, ref, h } from 'vue'
import { computed, reactive, ref, h, watch } from 'vue'
import { useConfigStore } from '@/stores/config';
import {
ReloadOutlined,
@ -124,8 +186,11 @@ import {
CodeOutlined,
ExceptionOutlined,
FolderOutlined,
DeleteOutlined,
EditOutlined,
} from '@ant-design/icons-vue';
import HeaderComponent from '@/components/HeaderComponent.vue';
import { notification, Button } from 'ant-design-vue';
const configStore = useConfigStore()
const items = computed(() => configStore.config._config_items)
@ -133,21 +198,19 @@ const modelNames = computed(() => configStore.config?.model_names)
const modelStatus = computed(() => configStore.config?.model_provider_status)
const modelProvider = computed(() => configStore.config?.model_provider)
const isNeedRestart = ref(false)
const customModel = reactive({
modelTitle: '添加自定义模型',
visible: false,
name: '',
api_key: '',
api_base: '',
edit_type: 'add',
})
const state = reactive({
loading: false,
section: 'base'
})
const colorMap = reactive({
siliconflow: '#FFECFF',
zhipu: '#EFF1FE',
qianfan: '#E8F5FE',
deepseek: '#D3DCFF',
openai: '#E5E7EB',
vllm: '#E5E7EB',
bailian: '#EFF1FE',
})
// modelStatus key
const modelKeys = computed(() => {
return Object.keys(modelStatus.value).filter(key => modelStatus.value[key])
@ -179,12 +242,73 @@ const handleChange = (key, e) => {
|| key == 'reranker') {
if (!isNeedRestart.value) {
isNeedRestart.value = true
message.info('修改配置后需要重启服务才能生效')
notification.info({
message: '需要重启服务',
description: '请点击右下角按钮重启服务',
placement: 'topLeft',
duration: 0,
btn: h(Button, { type: 'primary', onClick: sendRestart }, '立即重启')
})
}
}
configStore.setConfigValue(key, e)
}
const handleAddCustomModel = async () => {
if (!customModel.name || !customModel.api_base) {
message.error('请填写完整模型信息')
return
}
if (!configStore.config.custom_models) {
configStore.config.custom_models = []
}
if (configStore.config.custom_models.find(item => item.name == customModel.name)) {
message.error('模型名称已存在')
return
}
if (customModel.edit_type == 'add') {
configStore.config.custom_models.push(customModel)
} else {
configStore.config.custom_models = configStore.config.custom_models.map(item => {
if (item.name == customModel.name) {
return customModel
}
return item
})
}
customModel.visible = false
await configStore.setConfigValue('custom_models', configStore.config.custom_models)
configStore.refreshConfig()
message.success('添加自定义模型成功')
}
const handleDeleteCustomModel = (name) => {
configStore.config.custom_models = configStore.config.custom_models.filter(item => item.name != name)
configStore.setConfigValue('custom_models', configStore.config.custom_models)
configStore.refreshConfig()
}
const handleEditCustomModel = (item) => {
customModel.modelTitle = '编辑自定义模型'
customModel.name = item.name
customModel.api_key = item.api_key
customModel.api_base = item.api_base
customModel.visible = true
customModel.edit_type = 'edit'
}
const handleCancelCustomModel = () => {
customModel.name = ''
customModel.api_key = ''
customModel.api_base = ''
customModel.visible = false
}
const sendRestart = () => {
console.log('Restarting...')
message.loading({ content: '重新加载模型中', key: "restart", duration: 0 });
@ -213,8 +337,8 @@ const sendRestart = () => {
padding: 0;
box-sizing: border-box;
display: flex;
background: inherit;
position: relative;
min-height: 100%;
}
.sider {
@ -295,8 +419,8 @@ const sendRestart = () => {
}
.model-provider-card {
background-color: var(--gray-10);
border: 1px solid var(--gray-300);
background-color: white;
border-radius: 8px;
margin-bottom: 16px;
padding: 16px;
@ -328,9 +452,9 @@ const sendRestart = () => {
.success {
width: 1rem;
height: 1rem;
background-color: green;
background-color: rgb(91, 186, 91);
border-radius: 50%;
box-shadow: 0 0 10px 1px rgba( 0,128, 0, 0.5);
box-shadow: 0 0 10px 1px rgba( 0,128, 0, 0.2);
border: 2px solid white;
}
@ -351,7 +475,6 @@ const sendRestart = () => {
width: 100%;
border-radius: 8px;
border: 1px solid var(--gray-300);
background-color: var(--gray-50);
padding: 10px 16px;
display: flex;
gap: 6px;
@ -359,9 +482,9 @@ const sendRestart = () => {
align-items: center;
cursor: pointer;
box-sizing: border-box;
background-color: rgba(255, 255, 255, 0.6);
background-color: var(--gray-10);
transition: box-shadow 0.1s;
&:hover {
border-color: var(--gray-400);
box-shadow: 0 2px 4px rgba(0, 0, 0, 0.05);
}
.model_name {
@ -386,7 +509,60 @@ const sendRestart = () => {
border: 2px solid var(--main-color);
}
}
&.custom-model {
display: flex;
flex-direction: column;
align-items: flex-start;
padding-right: 8px;
gap: 10px;
.card-models__header {
width: 100%;
height: 32px;
display: flex;
justify-content: flex-start;
align-items: center;
.name {
color: var(--gray-1000);
font-weight: bold;
}
.action {
opacity: 0;
user-select: none;
margin-left: auto;
button {
padding: 4px 8px;
}
}
.custom-model-modal {
.ant-form-item {
margin-bottom: 10px;
}
}
}
.api_base {
font-size: 12px;
color: var(--gray-600);
}
&:hover {
.card-models__header {
.action {
opacity: 1;
}
}
}
}
&.model_selected.custom-model {
padding: 9px 7px 9px 15px;
.card-models__header {
.action {
opacity: 1;
}
}
}
}
}
}

View File

@ -1,27 +1,28 @@
import { fileURLToPath, URL } from 'node:url'
import { defineConfig } from 'vite'
import { defineConfig, loadEnv } from 'vite'
import vue from '@vitejs/plugin-vue'
// https://vitejs.dev/config/
export default defineConfig({
plugins: [vue()],
resolve: {
alias: {
'@': fileURLToPath(new URL('./src', import.meta.url))
}
},
// 定义代理
server: {
proxy: {
'^/api': {
target: 'http://127.0.0.1:5000', // 5000端口是flask的Debug模式默认端口, 8000是非Debug模式默认端口
changeOrigin: true,
rewrite: (path) => path.replace(/^\/api/, '')
export default defineConfig(({ mode }) => {
const env = loadEnv(mode, process.cwd(), '')
return {
plugins: [vue()],
resolve: {
alias: {
'@': fileURLToPath(new URL('./src', import.meta.url))
}
},
watch: {
ignored: ['**/node_modules/**', '**/dist/**'],
},
server: {
proxy: {
'^/api': {
target: env.VITE_API_URL || 'http://localhost:5000',
changeOrigin: true,
rewrite: (path) => path.replace(/^\/api/, '')
}
},
watch: {
ignored: ['**/node_modules/**', '**/dist/**'],
},
host: '0.0.0.0',
}
}
})