amanm10000/MLSC-Coherence-25-FAQ-Chatbot-API
0
1import os2import getpass3from groq import Groq4from langchain.chat_models import init_chat_model5from langchain_core.messages import HumanMessage, SystemMessage6from langchain_core.vectorstores import InMemoryVectorStore7from langchain_core.documents import Document8from langchain_text_splitters import RecursiveCharacterTextSplitter9from langchain_community.document_loaders import UnstructuredMarkdownLoader10from langchain_community.embeddings import HuggingFaceInferenceAPIEmbeddings11from langchain import hub12from langgraph.graph import START, StateGraph13from pydantic.main import BaseModel14from typing_extensions import List, TypedDict15 16from langchain_cohere import CohereEmbeddings17 18import re19# from dotenv import load_dotenv20from fastapi import FastAPI21from fastapi.middleware.cors import CORSMiddleware22from fastapi.responses import JSONResponse23 24'''25if not os.environ.get("GROQ_API_KEY"):26 os.environ["GROQ_API_KEY"] = getpass.getpass("Enter API key for Groq: ")27'''28 29# load_dotenv()30 31# print(f"GROQ_API_KEY: {os.getenv('GROQ_API_KEY')}")32# print(f"HUGGING_FACE_API_KEY: {os.getenv('HUGGING_FACE_API_KEY')}")33 34llm = init_chat_model("qwen-qwq-32b", model_provider="groq", api_key=os.environ["GROQ_API_KEY"])35'''36embeddings = HuggingFaceInferenceAPIEmbeddings(37 api_key = os.getenv('HUGGING_FACE_API_KEY'),38 model_name="sentence-transformers/all-MiniLM-L6-v2"39)40 41embeddings = HuggingFaceInferenceAPIEmbeddings(42 api_key=os.getenv('HUGGING_FACE_API_KEY'), model_name="sentence-transformers/all-MiniLM-L6-v2"43)'''44 45embeddings = CohereEmbeddings(46 cohere_api_key=os.environ['COHERE'],47 model="embed-english-v3.0", # Added this line48 user_agent="langchain-cohere-embeddings"49)50 51vector_store = InMemoryVectorStore(embedding=embeddings)52 53# Data - 1 and Data - 254data_1 = open(r'data_1.txt', 'r').read()55data_2 = open(r'data_2.txt', 'r').read()56data_3 = open(r'data_3.txt', 'r').read()57data_4 = open(r'data_4.txt', 'r').read()58 59comb = open(r'comb.txt', 'r').read()60 61md_loader = UnstructuredMarkdownLoader('comb.md')62 63text_splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=100)64# all_splits = text_splitter.split_text(data_1 + "\n\n" + data_2 + "\n\n" + data_3 + "\n\n" + data_4)65# all_splits = text_splitter.split_text(comb)66all_splits = text_splitter.split_documents(md_loader.load())67 68# docs = [Document(page_content=text) for text in all_splits]69docs = [Document(page_content=text.page_content, metadata=text.metadata) for text in all_splits]70_ = vector_store.add_documents(documents=docs)71 72 73prompt = hub.pull("rlm/rag-prompt")74 75# Replace with custom prompt76system_message = """You are a helpful and professional FAQ chatbot for the MLSC Coherence 25 Hackathon. Your role is to:771. Provide accurate and concise answers based on the provided context782. Be friendly but professional in tone793. If you don't know the answer, simply say "I don't have information about that"804. Keep responses brief and to the point815. Focus on providing factual information from the context826. Never mention "the provided context" or similar phrases in your responses837. Never explain why you don't know something - just state that you don't know848. Be direct and avoid unnecessary explanations"""85 86human_message_template = """Context: {context}87 88Question: {question}89 90Please provide a clear and concise answer based on the context above."""91 92class State(TypedDict):93 question: str94 context: List[Document]95 answer: str96 97def retrieve(state: State):98 retrieved_docs = vector_store.similarity_search(state["question"])99 return {"context": retrieved_docs}100 101def generate(state: State):102 docs_content = "\n\n".join(doc.page_content for doc in state["context"])103 messages = [104 SystemMessage(content=system_message),105 HumanMessage(content=human_message_template.format(106 context=docs_content,107 question=state["question"]108 ))109 ]110 print(messages)111 response = llm.invoke(messages)112 return {"answer": response.content}113 114graph_builder = StateGraph(State).add_sequence([retrieve, generate])115graph_builder.add_edge(START, "retrieve")116graph = graph_builder.compile()117'''118response = graph.invoke({"question": "Who should i contact for help ?"})119print(response["answer"])120'''121 122app = FastAPI()123 124origins = ["*"]125 126app.add_middleware(127 CORSMiddleware,128 allow_origins=origins,129 allow_credentials=True,130 allow_methods=["GET", "POST", "PUT", "DELETE"],131 allow_headers=["*"],132)133 134@app.get("/ping")135async def ping():136 return "Pong!"137 138class Query(BaseModel):139 question: str140 141@app.get("/chat")142async def chat(request: Query):143 response = graph.invoke({"question": request.question})144 response = response["answer"]145 response = re.sub(r'<think>.*?</think>', '', response, flags=re.DOTALL)146 # response = response[4:]147 return {"response": response}148 149@app.post("/chat")150async def chat(request: Query):151 response = graph.invoke({"question": request.question})152 response = response["answer"]153 response = re.sub(r'<think>.*?</think>', '', response, flags=re.DOTALL)154 # response = response[4:]155 return {"response": response}156 157@app.get("/")158async def root():159 return {"message": "Hello World"}160 