File size: 5,120 Bytes
04355ce
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7148f99
04355ce
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
import os
import streamlit as st
from langchain.chains import create_history_aware_retriever, create_retrieval_chain
from langchain.chains.combine_documents import create_stuff_documents_chain
from langchain_chroma import Chroma
from langchain_community.chat_message_histories import ChatMessageHistory
from langchain_core.chat_history import BaseChatMessageHistory
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_groq import ChatGroq
from langchain_core.runnables.history import RunnableWithMessageHistory
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_community.document_loaders import PyPDFLoader, WebBaseLoader
from chromadb.config import Settings

# Load secrets from environment (set in Hugging Face Space Settings > Secrets)
groq_api_key = os.environ.get("GROQ_API_KEY")
hf_token = os.environ.get("HUGGINGFACE_API_KEY")

# Initialize LLM and embeddings
embeddings = HuggingFaceEmbeddings(model_name="all-MiniLM-L6-v2")
llm = ChatGroq(groq_api_key=groq_api_key, model_name="llama3-70b-8192")

# Streamlit UI
st.title("Conversational RAG | PDF + Website")
st.write("Upload PDFs or enter a website URL to chat with their content.")

session_id = st.text_input("Session ID:", value="default_session")

# Initialize session state
if 'store' not in st.session_state:
    st.session_state.store = {}
if 'vectorstore' not in st.session_state:
    st.session_state.vectorstore = None
if 'document_source' not in st.session_state:
    st.session_state.document_source = ""
if 'documents' not in st.session_state:
    st.session_state.documents = []

# Input Mode Selector
input_mode = st.radio("Select Input Type:", ("PDF Upload", "Website URL"))

# --- Load PDF ---
if input_mode == "PDF Upload":
    uploaded_files = st.file_uploader("Upload PDF(s)", type="pdf", accept_multiple_files=True)
    if uploaded_files:
        for uploaded_file in uploaded_files:
            with open("temp.pdf", "wb") as f:
                f.write(uploaded_file.getvalue())
            loader = PyPDFLoader("temp.pdf")
            st.session_state.documents.extend(loader.load())
        st.session_state.document_source = "PDF"

# --- Load Website ---
elif input_mode == "Website URL":
    url = st.text_input("Enter Website URL")
    if st.button("Load Website") and url:
        loader = WebBaseLoader(web_paths=[url])
        st.session_state.documents = loader.load()
        st.session_state.document_source = url

# --- Build Vector Store and Chain ---
if st.session_state.documents:
    splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200)
    splits = splitter.split_documents(st.session_state.documents)

    vectorstore = Chroma.from_documents(
        documents=splits,
        embedding=embeddings,
        collection_name="rag_collection",
        persist_directory=None,
        client_settings=Settings(anonymized_telemetry=False)
    )
    st.session_state.vectorstore = vectorstore
    retriever = vectorstore.as_retriever()

    # Prompts
    contextualize_prompt = ChatPromptTemplate.from_messages([
        ("system", "Given a chat history and the latest user question which might reference context in the chat history, formulate a standalone question."),
        MessagesPlaceholder("chat_history"),
        ("human", "{input}")
    ])

    qa_prompt = ChatPromptTemplate.from_messages([
        ("system", "Use the following context to answer the question in 3 sentences or less.{context}"),
        MessagesPlaceholder("chat_history"),
        ("human", "{input}")
    ])

    # Chains
    history_aware_retriever = create_history_aware_retriever(llm, retriever, contextualize_prompt)
    qa_chain = create_stuff_documents_chain(llm, qa_prompt)
    rag_chain = create_retrieval_chain(history_aware_retriever, qa_chain)

    def get_session_history(session_id: str) -> BaseChatMessageHistory:
        if session_id not in st.session_state.store:
            st.session_state.store[session_id] = ChatMessageHistory()
        return st.session_state.store[session_id]

    conversational_rag_chain = RunnableWithMessageHistory(
        rag_chain,
        get_session_history,
        input_messages_key="input",
        history_messages_key="chat_history",
        output_messages_key="answer"
    )

    st.divider()
    st.write(f"You are chatting with: **{st.session_state.document_source}**")

    user_input = st.text_input("Your question:")
    if user_input:
        session_history = get_session_history(session_id)
        response = conversational_rag_chain.invoke(
            {"input": user_input},
            config={"configurable": {"session_id": session_id}}
        )
        st.write("**Assistant:**", response['answer'])

        with st.expander("📜 Chat History"):
            for msg in session_history.messages:
                role = "You" if msg.type == "human" else "Assistant"
                st.markdown(f"**{role}:** {msg.content}")

elif not groq_api_key:
    st.warning("❗ Please set your GROQ_API_KEY in the Space secrets.")