소스 검색

rag-gateway:注册,登录

zhangqian 1 년 전
커밋
3604703a9e
20개의 변경된 파일611개의 추가작업 그리고 0개의 파일을 삭제
  1. 10 0
      .gitignore
  2. 14 0
      Dockerfile
  3. 60 0
      README.md
  4. 10 0
      app/api/__init__.py
  5. 84 0
      app/api/auth.py
  6. 48 0
      app/config/config.py
  7. 9 0
      app/config/config.yaml
  8. 25 0
      app/models/base_model.py
  9. 51 0
      app/models/token_model.py
  10. 23 0
      app/models/user.py
  11. 10 0
      app/models/user_model.py
  12. 49 0
      app/service/auth.py
  13. 44 0
      app/service/bisheng.py
  14. 32 0
      app/service/ragflow.py
  15. 76 0
      app/utils/rsa_crypto.py
  16. 16 0
      main.py
  17. 38 0
      main.spec
  18. 1 0
      pip_install.sh
  19. BIN
      requirements.txt
  20. 11 0
      test_main.http

+ 10 - 0
.gitignore

@@ -0,0 +1,10 @@
+venv
+dist
+build
+main
+__pycache__
+.pytest_cache
+.idea
+.vscode
+.ipynb_checkpoints
+.pytest

+ 14 - 0
Dockerfile

@@ -0,0 +1,14 @@
+FROM python:3.11
+
+# 安装 PyInstaller
+RUN pip install pyinstaller
+
+# 复制项目文件到容器中
+COPY . /app
+WORKDIR /app
+
+# 安装项目依赖
+RUN pip install -r requirements.txt
+
+# 使用 PyInstaller 打包应用,并确保文件被复制到 dist 目录
+RUN pyinstaller -F main.py

+ 60 - 0
README.md

@@ -0,0 +1,60 @@
+# RAG Gateway Project
+
+## 项目简介
+
+RAG Gateway 是一个用于处理请求和响应的网关服务。它使用 FastAPI 和其他相关库来提供高效的 API 服务。
+
+## 目录结构
+rag-gateway/
+├── app/
+│   └── init.py
+│   └── main.py
+├── main.py
+├── requirements.txt
+├── test_main.http
+├── venv/
+└── README.txt
+
+
+
+## 运行
+
+### 1. 创建虚拟环境
+
+首先,创建一个虚拟环境并激活它:
+
+```bash
+python3 -m venv venv
+```
+```bash
+source venv/bin/activate  # 对于 Linux/Mac
+```
+或者
+```bash
+venv\Scripts\activate  # 对于 Windows
+```
+### 2. 安装依赖
+```bash
+pip install -r requirements.txt
+```
+
+
+### 3. 运行
+```bash
+python main.py
+```
+
+## 部署
+
+### 1. 打包成二进制文件
+
+#### 构建 Docker 镜像:
+```bash
+docker build -t my-python-app .
+```
+#### 2.运行 Docker 容器:
+
+```bash
+docker run --rm -v ${PWD}:/app -v ${PWD}/:/app/dist my-python-app
+```
+#### 3. 获取生成的二进制文件: 生成的二进制文件会出现在 dist 目录中

+ 10 - 0
app/api/__init__.py

@@ -0,0 +1,10 @@
+from fastapi import FastAPI
+from pydantic import BaseModel
+
+app = FastAPI()
+
+
+class Response(BaseModel):
+    code: int = 200
+    msg: str = ""
+    data: dict = {}

+ 84 - 0
app/api/auth.py

@@ -0,0 +1,84 @@
+from typing import Dict
+import json
+
+from fastapi import APIRouter, Depends, HTTPException
+from fastapi.security import OAuth2PasswordBearer
+from passlib.context import CryptContext
+from sqlalchemy.orm import Session
+
+from app.api import Response
+from app.config.config import settings
+from app.models.base_model import get_db
+from app.models.token_model import upsert_token
+from app.models.user import User, UserCreate, LoginData
+from app.models.user_model import UserModel
+from app.service.auth import authenticate_user, create_access_token
+from app.service.bisheng import BishengService
+from app.service.ragflow import RagflowService
+
+router = APIRouter()
+
+pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
+oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
+
+
+@router.post("/register", response_model=Response)
+async def register(user: UserCreate, db=Depends(get_db)):
+    db_user = db.query(UserModel).filter(UserModel.username == user.username).first()
+    if db_user:
+        return Response(code=200, msg="Username already registered")
+
+    bisheng_service = BishengService(settings.bisheng_base_url)
+    ragflow_service = RagflowService(settings.ragflow_base_url)
+
+    # 注册到毕昇
+    try:
+        await bisheng_service.register(user.username, user.password)
+    except Exception as e:
+        return Response(code=500, msg=f"Failed to register with Bisheng: {str(e)}")
+
+    # 注册到ragflow
+    try:
+        await ragflow_service.register(user.username, user.password)
+    except Exception as e:
+        return Response(code=500, msg=f"Failed to register with Ragflow: {str(e)}")
+
+    # 存储用户信息
+    hashed_password = pwd_context.hash(user.password)
+    db_user = UserModel(username=user.username, hashed_password=hashed_password)
+    db.add(db_user)
+    db.commit()
+    db.refresh(db_user)
+    return Response(code=200, msg="User registered successfully",data={"username": db_user.username})
+
+
+@router.post("/login", response_model=Response)
+async def login(login_data: LoginData, db: Session = Depends(get_db)):
+    user = authenticate_user(db, login_data.username, login_data.password)
+    if not user:
+        return Response(code=400, msg="Incorrect username or password")
+
+    bisheng_service = BishengService(settings.bisheng_base_url)
+    ragflow_service = RagflowService(settings.ragflow_base_url)
+
+    # 登录到毕昇
+    try:
+        bisheng_token = await bisheng_service.login(login_data.username, login_data.password)
+    except Exception as e:
+        return Response(code=500, msg=f"Failed to login with Bisheng: {str(e)}")
+
+    # 登录到ragflow
+    try:
+        ragflow_token = await ragflow_service.login(login_data.username, login_data.password)
+    except Exception as e:
+        return Response(code=500, msg=f"Failed to login with Ragflow: {str(e)}")
+
+    # 创建本地token
+    access_token = create_access_token(data={"sub": user.username})
+
+    upsert_token(db, user.id, access_token, bisheng_token, ragflow_token)
+
+    return Response(code=200, msg="Login successful", data={
+        "access_token": access_token,
+        "token_type": "bearer"
+    })

+ 48 - 0
app/config/config.py

@@ -0,0 +1,48 @@
+import os
+from pathlib import Path
+import yaml
+
+
+class Settings:
+    secret_key: str = ''
+    bisheng_base_url: str = ''
+    ragflow_base_url: str = ''
+    database_url: str = ''
+    PUBLIC_KEY: str
+    PRIVATE_KEY: str
+
+    def __init__(self, **kwargs):
+        # Check if all required fields are provided and set them
+        for field in self.__annotations__.keys():
+            if field not in kwargs:
+                raise ValueError(f"Missing setting: {field}")
+            setattr(self, field, kwargs[field])
+
+    def to_dict(self):
+        """Return the settings as a dictionary."""
+        return {k: getattr(self, k) for k in self.__annotations__.keys()}
+
+    def __repr__(self):
+        """Return a string representation of the settings."""
+        return f"Settings({self.to_dict()})"
+
+
+def load_yaml(file_path: Path) -> dict:
+    with file_path.open('r', encoding="utf-8") as fr:
+        try:
+            data = yaml.safe_load(fr)
+            return data
+        except yaml.YAMLError as e:
+            print(f"Error loading YAML file {file_path}: {e}")
+            return {}
+
+
+# Use pathlib to handle file paths
+config_yaml_path = Path(__file__).parent / 'config.yaml'
+settings_data = load_yaml(config_yaml_path)
+
+# Initialize settings object
+settings = Settings(**settings_data)
+
+# Print the loaded settings
+print(f"Loaded settings: {settings}")

+ 9 - 0
app/config/config.yaml

@@ -0,0 +1,9 @@
+secret_key: your-secret-key
+bisheng_base_url: http://192.168.20.119:13001
+ragflow_base_url: http://192.168.20.119:11080
+database_url: mysql+pymysql://root:infini_rag_flow@192.168.20.116:5455/rag_basic
+PUBLIC_KEY: |
+  -----BEGIN PUBLIC KEY-----
+  MIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEArq9XTUSeYr2+N1h3Afl/z8Dse/2yD0ZGrKwx+EEEcdsBLca9Ynmx3nIB5obmLlSfmskLpBo0UACBmB5rEjBp2Q2f3AG3Hjd4B+gNCG6BDaawuDlgANIhGnaTLrIqWrrcm4EMzJOnAOI1fgzJRsOOUEfaS318Eq9OVO3apEyCCt0lOQK6PuksduOjVxtltDav+guVAA068NrPYmRNabVKRNLJpL8w4D44sfth5RvZ3q9t+6RTArpEtc5sh5ChzvqPOzKGMXW83C95TxmXqpbK6olN4RevSfVjEAgCydH6HN6OhtOQEcnrU97r9H0iZOWwbw3pVrZiUkuRD1R56Wzs2wIDAQAB
+  -----END PUBLIC KEY-----
+PRIVATE_KEY: str

+ 25 - 0
app/models/base_model.py

@@ -0,0 +1,25 @@
+from sqlalchemy import create_engine
+from sqlalchemy.ext.declarative import declarative_base
+from sqlalchemy.orm import sessionmaker, Session
+
+from app.config.config import settings
+
+DATABASE_URL = settings.database_url
+
+engine = create_engine(DATABASE_URL)
+SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
+
+Base = declarative_base()
+
+
+# 创建所有表(如果有新的模型类,会自动创建相应的表)
+def init_db():
+    Base.metadata.create_all(bind=engine)
+
+
+def get_db():
+    db = SessionLocal()
+    try:
+        yield db
+    finally:
+        db.close()

+ 51 - 0
app/models/token_model.py

@@ -0,0 +1,51 @@
+from datetime import datetime
+
+from sqlalchemy import Column, Integer, String, DateTime, Text
+from sqlalchemy.orm import Session
+
+from app.models.base_model import Base
+
+
+class TokenModel(Base):
+    __tablename__ = "token"
+    id = Column(Integer, primary_key=True, index=True)
+    user_id = Column(Integer, index=True)
+    token = Column(Text(10000), unique=True, index=True)
+    bisheng_token = Column(Text(10000), unique=True, index=True)
+    ragflow_token = Column(Text(10000), unique=True, index=True)
+    created_at = Column(DateTime, default=datetime.utcnow)
+
+
+def upsert_token(db: Session, user_id: int, access_token: str, bisheng_token: str, ragflow_token: str):
+    # 参数验证
+    if not isinstance(user_id, int) or user_id <= 0:
+        return
+    if not access_token or not bisheng_token or not ragflow_token:
+        return
+    db_token = None
+    try:
+        # 查询现有记录
+        existing_token = db.query(TokenModel).filter_by(user_id=user_id).first()
+
+        if existing_token:
+            # 记录存在,进行更新
+            existing_token.token = access_token
+            existing_token.bisheng_token = bisheng_token
+            existing_token.ragflow_token = ragflow_token
+        else:
+            # 记录不存在,进行插入
+            db_token = TokenModel(
+                user_id=user_id,
+                token=access_token,
+                bisheng_token=bisheng_token,
+                ragflow_token=ragflow_token
+            )
+            db.add(db_token)
+
+        # 提交事务
+        db.commit()
+        db.refresh(db_token)
+
+    except Exception as e:
+        # 异常处理
+        db.rollback()  # 回滚事务

+ 23 - 0
app/models/user.py

@@ -0,0 +1,23 @@
+from pydantic import BaseModel
+
+
+class UserCreate(BaseModel):
+    username: str
+    password: str
+
+
+# 定义请求体模型
+class LoginData(BaseModel):
+    username: str
+    password: str
+
+
+class User(BaseModel):
+    username: str
+
+
+class Token(BaseModel):
+    access_token: str
+    token_type: str
+    bisheng_token: str
+    ragflow_token: str

+ 10 - 0
app/models/user_model.py

@@ -0,0 +1,10 @@
+from sqlalchemy import Column, Integer, String
+
+from app.models.base_model import Base
+
+
+class UserModel(Base):
+    __tablename__ = "user"
+    id = Column(Integer, primary_key=True, index=True)
+    username = Column(String(255), unique=True, index=True)
+    hashed_password = Column(String(255))

+ 49 - 0
app/service/auth.py

@@ -0,0 +1,49 @@
+from datetime import datetime, timedelta
+from jwt import encode, decode, exceptions
+from passlib.context import CryptContext
+from fastapi import HTTPException, status
+
+from app.config.config import settings
+from app.models.user_model import UserModel
+
+SECRET_KEY = settings.secret_key
+ALGORITHM = "HS256"
+ACCESS_TOKEN_EXPIRE_MINUTES = 30
+
+pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
+
+
+def verify_password(plain_password, hashed_password):
+    return pwd_context.verify(plain_password, hashed_password)
+
+
+def get_password_hash(password):
+    return pwd_context.hash(password)
+
+
+def authenticate_user(db, username: str, password: str):
+    user = db.query(UserModel).filter(UserModel.username == username).first()
+    if not user:
+        return False
+    if not verify_password(password, user.hashed_password):
+        return False
+    return user
+
+
+def create_access_token(data: dict, expires_delta: timedelta = None):
+    to_encode = data.copy()
+    if expires_delta:
+        expire = datetime.utcnow() + expires_delta
+    else:
+        expire = datetime.utcnow() + timedelta(minutes=15)
+    to_encode.update({"exp": expire})
+    encoded_jwt = encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
+    return encoded_jwt
+
+
+def decode_access_token(token: str):
+    try:
+        payload = decode(token, SECRET_KEY, algorithms=[ALGORITHM])
+        return payload
+    except exceptions.DecodeError:
+        raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Could not validate credentials")

+ 44 - 0
app/service/bisheng.py

@@ -0,0 +1,44 @@
+import httpx
+
+from app.config.config import settings
+from app.utils.rsa_crypto import BishengCrypto
+
+
+class BishengService:
+    def __init__(self, base_url: str):
+        self.base_url = base_url
+
+    async def register(self, username: str, password: str):
+        public_key = await self.get_public_key_api()
+        password = BishengCrypto(public_key, settings.PRIVATE_KEY).encrypt(password)
+        async with httpx.AsyncClient() as client:
+            response = await client.post(
+                f"{self.base_url}/api/v1/user/regist",
+                json={"user_name": username, "password": password},
+                headers={'Content-Type': 'application/json'}
+            )
+            if response.status_code != 200 and response.status_code != 201:
+                raise Exception(f"Bisheng registration failed: {response.text}")
+
+    async def login(self, username: str, password: str) -> str:
+        public_key = await self.get_public_key_api()
+        password = BishengCrypto(public_key, settings.PRIVATE_KEY).encrypt(password)
+        async with httpx.AsyncClient() as client:
+            response = await client.post(
+                f"{self.base_url}/api/v1/user/login",
+                json={"user_name": username, "password": password},
+                headers={'Content-Type': 'application/json'}
+            )
+            if response.status_code != 200 and response.status_code != 201:
+                raise Exception(f"Bisheng login failed: {response.text}")
+            return response.json().get('data', {}).get('access_token')
+
+    async def get_public_key_api(self) -> dict:
+        async with httpx.AsyncClient() as client:
+            response = await client.get(
+                f"{self.base_url}/api/v1/user/public_key",
+                headers={'Content-Type': 'application/json'}
+            )
+            if response.status_code != 200:
+                raise Exception(f"Failed to get public key: {response.text}")
+            return response.json().get('data', {}).get('public_key')

+ 32 - 0
app/service/ragflow.py

@@ -0,0 +1,32 @@
+import httpx
+
+from app.config.config import settings
+from app.utils.rsa_crypto import RagflowCrypto
+
+
+class RagflowService:
+    def __init__(self, base_url: str):
+        self.base_url = base_url
+
+    async def register(self, username: str, password: str):
+        password = RagflowCrypto(settings.PUBLIC_KEY, settings.PRIVATE_KEY).encrypt(password)
+        async with httpx.AsyncClient() as client:
+            response = await client.post(
+                f"{self.base_url}/v1/user/register",
+                json={"nickname": username, "email": f"{username}@example.com", "password": password},
+                headers={'Content-Type': 'application/json'}
+            )
+            if response.status_code != 200:
+                raise Exception(f"Ragflow registration failed: {response.text}")
+
+    async def login(self, username: str, password: str) -> str:
+        password = RagflowCrypto(settings.PUBLIC_KEY, settings.PRIVATE_KEY).encrypt(password)
+        async with httpx.AsyncClient() as client:
+            response = await client.post(
+                f"{self.base_url}/v1/user/login",
+                json={"email": f"{username}@example.com", "password": password},
+                headers={'Content-Type': 'application/json'}
+            )
+            if response.status_code != 200:
+                raise Exception(f"Ragflow login failed: {response.text}")
+            return response.json().get('data', {}).get('access_token')

+ 76 - 0
app/utils/rsa_crypto.py

@@ -0,0 +1,76 @@
+from abc import ABC, abstractmethod
+from Cryptodome.PublicKey import RSA
+from Cryptodome.Cipher import PKCS1_v1_5
+import base64
+import rsa
+
+
+# 定义抽象基类
+class RSACrypto(ABC):
+
+    @abstractmethod
+    def encrypt(self, password: str) -> str:
+        pass
+
+    @abstractmethod
+    def decrypt(self, encrypted_password: str) -> str:
+        pass
+
+
+# 实现 RagflowCrypto 类
+class RagflowCrypto(RSACrypto):
+
+    def __init__(self, public_key: str, private_key: str):
+        self.public_key = public_key
+        self.private_key = private_key
+
+    def encrypt(self, password: str) -> str:
+        rsa_key = RSA.importKey(self.public_key)
+        cipher = PKCS1_v1_5.new(rsa_key)
+        encrypted_password = cipher.encrypt(base64.b64encode(password.encode('utf-8')))
+        return base64.b64encode(encrypted_password).decode('utf-8')
+
+    def decrypt(self, encrypted_password: str) -> str:
+        rsa_key = RSA.importKey(self.private_key)
+        cipher = PKCS1_v1_5.new(rsa_key)
+        encrypted_password_bytes = base64.b64decode(encrypted_password)
+        decoded_password = cipher.decrypt(encrypted_password_bytes, "Fail to decrypt password!")
+        return base64.b64decode(decoded_password).decode('utf-8')
+
+
+# 实现 BishengCrypto 类
+# class BishengCrypto(RSACrypto):
+#
+#     def __init__(self, public_key, private_key: str):
+#         self.public_key = public_key
+#         self.private_key = private_key
+#
+#     def encrypt(self, password: str) -> str:
+#         rsa_key = RSA.importKey(self.public_key)
+#         cipher = PKCS1_v1_5.new(rsa_key)
+#         encrypted_password = cipher.encrypt(password.encode('utf-8'))
+#         return base64.b64encode(encrypted_password).decode('utf-8')
+#
+#     def decrypt(self, encrypted_password: str) -> str:
+#         rsa_key = RSA.importKey(self.private_key)
+#         cipher = PKCS1_v1_5.new(rsa_key)
+#         encrypted_password_bytes = base64.b64decode(encrypted_password)
+#         decoded_password = cipher.decrypt(encrypted_password_bytes, "Fail to decrypt password!")
+#         return decoded_password.decode('utf-8')
+
+
+class BishengCrypto:
+
+    def __init__(self, public_key: str, private_key: str):
+        self.public_key = rsa.PublicKey.load_pkcs1(public_key.encode('utf-8'))
+
+    def encrypt(self, password: str) -> str:
+        encrypted_password = rsa.encrypt(password.encode('utf-8'), self.public_key)
+        return base64.b64encode(encrypted_password).decode('utf-8')
+
+    @classmethod
+    def decrypt(cls, password: str, private_key: str) -> str:
+        private_key = rsa.PrivateKey.load_pkcs1(private_key.encode('utf-8'))
+        encrypted_password_bytes = base64.b64decode(password)
+        return rsa.decrypt(encrypted_password_bytes, private_key).decode('utf-8')
+

+ 16 - 0
main.py

@@ -0,0 +1,16 @@
+from fastapi import FastAPI
+from app.api.auth import router as auth_router
+from app.models.base_model import init_db
+
+init_db()
+app = FastAPI(
+  title="basic_rag_gateway",
+  version="0.1",
+  description="",
+)
+
+app.include_router(auth_router, prefix='/auth', tags=["auth"])
+
+if __name__ == "__main__":
+    import uvicorn
+    uvicorn.run(app, host="0.0.0.0", port=9201)

+ 38 - 0
main.spec

@@ -0,0 +1,38 @@
+# -*- mode: python ; coding: utf-8 -*-
+
+
+a = Analysis(
+    ['main.py'],
+    pathex=[],
+    binaries=[],
+    datas=[],
+    hiddenimports=[],
+    hookspath=[],
+    hooksconfig={},
+    runtime_hooks=[],
+    excludes=[],
+    noarchive=False,
+    optimize=0,
+)
+pyz = PYZ(a.pure)
+
+exe = EXE(
+    pyz,
+    a.scripts,
+    a.binaries,
+    a.datas,
+    [],
+    name='main',
+    debug=False,
+    bootloader_ignore_signals=False,
+    strip=False,
+    upx=True,
+    upx_exclude=[],
+    runtime_tmpdir=None,
+    console=True,
+    disable_windowed_traceback=False,
+    argv_emulation=False,
+    target_arch=None,
+    codesign_identity=None,
+    entitlements_file=None,
+)

+ 1 - 0
pip_install.sh

@@ -0,0 +1 @@
+pip install PyMySQL & pip install fastapi & pip install sqlalchemy & pip install PyJWT & pip install rsa & pip install httpx & pip install uvicorn & pip install bcrypt & pip install PyYAML & pip install pycryptodomex & pip install passlib

BIN
requirements.txt


+ 11 - 0
test_main.http

@@ -0,0 +1,11 @@
+# Test your FastAPI endpoints
+
+GET http://127.0.0.1:8000/
+Accept: application/json
+
+###
+
+GET http://127.0.0.1:8000/hello/User
+Accept: application/json
+
+###