Rocky Linux 私有 RAG 知识库进阶|LDAP 域控登录适配 + 企业级权限落地(无缝升级完整版)

本次基于原有豆包RAG私有知识库方案做全方位增强,新增网页文件上传、打字机流式输出、问答历史本地存储、账号密码登录鉴权、答案溯源展示五大功能。全程保留原有项目结构与部署方式,仅小幅新增依赖,兼容性极强。

一、新增/修改项目依赖

在原有依赖基础上,执行以下命令安装新增依赖包:

pip install python-jose[cryptography] python-multipart

二、配置文件修改(.env)

保留原有所有配置,新增 JWT 鉴权、登录账号密码配置,自行替换自定义密钥、账号密码:

# 原有核心配置无需修改
ARK_API_KEY=你的方舟API密钥
ARK_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
ARK_MODEL_ID=你的接入点EndpointID
CHROMA_DB_PATH=./chroma_storage
TEMPERATURE=0.1
CHUNK_SIZE=600
CHUNK_OVERLAP=100
RETRIEVAL_TOP_K=4

# 新增登录鉴权配置
LOGIN_USER=admin
LOGIN_PASS=admin123
SECRET_KEY=自定义随机长字符串
ALGORITHM=HS256
ACCESS_TOKEN_EXPIRE_MINUTES=480

三、新增核心模块文件

1. 登录鉴权模块(auth.py)

实现 JWT 令牌生成、用户校验、接口全局鉴权拦截功能:

# auth.py
import os
from datetime import datetime, timedelta
from typing import Optional
from jose import JWTError, jwt
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from dotenv import load_dotenv

load_dotenv()

SECRET_KEY = os.getenv("SECRET_KEY")
ALGORITHM = os.getenv("ALGORITHM")
ACCESS_TOKEN_EXPIRE_MINUTES = int(os.getenv("ACCESS_TOKEN_EXPIRE_MINUTES", 480))
LOGIN_USER = os.getenv("LOGIN_USER")
LOGIN_PASS = os.getenv("LOGIN_PASS")

oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/login")

def create_access_token(data: dict, expires_delta: Optional[timedelta] = None):
    to_encode = data.copy()
    expire = datetime.utcnow() + (expires_delta or timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES))
    to_encode.update({"exp": expire})
    return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)

def verify_user(username: str, password: str):
    if username == LOGIN_USER and password == LOGIN_PASS:
        return True
    return False

async def get_current_user(token: str = Depends(oauth2_scheme)):
    credentials_exception = HTTPException(
        status_code=status.HTTP_401_UNAUTHORIZED,
        detail="Could not validate credentials",
        headers={"WWW-Authenticate": "Bearer"},
    )
    try:
        payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
        username: str = payload.get("sub")
        if username != LOGIN_USER:
            raise credentials_exception
    except JWTError:
        raise credentials_exception
    return username

2. 问答历史存储模块(database.py)

基于 SQLite 实现问答历史本地持久化,支持查询、清空历史记录:

# database.py
import sqlite3
import datetime

DB_FILE = "chat_history.db"

def init_db():
    conn = sqlite3.connect(DB_FILE)
    c = conn.cursor()
    c.execute('''CREATE TABLE IF NOT EXISTS history
                 (id INTEGER PRIMARY KEY AUTOINCREMENT,
                  username TEXT,
                  question TEXT,
                  answer TEXT,
                  timestamp TEXT)''')
    conn.commit()
    conn.close()

def add_history(username, question, answer):
    conn = sqlite3.connect(DB_FILE)
    c = conn.cursor()
    c.execute("INSERT INTO history (username, question, answer, timestamp) VALUES (?,?,?,?)",
              (username, question, answer, datetime.datetime.now().isoformat()))
    conn.commit()
    conn.close()

def get_history(username, limit=50):
    conn = sqlite3.connect(DB_FILE)
    c = conn.cursor()
    c.execute("SELECT question, answer, timestamp FROM history WHERE username=? ORDER BY id DESC LIMIT ?", (username, limit))
    rows = c.fetchall()
    conn.close()
    return [{"question": q, "answer": a, "time": t} for q, a, t in rows]

def clear_history(username):
    conn = sqlite3.connect(DB_FILE)
    c = conn.cursor()
    c.execute("DELETE FROM history WHERE username=?", (username,))
    conn.commit()
    conn.close()

四、核心RAG逻辑升级(doubao_rag_core.py)

新增流式生成接口答案溯源解析,保留原有非流式问答接口,兼容钉钉等第三方调用:

# doubao_rag_core.py
__import__('pysqlite3')
import sys
sys.modules['sqlite3'] = sys.modules.pop('pysqlite3')

import os
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
from dotenv import load_dotenv
from openai import OpenAI
from langchain_community.vectorstores import Chroma
from langchain_community.embeddings import HuggingFaceEmbeddings
from typing import List, Dict, Generator

load_dotenv()

client = OpenAI(
    api_key=os.getenv("ARK_API_KEY"),
    base_url=os.getenv("ARK_BASE_URL")
)

embedding = HuggingFaceEmbeddings(
    model_name="BAAI/bge-small-zh-v1.5",
    cache_folder="./model_cache",
    model_kwargs={"trust_remote_code": True}
)

vectordb = Chroma(
    persist_directory=os.getenv("CHROMA_DB_PATH"),
    embedding_function=embedding
)
retriever = vectordb.as_retriever(search_kwargs={"k": int(os.getenv("RETRIEVAL_TOP_K"))})

def extract_metadata(docs) -> List[Dict]:
    """从检索到的文档中提取文件名和页码"""
    refs = []
    for doc in docs:
        metadata = doc.metadata
        source = metadata.get("source", "未知文件")
        page = metadata.get("page", None)
        refs.append({
            "source": os.path.basename(source),
            "page": page + 1 if page is not None else None  # 页码从0开始,转为1-based
        })
    return refs

def get_answer(user_question: str):
    # 检索
    related_docs = retriever.get_relevant_documents(user_question)
    context_text = "\n====文档片段====\n".join([d.page_content for d in related_docs])
    refs = extract_metadata(related_docs)

    sys_prompt = f"""
你是企业内部知识库问答助手,严格依据下方【内部参考文档】回答员工问题。
禁止编造文档不存在的信息,无匹配资料直接回复:【暂无该问题对应的内部资料】
回答简洁准确,不要多余话术。
【内部参考文档】
{context_text}
"""
    resp = client.chat.completions.create(
        model=os.getenv("ARK_MODEL_ID"),
        messages=[
            {"role": "system", "content": sys_prompt},
            {"role": "user", "content": user_question}
        ],
        temperature=float(os.getenv("TEMPERATURE")),
    )
    answer = resp.choices[0].message.content
    return answer, refs

def get_answer_stream(user_question: str) -> Generator:
    """流式生成器,每次yield一个token,最后yield引用列表"""
    related_docs = retriever.get_relevant_documents(user_question)
    context_text = "\n====文档片段====\n".join([d.page_content for d in related_docs])
    refs = extract_metadata(related_docs)

    sys_prompt = f"""
你是企业内部知识库问答助手,严格依据下方【内部参考文档】回答员工问题。
禁止编造文档不存在的信息,无匹配资料直接回复:【暂无该问题对应的内部资料】
回答简洁准确,不要多余话术。
【内部参考文档】
{context_text}
"""
    stream = client.chat.completions.create(
        model=os.getenv("ARK_MODEL_ID"),
        messages=[
            {"role": "system", "content": sys_prompt},
            {"role": "user", "content": user_question}
        ],
        temperature=float(os.getenv("TEMPERATURE")),
        stream=True
    )
    for chunk in stream:
        if chunk.choices[0].delta.content:
            yield chunk.choices[0].delta.content
    # 流结束后发送一个特殊标记,携带引用信息
    yield f"\n\n<references>{refs}</references>"

五、主服务重构(main_service.py)

整合登录、文件上传、流式问答、历史记录、鉴权拦截全功能,保留原有钉钉回调接口:

# main_service.py
__import__('pysqlite3')
import sys
sys.modules['sqlite3'] = sys.modules.pop('pysqlite3')

import os
import shutil
import threading
import json
from fastapi import FastAPI, Request, HTTPException, Depends, File, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel
from typing import List
import uvicorn

from doubao_rag_core import get_answer, get_answer_stream
from ding_bot import get_rag_answer, send_ding_msg
from auth import create_access_token, verify_user, get_current_user
from database import init_db, add_history, get_history, clear_history

init_db()  # 初始化历史记录表

app = FastAPI(title="豆包私有知识库API (增强版)")

# 跨域配置
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)

class LoginRequest(BaseModel):
    username: str
    password: str

class QuestionRequest(BaseModel):
    question: str

# 登录接口(无需鉴权)
@app.post("/login")
def login(req: LoginRequest):
    if not verify_user(req.username, req.password):
        raise HTTPException(status_code=401, detail="Incorrect username or password")
    access_token = create_access_token(data={"sub": req.username})
    return {"access_token": access_token, "token_type": "bearer"}

# 文件上传接口(需登录鉴权)
@app.post("/upload")
async def upload_file(file: UploadFile = File(...), user: str = Depends(get_current_user)):
    if not file.filename:
        raise HTTPException(status_code=400, detail="No file selected")
    # 限制合法文档格式
    allowed_ext = (".pdf", ".docx", ".txt", ".md")
    if not file.filename.lower().endswith(allowed_ext):
        raise HTTPException(status_code=400, detail="File type not allowed")
    save_path = os.path.join("docs", file.filename)
    with open(save_path, "wb") as buffer:
        shutil.copyfileobj(file.file, buffer)

    # 后台异步入库,不阻塞前端
    def run_ingest():
        from kb_loader import build_vector_db
        build_vector_db()
    threading.Thread(target=run_ingest, daemon=True).start()

    return {"message": f"文件 {file.filename} 已上传,后台入库中..."}

# 问答历史记录接口
@app.get("/history")
def read_history(user: str = Depends(get_current_user)):
    return get_history(user)

@app.delete("/history")
def delete_history(user: str = Depends(get_current_user)):
    clear_history(user)
    return {"message": "History cleared"}

# 原有普通问答接口(钉钉专用,无需鉴权)
@app.post("/api/qa")
async def api_qa(q: str):
    ans, refs = get_answer(q)
    return {"code": 200, "question": q, "answer": ans, "references": refs}

# 增强流式问答接口(需登录,自动保存历史)
@app.post("/api/stream/qa")
async def stream_qa(req: QuestionRequest, user: str = Depends(get_current_user)):
    question = req.question.strip()
    if not question:
        raise HTTPException(status_code=400, detail="Question cannot be empty")

    def generate():
        full_answer = ""
        refs_str = ""
        for chunk in get_answer_stream(question):
            if chunk.startswith("\n\n<references>") and chunk.endswith("</references>"):
                refs_str = chunk.replace("\n\n<references>", "").replace("</references>", "")
                try:
                    import ast
                    refs = ast.literal_eval(refs_str)
                    ref_text = "\n\n**参考来源:**\n" + "\n".join(
                        [f"- {r['source']}" + (f" 第{r['page']}页" if r['page'] else "") for r in refs]
                    )
                    full_answer += ref_text
                    yield ref_text
                except:
                    pass
            else:
                full_answer += chunk
                yield chunk
        # 保存问答历史
        add_history(user, question, full_answer)
    return StreamingResponse(generate(), media_type="text/plain")

# 钉钉机器人回调接口(保持不变)
@app.post("/ding/callback")
async def ding_callback(req: Request):
    data = await req.json()
    msg_type = data.get("Msgtype")
    if msg_type != "text":
        return {"msg": "ignore"}
    question = data["text"]["content"].strip()
    if not question:
        return {"msg": "empty"}
    ans, _ = get_answer(question)
    md_text = f"""
**用户提问:**
{question}

**知识库回复:**
{ans}
"""
    send_ding_msg(md_text)
    return {"msg": "success"}

# 前端首页入口
@app.get("/")
async def index():
    return FileResponse("index.html")

if __name__ == "__main__":
    uvicorn.run(app, host="0.0.0.0", port=8000)

六、前端页面(index.html)

全新重构深色主题前端页面,集成登录、上传、流式对话、历史记录、溯源展示全功能:

<!DOCTYPE html>
<html lang="zh-CN">
<head>
  <meta charset="UTF-8">
  <meta name="viewport" content="width=device-width, initial-scale=1.0">
  <title>企业私有知识库 · 字节豆包RAG</title>
  <script src="https://cdn.tailwindcss.com"></script>
  <link href="https://cdn.jsdelivr.net/npm/font-awesome@4.7.0/css/font-awesome.min.css" rel="stylesheet">
  <script>
    tailwind.config = {
      theme: {
        extend: {
          colors: {
            primary: '#165DFF',
            darkBg: '#17171a',
            sideBg: '#232329',
            chatBg: '#2c2c34',
            textGray: '#a1a1aa'
          }
        }
      }
    }
  </script>
  <style>
    .message-box::-webkit-scrollbar { width: 4px }
    .message-box::-webkit-scrollbar-thumb { background: #444; border-radius: 4px }
    .typing-dot { animation: blink 1.4s infinite both; }
    @keyframes blink { 0%,80%,100%{opacity:0.3} 40%{opacity:1} }
  </style>
</head>
<body class="bg-darkBg text-white h-screen flex overflow-hidden">

  <!-- 登录浮层 -->
  <div id="loginOverlay" class="absolute inset-0 bg-black bg-opacity-70 flex items-center justify-center z-50">
    <div class="bg-sideBg p-8 rounded-lg w-80">
      <h2 class="text-xl mb-4">登录知识库</h2>
      <input id="username" type="text" placeholder="账号" class="w-full p-2 mb-2 bg-chatBg rounded border border-gray-600 focus:border-primary outline-none">
      <input id="password" type="password" placeholder="密码" class="w-full p-2 mb-4 bg-chatBg rounded border border-gray-600 focus:border-primary outline-none">
      <button onclick="doLogin()" class="w-full bg-primary hover:bg-primary/80 text-white py-2 rounded">登 录</button>
      <p id="loginError" class="text-red-400 text-sm mt-2 hidden">账号或密码错误</p>
    </div>
  </div>

  <!-- 主界面(登录后可见) -->
  <div id="mainApp" class="hidden flex w-full h-full">
    <!-- 侧边栏 -->
    <aside class="w-60 bg-sideBg flex flex-col shrink-0">
      <div class="px-4 py-5 border-b border-gray-700">
        <div class="flex items-center gap-2">
          <div class="w-8 h-8 rounded-md bg-primary flex items-center justify-center"><i class="fa fa-book text-white"></i></div>
          <span class="font-bold text-lg">私有知识库助手</span>
        </div>
        <p class="text-textGray text-xs mt-1">基于字节豆包 · 本地RAG</p>
      </div>
      <div class="flex-1 p-3 overflow-auto">
        <div class="flex justify-between items-center mb-2">
          <span class="text-textGray text-xs">对话记录</span>
          <button onclick="clearHistory()" class="text-xs text-red-400 hover:underline">清空</button>
        </div>
        <div id="historyList" class="space-y-1"></div>
      </div>
      <div class="p-3 border-t border-gray-700">
        <span class="text-textGray text-xs">系统状态</span>
        <div class="text-sm px-2 py-2 rounded bg-chatBg text-textGray">
          <div class="flex items-center gap-2"><i class="fa fa-check text-green-400"></i>向量库正常</div>
          <div class="flex items-center gap-2 mt-1"><i class="fa fa-check text-green-400"></i>豆包API已连通</div>
        </div>
      </div>
    </aside>

    <!-- 主对话区 -->
    <main class="flex-1 flex flex-col">
      <header class="h-14 border-b border-gray-700 flex items-center px-6 justify-between">
        <h2 class="font-medium">智能问答对话</h2>
        <div class="flex gap-3">
          <!-- 上传文件按钮 -->
          <label class="px-3 py-1 rounded bg-chatBg hover:bg-gray-600 text-sm cursor-pointer">
            <i class="fa fa-upload mr-1"></i>上传文档
            <input type="file" id="fileUpload" accept=".pdf,.docx,.txt,.md" class="hidden" onchange="uploadFile()">
          </label>
          <button onclick="clearChat()" class="px-3 py-1 rounded bg-chatBg hover:bg-gray-600 text-sm">
            <i class="fa fa-trash-o mr-1"></i>清空对话
          </button>
        </div>
      </header>

      <!-- 聊天容器 -->
      <div id="chatBox" class="message-box flex-1 overflow-y-auto p-6 space-y-6 bg-[#1c1c22]">
        <div class="flex gap-3">
          <div class="w-9 h-9 rounded-md bg-primary shrink-0 flex items-center justify-center"><i class="fa fa-robot"></i></div>
          <div class="bg-chatBg rounded-lg px-4 py-3 max-w-[75%]">
            <p>你好,我是本地知识库问答助手。<br>已加载本地文档资料,你的提问只会检索内部文档资料进行回答。</p>
          </div>
        </div>
      </div>

      <!-- 输入区 -->
      <div class="p-4 border-t border-gray-700">
        <div class="flex items-end gap-3 max-w-[90%] mx-auto">
          <textarea id="userInput" rows="3" placeholder="输入你的问题,回车发送..."
            class="flex-1 bg-chatBg rounded-lg px-4 py-3 outline-none resize-none border border-gray-600 focus:border-primary"></textarea>
          <button id="sendBtn" onclick="sendQuestion()" class="bg-primary hover:bg-primary/80 rounded-lg h-[74px] w-12 flex items-center justify-center">
            <i class="fa fa-paper-plane"></i>
          </button>
        </div>
        <div class="text-textGray text-xs text-center mt-2">数据全部本地向量存储,不会上传你的文档</div>
      </div>
    </main>
  </div>

  <script>
    // 全局变量
    let token = localStorage.getItem('token') || '';
    const chatBox = document.getElementById('chatBox');
    const userInput = document.getElementById('userInput');
    const sendBtn = document.getElementById('sendBtn');
    let loading = false;

    // 登录逻辑
    async function doLogin() {
      const username = document.getElementById('username').value;
      const password = document.getElementById('password').value;
      try {
        const res = await fetch('/login', {
          method: 'POST',
          headers: { 'Content-Type': 'application/json' },
          body: JSON.stringify({ username, password })
        });
        if (!res.ok) throw new Error('Login failed');
        const data = await res.json();
        token = data.access_token;
        localStorage.setItem('token', token);
        document.getElementById('loginOverlay').classList.add('hidden');
        document.getElementById('mainApp').classList.remove('hidden');
        loadHistory();
      } catch (e) {
        document.getElementById('loginError').classList.remove('hidden');
      }
    }

    // 初始化登录状态校验
    if (token) {
      fetch('/history', { headers: { 'Authorization': `Bearer ${token}` } })
        .then(res => {
          if (res.ok) {
            document.getElementById('loginOverlay').classList.add('hidden');
            document.getElementById('mainApp').classList.remove('hidden');
            loadHistory();
          } else {
            localStorage.removeItem('token');
            token = '';
          }
        });
    }

    function authHeaders() {
      return { 'Authorization': `Bearer ${token}`, 'Content-Type': 'application/json' };
    }

    // 文件上传
    async function uploadFile() {
      const fileInput = document.getElementById('fileUpload');
      if (!fileInput.files.length) return;
      const formData = new FormData();
      formData.append('file', fileInput.files[0]);
      try {
        const res = await fetch('/upload', {
          method: 'POST',
          headers: { 'Authorization': `Bearer ${token}` },
          body: formData
        });
        const data = await res.json();
        alert(data.message);
        fileInput.value = '';
      } catch (err) {
        alert('上传失败');
        console.error(err);
      }
    }

    // 加载历史记录
    async function loadHistory() {
      try {
        const res = await fetch('/history', { headers: authHeaders() });
        const history = await res.json();
        const list = document.getElementById('historyList');
        list.innerHTML = history.map((item, idx) => `
          <div class="text-sm px-2 py-1 rounded hover:bg-gray-700 cursor-pointer truncate" onclick="loadHistoryItem('${item.question.replace(/'/g, "\\'")}')">
            ${item.question.length > 20 ? item.question.slice(0,20)+'...' : item.question}
          </div>
        `).join('');
      } catch (err) {
        console.error('加载历史失败', err);
      }
    }

    function loadHistoryItem(question) {
      userInput.value = question;
    }

    async function clearHistory() {
      if (!confirm('确定清空所有对话历史吗?')) return;
      await fetch('/history', { method: 'DELETE', headers: authHeaders() });
      loadHistory();
    }

    // 流式问答发送
    async function sendQuestion() {
      const question = userInput.value.trim();
      if (!question || loading) return;
      userInput.value = '';
      loading = true;
      sendBtn.disabled = true;

      // 渲染用户消息
      chatBox.innerHTML += `
        <div class="flex gap-3 justify-end">
          <div class="bg-primary/80 rounded-lg px-4 py-3 max-w-[75%]">${question}</div>
          <div class="w-9 h-9 rounded-md bg-gray-500 shrink-0 flex items-center justify-center"><i class="fa fa-user"></i></div>
        </div>
      `;

      // 初始化AI回复容器
      const respWrap = document.createElement('div');
      respWrap.className = "flex gap-3";
      respWrap.innerHTML = `
        <div class="w-9 h-9 rounded-md bg-primary shrink-0 flex items-center justify-center"><i class="fa fa-robot"></i></div>
        <div class="bg-chatBg rounded-lg px-4 py-3 max-w-[75%]"><span class="ai-text"></span></div>
      `;
      chatBox.appendChild(respWrap);
      const aiTextSpan = respWrap.querySelector('.ai-text');
      chatBox.scrollTop = chatBox.scrollHeight;

      try {
        const response = await fetch('/api/stream/qa', {
          method: 'POST',
          headers: authHeaders(),
          body: JSON.stringify({ question })
        });
        const reader = response.body.getReader();
        const decoder = new TextDecoder();
        let fullText = '';

        while (true) {
          const { done, value } = await reader.read();
          if (done) break;
          const chunk = decoder.decode(value, { stream: true });
          fullText += chunk;
          aiTextSpan.innerHTML = fullText.replace(/\n/g, '<br>');
          chatBox.scrollTop = chatBox.scrollHeight;
        }
      } catch (err) {
        aiTextSpan.innerHTML = '请求异常,请检查配置';
        console.error(err);
      }
      loading = false;
      sendBtn.disabled = false;
      setTimeout(loadHistory, 500);
    }

    // 清空对话
    function clearChat() {
      chatBox.innerHTML = `
        <div class="flex gap-3">
          <div class="w-9 h-9 rounded-md bg-primary shrink-0 flex items-center justify-center"><i class="fa fa-robot"></i></div>
          <div class="bg-chatBg rounded-lg px-4 py-3 max-w-[75%]">
            <p>对话已清空,请重新提问</p>
          </div>
        </div>
      `;
    }

    // 回车发送快捷键
    userInput.addEventListener("keydown", (e) => {
      if (e.key === "Enter" && !e.shiftKey) {
        e.preventDefault();
        sendQuestion();
      }
    });
  </script>
</body>
</html>

七、项目部署步骤

  1. 安装新增依赖:在项目根目录执行依赖安装命令

  2. 修改配置文件:编辑 .env,自定义登录账号、密码、JWT密钥

  3. 检查目录:确保项目根目录存在 docs/ 文档存储文件夹

  4. 启动服务:执行启动命令运行项目

  5. 访问使用:浏览器打开 http://服务器IP:8000,登录后即可使用全部功能

# 启动项目
python3 main_service.py

八、核心功能验证清单

  • 账号鉴权:未登录状态无法访问任何业务接口,前端强制拦截

  • 文档上传:支持 PDF、DOCX、TXT、MD 格式,后台自动异步入库向量化

  • 流式输出:问答实现打字机实时逐字输出,支持换行格式化

  • 历史存储:本地SQLite持久化问答记录,支持查看、清空历史对话

  • 答案溯源:回答末尾自动展示参考文档名称、对应页码,有据可查

九、进阶升级:无缝适配AD域控/LDAP登录认证

原有账号密码体系为静态本地账号,支持无缝升级企业AD域控LDAP认证。本次升级改动量极小,仅替换用户校验逻辑、新增少量依赖与配置,JWT签发、前端页面、接口鉴权、历史记录等所有原有功能完全无需改动,同时保留本地账号兼容模式,可随时切换。

1. 安装AD认证依赖包

核心依赖 python-ldap,Linux系统若编译报错,需提前安装系统编译依赖:

# Python依赖安装
pip install python-ldap

# Rocky Linux 专属编译依赖(报错时执行)
dnf install openldap-devel python3-devel gcc -y

2. 扩容.env配置文件

注释/删除原有静态账号密码配置,新增LDAP域控配置,按需替换为企业实际域控信息,支持双模式切换:

# 原有鉴权配置 - 静态账号(注释停用,保留兼容)
# LOGIN_USER=admin
# LOGIN_PASS=admin123

# 新增AD/LDAP域控认证配置
AUTH_MODE=ldap
LDAP_SERVER=ldap://192.168.1.10
LDAP_PORT=389
LDAP_BASE_DN=dc=company,dc=com
LDAP_DOMAIN=company.com
# 可选:指定用户搜索OU,不填默认使用根域
LDAP_USER_SEARCH_BASE=ou=Users,dc=company,dc=com

3. 重写auth.py认证模块(核心改动)

仅重写用户验证逻辑,保留原有JWT签发、全局鉴权、令牌校验全部逻辑,同时兼容静态账号+LDAP域控双模式:

# auth.py
import os
import ldap
from datetime import datetime, timedelta
from typing import Optional
from jose import JWTError, jwt
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from dotenv import load_dotenv

load_dotenv()

SECRET_KEY = os.getenv("SECRET_KEY")
ALGORITHM = os.getenv("ALGORITHM")
ACCESS_TOKEN_EXPIRE_MINUTES = int(os.getenv("ACCESS_TOKEN_EXPIRE_MINUTES", 480))
AUTH_MODE = os.getenv("AUTH_MODE", "simple").lower()

oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/login")

def create_access_token(data: dict, expires_delta: Optional[timedelta] = None):
    to_encode = data.copy()
    expire = datetime.utcnow() + (expires_delta or timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES))
    to_encode.update({"exp": expire})
    return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)

# 原静态账号密码验证(兼容保留)
def _verify_simple(username: str, password: str):
    return username == os.getenv("LOGIN_USER") and password == os.getenv("LOGIN_PASS")

# AD/LDAP域控账号验证
def _verify_ldap(username: str, password: str):
    server = os.getenv("LDAP_SERVER")
    base_dn = os.getenv("LDAP_BASE_DN")
    domain = os.getenv("LDAP_DOMAIN")
    search_base = os.getenv("LDAP_USER_SEARCH_BASE", base_dn)

    # 拼接AD用户登录主体:用户名@域名
    user_dn = f"{username}@{domain}"

    try:
        conn = ldap.initialize(server)
        conn.set_option(ldap.OPT_REFERRALS, 0)  # 关闭引用追踪,避免连接异常
        conn.simple_bind_s(user_dn, password)   # 域控账号密码绑定认证
        conn.unbind_s()
        return True
    except ldap.INVALID_CREDENTIALS:
        # 账号密码错误
        return False
    except Exception as e:
        # 域控连接异常、配置错误等
        print(f"LDAP连接异常:{e}")
        return False

# 统一认证入口,自动适配双模式
def verify_user(username: str, password: str):
    if AUTH_MODE == "ldap":
        return _verify_ldap(username, password)
    else:
        return _verify_simple(username, password)

# 全局接口鉴权逻辑(完全保留、无需修改)
async def get_current_user(token: str = Depends(oauth2_scheme)):
    credentials_exception = HTTPException(
        status_code=status.HTTP_401_UNAUTHORIZED,
        detail="Could not validate credentials",
        headers={"WWW-Authenticate": "Bearer"},
    )
    try:
        payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
        username: str = payload.get("sub")
        if not username:
            raise credentials_exception
    except JWTError:
        raise credentials_exception
    return username

4. 认证模式切换规则

  • 域控认证模式.env 配置AUTH_MODE=ldap,读取LDAP配置,使用企业AD账号登录

  • 静态账号模式:不配置AUTH_MODE 或设为 simple,沿用原有 LOGIN_USER/LOGIN_PASS,适配测试环境快速切换

5. 核心升级优势说明

  • 零前端改动:保留原有登录页面,用户直接输入企业AD账号密码即可登录,前端体验完全不变

  • 高性能鉴权:仅登录时请求一次域控,认证成功后签发JWT令牌,后续所有接口鉴权均基于本地令牌,不重复访问域控,无性能损耗

  • 高可扩展性:如需限制指定部门、指定AD用户组访问,可在 _verify_ldap 函数中新增 memberOf 组过滤逻辑,仅需少量代码改造

  • 双向兼容:新旧认证模式无缝切换,生产、测试环境可灵活适配

6. 接口测试验证

重启项目后,通过curl命令测试AD域账号登录,返回 access_token 即代表认证成功:

# 替换为你的AD真实账号密码
curl -X POST http://localhost:8000/login \
  -H "Content-Type: application/json" \
  -d '{"username":"testuser","password":"Test1234"}'

7. 升级总结

本次AD域控适配升级,仅新增依赖、修改配置、替换核心校验函数,项目原有RAG问答、文件上传、流式输出、历史记录、答案溯源等全部功能完全兼容,快速实现企业私有化域账号统一登录,适配企业内部权限管控规范。

最终效果呈现

Snipaste_2026-08-05_10-57-22.png

Rocky Linux 从零搭建私有RAG知识库|FastAPI+豆包方舟+现代化Web界面(小白零门槛完整版) 2026-07-30
私有RAG知识库踩坑实录|对话记录点击无法还原问答内容问题排查与完整修复 2026-07-31