PDF-Chatbot / app.py
Sachin B M
block setting
9b5201a
Raw History Blame Contribute Delete
13.4 kB
import os
import gradio as gr
import spaces
from dotenv import load_dotenv
from PyPDF2 import PdfReader
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_cohere import CohereEmbeddings, ChatCohere
from langchain_community.vectorstores import FAISS
from langchain_classic.chains import ConversationalRetrievalChain
load_dotenv()
# =========================================================
# 1. ZERO-GPU STARTUP CHECK
# =========================================================
@spaces.GPU(duration=1)
def zerogpu_check():
return None
# =========================================================
# 2. EXTRACT TEXT FROM PDF FILES
# =========================================================
def get_pdf_text(pdf_files):
text = ""
for pdf_file in pdf_files:
pdf_reader = PdfReader(pdf_file)
for page in pdf_reader.pages:
text += page.extract_text() or ""
return text
# =========================================================
# 3. SPLIT TEXT INTO CHUNKS
# =========================================================
def get_text_chunks(text):
text_splitter = RecursiveCharacterTextSplitter(
separators=["\n\n", "\n", " ", ""],
chunk_size=1000,
chunk_overlap=200,
length_function=len
)
return text_splitter.split_text(text)
# =========================================================
# 4. CREATE COHERE EMBEDDINGS
# =========================================================
def get_embeddings():
return CohereEmbeddings(
model="embed-english-v3.0",
user_agent="langchain"
)
# =========================================================
# 5. CREATE CONVERSATIONAL RETRIEVAL CHAIN
# =========================================================
def get_conversational_chain(vectorstore):
llm = ChatCohere(
model="command-r-08-2024",
temperature=0
)
conversational_chain = ConversationalRetrievalChain.from_llm(
llm=llm,
retriever=vectorstore.as_retriever(),
return_source_documents=False
)
return conversational_chain
# =========================================================
# 6. PROCESS UPLOADED DOCUMENTS
# =========================================================
def process_documents(pdf_files):
if not pdf_files:
return (
None,
[],
"Please upload at least one PDF."
)
if not os.getenv("COHERE_API_KEY"):
return (
None,
[],
"Error: Cohere API key is missing."
)
try:
raw_text = get_pdf_text(pdf_files)
if not raw_text.strip():
return (
None,
[],
"No extractable text found in the PDFs."
)
chunks = get_text_chunks(raw_text)
if not chunks:
return (
None,
[],
"No text chunks were created."
)
embeddings = get_embeddings()
vectorstore = FAISS.from_texts(
texts=chunks,
embedding=embeddings
)
chain = get_conversational_chain(vectorstore)
return (
chain,
[],
"Documents processed successfully! You can now ask questions."
)
except Exception as e:
return (
None,
[],
f"Error processing documents: {str(e)}"
)
# =========================================================
# 7. CONVERT GRADIO CHAT HISTORY TO LANGCHAIN FORMAT
# =========================================================
def convert_chat_history(history):
chat_history = []
pending_question = None
for message in history or []:
if isinstance(message, dict):
role = message.get("role")
content = message.get("content", "")
if isinstance(content, list):
content = " ".join(
str(item.get("text", ""))
if isinstance(item, dict)
else str(item)
for item in content
)
content = str(content)
if role == "user":
pending_question = content
elif role == "assistant" and pending_question is not None:
chat_history.append(
(pending_question, content)
)
pending_question = None
elif isinstance(message, (list, tuple)) and len(message) == 2:
user_message, assistant_message = message
if user_message:
chat_history.append(
(
str(user_message),
str(assistant_message or "")
)
)
return chat_history
# =========================================================
# 8. HANDLE USER QUESTIONS
# =========================================================
def handle_user_question(question, history, chain):
history = list(history or [])
if not question or not question.strip():
return history, ""
question = question.strip()
if chain is None:
history.append({
"role": "user",
"content": question
})
history.append({
"role": "assistant",
"content": "Please upload and process your PDFs first."
})
return history, ""
try:
chat_history = convert_chat_history(history)
response = chain.invoke({
"question": question,
"chat_history": chat_history
})
answer = response.get(
"answer",
"No answer was generated."
)
if isinstance(answer, list):
answer = " ".join(
str(item.get("text", ""))
if isinstance(item, dict)
else str(item)
for item in answer
)
else:
answer = str(answer)
history.append({
"role": "user",
"content": question
})
history.append({
"role": "assistant",
"content": answer
})
return history, ""
except Exception as e:
history.append({
"role": "user",
"content": question
})
history.append({
"role": "assistant",
"content": f"Error generating answer: {str(e)}"
})
return history, ""
# =========================================================
# 9. CLEAR CHAT
# =========================================================
def clear_chat():
return []
# =========================================================
# 10. CUSTOM CSS
# =========================================================
custom_css = """
.gradio-container {
background: #0e0f14 !important;
color: #f8fafc !important;
max-width: 100% !important;
min-height: 100vh;
padding: 16px !important;
}
#main-row {
display: flex !important;
align-items: flex-start !important;
gap: 20px !important;
max-width: 1500px;
margin: 0 auto !important;
}
#sidebar {
background: #24252f !important;
border: 1px solid #343642 !important;
border-radius: 14px !important;
padding: 16px !important;
height: fit-content !important;
min-height: 0 !important;
align-self: flex-start !important;
margin-top: 40px;
}
#sidebar h2 {
font-size: 1.15rem !important;
margin: 0 0 12px 0 !important;
color: #ffffff !important;
}
#sidebar [data-testid="file-upload"] {
min-height: 110px !important;
}
#process-button {
margin-top: 10px !important;
border-radius: 8px !important;
transition: transform 0.2s ease, box-shadow 0.2s ease;
}
#process-button:hover {
transform: translateY(-1px);
box-shadow: 0 4px 12px rgba(37, 99, 235, 0.3);
}
#status {
margin-top: 10px !important;
}
#main-panel {
padding: 0 8px !important;
min-width: 0 !important;
}
#hero {
background: linear-gradient(120deg, #0f172a 0%, #1d4ed8 100%);
padding: 28px 32px;
border-radius: 18px;
box-shadow: 0 10px 30px rgba(29, 78, 216, 0.18);
margin: 24px 0 18px 0;
transition: box-shadow 0.2s ease;
}
#hero:hover {
box-shadow: 0 12px 34px rgba(29, 78, 216, 0.28);
}
#hero h1 {
color: white !important;
font-size: clamp(1.5rem, 2.5vw, 2.3rem) !important;
font-weight: 800 !important;
line-height: 1.25 !important;
margin: 0 !important;
}
#hero p {
color: #dbeafe !important;
font-size: 1rem !important;
margin: 10px 0 0 0 !important;
}
#chatbot {
border-radius: 12px !important;
min-height: 280px !important;
height: calc(100vh - 390px) !important;
max-height: 650px !important;
overflow: auto !important;
}
#question-box {
margin-top: 12px !important;
}
#question-box textarea {
min-height: 48px !important;
border-radius: 10px !important;
}
#main-panel button {
border-radius: 8px !important;
transition: background 0.2s ease;
}
footer {
display: none !important;
}
@media (max-width: 900px) {
.gradio-container {
padding: 12px !important;
}
#main-row {
gap: 14px !important;
}
#sidebar {
margin-top: 30px;
padding: 12px !important;
}
#hero {
padding: 22px;
}
#chatbot {
height: 55vh !important;
}
}
@media (max-width: 700px) {
.gradio-container {
padding: 8px !important;
}
#main-row {
flex-direction: column !important;
gap: 12px !important;
}
#sidebar {
width: 100% !important;
max-width: 100% !important;
margin-top: 0 !important;
padding: 14px !important;
}
#main-panel {
width: 100% !important;
max-width: 100% !important;
padding: 0 !important;
}
#hero {
margin: 8px 0 14px 0;
padding: 20px;
border-radius: 14px;
}
#hero h1 {
font-size: 1.5rem !important;
}
#hero p {
font-size: 0.9rem !important;
}
#chatbot {
height: 55vh !important;
min-height: 300px !important;
max-height: none !important;
}
#question-box {
margin-top: 10px !important;
}
}
"""
# =========================================================
# 11. BUILD GRADIO INTERFACE
# =========================================================
with gr.Blocks(
title="Chat with Documents"
) as demo:
# State to store the conversational chain
chain_state = gr.State(None)
# Hidden event to satisfy ZeroGPU startup detection
hidden_button = gr.Button(visible=False)
hidden_button.click(
fn=zerogpu_check,
inputs=[],
outputs=[]
)
with gr.Row(
elem_id="main-row",
equal_height=False
):
# LEFT SIDEBAR
with gr.Column(
scale=1,
min_width=250,
elem_id="sidebar"
):
gr.Markdown("## Upload Documents")
pdf_upload = gr.File(
file_types=[".pdf"],
file_count="multiple",
type="filepath",
height=120
)
process_button = gr.Button(
"Process Documents",
variant="primary",
elem_id="process-button"
)
status = gr.Textbox(
label="Status",
interactive=False,
lines=2,
elem_id="status"
)
# RIGHT CHAT AREA
with gr.Column(
scale=4,
min_width=300,
elem_id="main-panel"
):
gr.HTML("""
<div id="hero">
<h1>AI Chatbot for Your Documents</h1>
<p>
Upload a PDF and ask questions in natural language.
</p>
</div>
""")
chatbot = gr.Chatbot(
label="Chat with your PDFs",
height=500,
elem_id="chatbot"
)
question = gr.Textbox(
label="Ask a question",
placeholder="Ask a question about your PDFs",
lines=1,
max_lines=1,
elem_id="question-box"
)
clear_button = gr.Button("Clear Chat")
# Process PDF documents
process_button.click(
fn=process_documents,
inputs=[pdf_upload],
outputs=[chain_state, chatbot, status]
)
# Press Enter to send a question
question.submit(
fn=handle_user_question,
inputs=[question, chatbot, chain_state],
outputs=[chatbot, question]
)
# Clear conversation
clear_button.click(
fn=clear_chat,
inputs=[],
outputs=[chatbot]
)
# =========================================================
# 12. LAUNCH APPLICATION
# =========================================================
if __name__ == "__main__":
demo.queue().launch(
theme=gr.themes.Base(),
css=custom_css,
server_name="0.0.0.0",
server_port=7860,
ssr_mode=False
)