Team Ai
Apppublic

githubear/oceanbase_chat_pdf

sourceHugging Facemitupdated 3y agoView on Hugging Face
1likes
app.py223 linesDownload Raw Back to root
1import os2import requests3import bcrypt4import pymysql5import streamlit as st6# vectordb7from langchain.vectorstores import OceanBase8# embeddings9# from langchain_openai import OpenAIEmbeddings10from langchain_community.embeddings import JinaEmbeddings11# PDF12from PyPDF2 import PdfReader13from langchain.text_splitter import RecursiveCharacterTextSplitter14# LLM15# from langchain_openai import ChatOpenAI16from langchain_community.llms import Tongyi17from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder18from langchain_core.output_parsers import StrOutputParser19from langchain_core.messages import HumanMessage20from langchain_core.runnables import RunnablePassthrough21from langchain.chains.combine_documents import create_stuff_documents_chain22# env23from dotenv import load_dotenv24 25load_dotenv()26if 'login_status' not in st.session_state:27    st.session_state['login_status'] = False28if 'username' not in st.session_state:29    st.session_state['username'] = ''30if 'user_id' not in st.session_state:31    st.session_state['user_id'] = -132 33# 配置数据库连接34def create_db_connection():35    connection = None36    try:37        connection = pymysql.connect(38            host=os.getenv("OB_HOST", "localhost"),39            port=int(os.getenv("OB_PORT", 2881)),40            db=os.getenv("OB_DATABASE", "test"),41            user=os.getenv("OB_USER", "root"),42            passwd=os.getenv("OB_PASSWORD", ""),43            charset='utf8mb4',44            cursorclass=pymysql.cursors.DictCursor45        )46    except pymysql.MySQLError as e:47        st.error(f"The error '{e}' occurred")48    return connection49 50# 用户注册功能51def register_user(connection, username, password):52    with connection.cursor() as cursor:53        try:54            password_hash = bcrypt.hashpw(password.encode('utf-8'), bcrypt.gensalt())55            insert_query = """56                INSERT INTO chat_users (username, password)57                VALUES (%s, %s)58            """59            cursor.execute(insert_query, (username, password_hash))60            connection.commit()61            st.success("User registered successfully.")62        except pymysql.MySQLError as e:63            st.error(f"The error '{e}' occurred")64 65# 用户登录功能66def login_user(connection, username, password):67    with connection.cursor() as cursor:68        cursor.execute("SELECT user_id,password FROM chat_users WHERE username = %s", (username,))69        record = cursor.fetchone()70        if record and bcrypt.checkpw(password.encode('utf-8'), record['password'].encode('utf-8')):71            return record['user_id']72        return -173 74def get_pdf_text(pdf_docs):75    text = ""76    for pdf in pdf_docs:77        pdf_reader = PdfReader(pdf)78        for page in pdf_reader.pages:79            text += page.extract_text()80    return text81 82def load_text_chunks(pdf_docs):83    text = get_pdf_text(pdf_docs)84    text_splitter = RecursiveCharacterTextSplitter(chunk_size=1024, chunk_overlap=0)85    return text_splitter.split_text(text)86 87## create vectore store 88def get_oceanbase() -> OceanBase:89    connection_str = OceanBase.connection_string_from_db_params(90        host=os.getenv("OB_HOST", "localhost"),91        port=os.getenv("OB_PORT", "2881"),92        database=os.getenv("OB_DATABASE", "test"),93        user=os.getenv("OB_USER", "root"),94        password=os.getenv("OB_PASSWORD", ""),95    )96    # Embeddings: can be changed97    # embeddings = OpenAIEmbeddings(api_key=os.getenv("OPENAI_API_KEY", ""), openai_proxy='https://api.chatgptid.net/v1')98    embeddings = JinaEmbeddings(99        jina_api_key=os.getenv("JINA_AI_API", ""), model_name="jina-embeddings-v2-base-zh"100    )101    # create oceanbase102    collection_name = f"langchain_document{st.session_state['user_id']}"103    oceanbase = OceanBase(104        connection_string=connection_str,105        embedding_function=embeddings,106        collection_name=collection_name,107        # pre_delete_collection=True,  # TODO108    )109    return oceanbase110 111def get_texts_summary(texts):112    prompt_text = """您是一名助理,负责总结文本以供检索。 \113    这些摘要将用于embedding并用于检索原始文本。 \114    现在请给出针对检索进行优化的简洁摘要。 以下是原始文本: {element} """115    prompt = ChatPromptTemplate.from_template(prompt_text)116    117    llm = Tongyi()118    # llm = ChatOpenAI(api_key=os.getenv("OPENAI_API_KEY", ""), model="gpt-3.5-turbo-1106", openai_proxy='https://api.chatgptid.net/v1')119    summarize_chain = {"element": lambda x: x} | prompt | llm | StrOutputParser()120    # print(texts)121    return summarize_chain.batch(texts, {"max_concurrency": 5})122 123def text_rag_chain(retriever):124    llm = Tongyi()125    # llm = ChatOpenAI(api_key=os.getenv("OPENAI_API_KEY", ""), model="gpt-3.5-turbo-1106", openai_proxy='https://api.chatgptid.net/v1')126    SYSTEM_TEMPLATE = """127    根据下面给出的上下文回答用户的问题。128    如果下面的上下文中不包含与问题相关的任何信息,请不要编造内容,仅仅回复”我不知道“129 130    <context>131    {context}132    </context>133    """134    question_answering_prompt = ChatPromptTemplate.from_messages(135        [136            (137                "system",138                SYSTEM_TEMPLATE,139            ),140            MessagesPlaceholder(variable_name="messages"),141        ]142    )143    document_chain = create_stuff_documents_chain(llm, question_answering_prompt)144 145    def parse_retriever_input(params):146        return params["messages"][-1].content147    retrieval_chain = RunnablePassthrough.assign(148        context=parse_retriever_input | retriever,149    ).assign(150        answer=document_chain,151    )152    return retrieval_chain153 154def run_pipeline(oceanbase, user_question):155    retriever = oceanbase.as_retriever(k=5)156    chain = text_rag_chain(retriever)157    response = chain.invoke(158        {159            "messages": [160                HumanMessage(content=user_question)161            ],162        }163    )164    st.write(response["answer"])165 166def main():167    if not st.session_state['login_status']:    168        # Streamlit 界面169        st.title("Login to chat with OceanBase")170        tab = st.radio("Choose a tab:", ["Login", "Register"])171 172        # 注册界面173        if tab == "Login":174            login_username = st.text_input("Username", key="login_username")175            login_password = st.text_input("Password", type="password", key="login_password")176            login_button = st.button("Login")177            178            if login_button:179                conn = create_db_connection()180                st.session_state['user_id'] = login_user(conn, login_username, login_password)181                if conn and st.session_state['user_id'] != -1:182                    st.session_state['login_status'] = True183                    st.session_state['username'] = login_username184                    st.success(f"Welcome {login_username}!")185                    conn.close()186                elif conn:187                    st.error("Incorrect username/password")188                    conn.close()189        elif tab == "Register":190            new_username = st.text_input("Username", key="register_username")191            new_password = st.text_input("Password", type="password", key="register_password")192            register_button = st.button("Register")193            194            if register_button:195                conn = create_db_connection()196                if conn:197                    register_user(conn, new_username, new_password)198                    conn.close()199    200    elif st.session_state['login_status']:201        oceanbase = get_oceanbase()202        st.set_page_config("Chat PDF")203        st.header("Chat with PDF")204 205        user_question = st.text_input("Ask a Question from the PDF Files")206 207        if user_question:208            run_pipeline(oceanbase, user_question)209 210        with st.sidebar:211            st.title("Menu:")212            pdf_docs = st.file_uploader("Upload PDF Files", accept_multiple_files=True)213            if st.button("Submit & Process"):214                with st.spinner("Processing..."):215                    if pdf_docs:216                        texts = load_text_chunks(pdf_docs)217                        # summary = get_texts_summary(texts)218                        # texts.append(summary)219                        oceanbase.add_texts(texts=texts)220                        st.success("Done")221 222if __name__ == "__main__":223    main()