Procházet zdrojové kódy

首页知识库对话-rg

zhaoqingang před 1 rokem
rodič
revize
282a631b9c

+ 1 - 1
app/api/v2/chat.py

@@ -42,7 +42,7 @@ async def api_chat_dialog(chatId:str, dialog: ChatData, current_user: UserModel
             return StreamingResponse(f"data: {error_msg}\n\n",
                                      media_type="text/event-stream")
         session_id = session.get("data", {}).get("id")
-    return StreamingResponse(service_chat_dialog(db, chatId, dialog.query, session_id, current_user.id, chat_info.mode),
+    return StreamingResponse(service_chat_dialog(db, chatId, dialog.query, session_id, current_user.id, chat_info.mode, chat_info.get_kb_ids()),
                              media_type="text/event-stream")
 
 @chat_router_v2.post("/agent/{chatId}/completions")

+ 10 - 1
app/api/v2/mindmap.py

@@ -14,7 +14,7 @@ from app.models.base_model import get_db
 from app.models.v2.chat import RetrievalRequest, ComplexChatDao
 from app.models.v2.mindmap import MindmapRequest
 from app.models.v2.session_model import ChatData
-from app.service.v2.mindmap import service_chat_mindmap
+from app.service.v2.mindmap import service_chat_mindmap, service_message_mindmap_parse
 
 mind_map_router = APIRouter()
 
@@ -28,4 +28,13 @@ async def api_chat_mindmap(mindmap: MindmapRequest, current_user: UserModel = De
             return Response(code=500, msg="create failure", data={})
     else:
         return Response(code=500, msg="网络异常!failure", data={})
+    return Response(code=200, msg="create success", data=data)
+
+
+@mind_map_router.get("/{messageId}/parse", response_model=Response)
+async def api_chat_mindmap(messageId: str, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): #  current_user: UserModel = Depends(get_current_user)
+
+    data = await service_message_mindmap_parse(db, messageId, current_user.id)
+    if not data:
+        return Response(code=500, msg="create failure", data={})
     return Response(code=200, msg="create success", data=data)

+ 5 - 0
app/models/dialog_model.py

@@ -1,3 +1,4 @@
+import json
 from datetime import datetime
 from typing import Optional
 
@@ -24,6 +25,7 @@ class DialogModel(Base):
     # agent_id = Column(String(36))
     mode = Column(String(36))
     parameters = Column(Text)
+    kb_ids = Column(String(128))
 
     def get_id(self):
         return str(self.id)
@@ -43,6 +45,9 @@ class DialogModel(Base):
             'mode': self.mode,
         }
 
+    def get_kb_ids(self):
+        return json.loads(self.kb_ids) if self.kb_ids else []
+
 
 class ConversationModel(Base):
     __tablename__ = 'conversation'

+ 10 - 4
app/models/v2/chat.py

@@ -6,7 +6,7 @@ from pydantic import BaseModel
 from sqlalchemy import Column, Integer, String, BigInteger, ForeignKey, DateTime, Text, TEXT
 from sqlalchemy.orm import Session
 
-from app.config.const import Dialog_STATSU_DELETE, Dialog_STATSU_ON
+from app.config.const import Dialog_STATSU_DELETE, Dialog_STATSU_ON, complex_knowledge_chat
 from app.models.base_model import Base
 from app.utils.common import current_time
 
@@ -187,14 +187,20 @@ class ComplexChatSessionModel(Base):
             query = {}
             if self.query:
                 query = json.loads(self.query)
-            return {
+
+            res = {
                 'id': self.id,
                 'role': "assistant",
                 'answer': self.content,
                 'chat_mode': self.chat_mode,
-                'node_list': json.loads(self.node_data) if self.node_data else [],
-                "parentId": query.get("parentId")
+                "parentId": query.get("parentId"),
+                "isDeep": query.get("isDeep", 1),
             }
+            if self.chat_mode == complex_knowledge_chat:
+                res['reference'] = json.loads(self.node_data) if self.node_data else {}
+            else:
+                res['node_list'] = json.loads(self.node_data) if self.node_data else []
+            return res
 
 
 class ComplexChatSessionDao:

+ 1 - 0
app/service/dialog.py

@@ -245,6 +245,7 @@ async def sync_dialog_service(db, dialog_id):
             if app_dialog:
                 dialog.name = app_dialog["name"]
                 dialog.description = app_dialog["description"]
+                dialog.kb_ids = app_dialog["kb_ids"]
                 dialog.update_date = datetime.now()
                 db.add(dialog)
                 db.commit()

+ 2 - 2
app/service/knowledge.py

@@ -17,8 +17,8 @@ async def get_knowledge_list(db, user_id, keyword, page_size, page_index, status
         klg_list = [j.id for i in user.groups for j in i.knowledges]
         query = query.filter(or_(KnowledgeModel.id.in_(klg_list), KnowledgeModel.tenant_id == str(user_id)))
 
-    if location:
-        query = query.filter(or_(KnowledgeModel.permission == "team", KnowledgeModel.tenant_id == str(user_id)))
+        if location:
+            query = query.filter(or_(KnowledgeModel.permission == "team", KnowledgeModel.tenant_id == str(user_id)))
 
     if keyword:
         query = query.filter(KnowledgeModel.name.like('%{}%'.format(keyword)))

+ 32 - 3
app/service/v2/chat.py

@@ -6,6 +6,7 @@ import uuid
 
 import fitz
 from fastapi import HTTPException
+from sqlalchemy import or_
 
 from Log import logger
 from app.config.agent_base_url import RG_CHAT_DIALOG, DF_CHAT_AGENT, DF_CHAT_PARAMETERS, RG_CHAT_SESSIONS, \
@@ -13,7 +14,7 @@ from app.config.agent_base_url import RG_CHAT_DIALOG, DF_CHAT_AGENT, DF_CHAT_PAR
 from app.config.config import settings
 from app.config.const import *
 from app.models import DialogModel, ApiTokenModel, UserTokenModel, ComplexChatSessionDao, ChatDataRequest, \
-    ComplexChatDao
+    ComplexChatDao, KnowledgeModel, UserModel
 from app.models.v2.session_model import ChatSessionDao, ChatData
 from app.service.v2.app_driver.chat_agent import ChatAgent
 from app.service.v2.app_driver.chat_data import ChatBaseApply
@@ -87,17 +88,45 @@ async def get_chat_object(mode):
         return ChatAgent(), url
 
 
-async def service_chat_dialog(db, chat_id: str, question: str, session_id: str, user_id, mode: str):
+
+async def get_user_kb(db, user_id: int, kb_ids: list) -> list:
+    res = []
+    user = db.query(UserModel).filter(UserModel.id == user_id).first()
+    if user is None:
+        return res
+    query = db.query(KnowledgeModel)
+    if user.permission != "admin":
+        klg_list = [j.id for i in user.groups for j in i.knowledges]
+        query = query.filter(or_(KnowledgeModel.id.in_(klg_list), KnowledgeModel.tenant_id == str(user_id)))
+        kb_list= query.all()
+        for kb in kb_list:
+            if kb.id in kb_ids:
+                if kb.permission == "team":
+                    res.append(kb.id)
+                elif kb.tenant_id == str(user_id):
+                    res.append(kb.id)
+        return res
+    else:
+        return kb_ids
+
+
+async def service_chat_dialog(db, chat_id: str, question: str, session_id: str, user_id: int, mode: str, kb_ids: list):
     conversation_id = ""
     token = await get_chat_token(db, rg_api_token)
     url = settings.fwr_base_url + RG_CHAT_DIALOG.format(chat_id)
+    kb_id = await get_user_kb(db, user_id, kb_ids)
+    if not kb_id:
+        yield "data: " + json.dumps({"message": smart_message_error,
+                                     "error": "\n**ERROR**: The agent has no knowledge base to work with!", "status": http_400},
+                                    ensure_ascii=False) + "\n\n"
+        return
     chat = ChatDialog()
     session = await add_session_log(db, session_id, question, chat_id, user_id, mode, session_id, RG_TYPE)
     if session:
         conversation_id = session.conversation_id
     message = {"role": "assistant", "answer": "", "reference": {}}
     try:
-        async for ans in chat.chat_completions(url, await chat.request_data(question, conversation_id),
+        async for ans in chat.chat_completions(url, await chat.complex_request_data(question, kb_id, conversation_id),
                                                await chat.get_headers(token)):
             data = {}
             error = ""

Rozdílová data souboru nebyla zobrazena, protože soubor je příliš velký
+ 52 - 19
app/service/v2/mindmap.py


+ 7 - 5
app/task/fetch_agent.py

@@ -43,6 +43,7 @@ class Dialog(Base):
     status = Column(String(1), nullable=False)
     description = Column(String(255), nullable=False)
     tenant_id = Column(String(36), nullable=False)
+    kb_ids = Column(String(128), nullable=False)
 
 
 class DfApps(Base):
@@ -257,13 +258,13 @@ def get_data_from_ragflow_v2(base_db, names: List[str], tenant_id) -> List[Dict]
             query = db.query(Dialog.id, Dialog.name, Dialog.description, Dialog.status, Dialog.tenant_id) \
                 .filter(Dialog.name.in_(names), Dialog.status == "1")
         else:
-            query = db.query(Dialog.id, Dialog.name, Dialog.description, Dialog.status, Dialog.tenant_id).filter(
+            query = db.query(Dialog.id, Dialog.name, Dialog.description, Dialog.status, Dialog.tenant_id, Dialog.kb_ids).filter(
                 Dialog.status == "1", Dialog.tenant_id == tenant_id)
 
         results = query.all()
         formatted_results = [
             {"id": row[0], "name": row[1], "description": row[2], "status": "1" if row[3] == "1" else "2",
-             "user_id": str(row[4]), "mode": "agent-dialog", "parameters": para} for row in results if row[0] not in chat_ids]
+             "user_id": str(row[4]), "mode": "agent-dialog", "parameters": para, "kb_ids": row[5]} for row in results if row[0] not in chat_ids]
         return formatted_results
     finally:
         db.close()
@@ -301,13 +302,14 @@ def update_ids_in_local_v2(data: List[Dict], dialog_type: str):
                 existing_agent.name = row["name"]
                 existing_agent.description = row["description"]
                 existing_agent.mode = row["mode"]
+                existing_agent.kb_ids = row.get("kb_ids", "")
                 if existing_agent.status == Dialog_STATSU_DELETE:
                     existing_agent.status = Dialog_STATSU_ON
                 if row["parameters"]:
                     existing_agent.parameters = json.dumps(row["parameters"])
             else:
                 existing = DialogModel(id=row["id"], status=row["status"], name=row["name"],
-                                       description=row["description"],
+                                       description=row["description"], kb_ids=row.get("kb_ids", ""),
                                        tenant_id=get_rag_user_id(db, row["user_id"], type_dict[dialog_type]),
                                        dialog_type=dialog_type, mode=row["mode"], parameters=json.dumps(row["parameters"]))
                 db.add(existing)
@@ -411,10 +413,10 @@ def get_one_from_ragflow_knowledge(klg_id):
 def get_one_from_ragflow_dialog(dialog_id):
     db = SessionRagflow()
     try:
-        row = db.query(Dialog.id, Dialog.name, Dialog.description, Dialog.status, Dialog.tenant_id) \
+        row = db.query(Dialog.id, Dialog.name, Dialog.description, Dialog.status, Dialog.tenant_id, Dialog.kb_ids) \
             .filter(Dialog.id==dialog_id).first()
         return {"id": row[0], "name": row[1], "description": row[2], "status": str(row[3]),
-                "user_id": str(row[4])} if row else {}
+                "user_id": str(row[4]), "kb_ids": row[5]} if row else {}
     finally:
         db.close()
 

Některé soubory nejsou zobrazeny, neboť je v těchto rozdílových datech změněno mnoho souborů