|
|
| |
| |
| |
| |
| import os |
| import chromadb |
| from dotenv import load_dotenv |
| from typing import Dict, List, Any, TypedDict |
| from datetime import datetime |
| import streamlit as st |
| import httpx |
|
|
| from langchain_core.runnables import RunnablePassthrough |
| from langchain.prompts import ChatPromptTemplate |
| from langchain_core.output_parsers import StrOutputParser |
| from langchain_core.tools import tool |
| from langchain.agents import create_tool_calling_agent, AgentExecutor |
| from langchain_openai import OpenAIEmbeddings, ChatOpenAI |
| from langchain_community.vectorstores import Chroma |
|
|
| from langgraph.graph import StateGraph, END, START |
| from pydantic import BaseModel |
|
|
| from groq import Groq |
| from mem0 import MemoryClient |
|
|
| |
| |
| |
| |
| |
| |
| load_dotenv() |
|
|
| openai_api_key = os.environ.get("OPENAI_API_KEY") |
| openai_api_base = os.environ.get("OPENAI_API_BASE") |
| groq_api_key = os.environ.get('GROQ_API_KEY') |
| mem0_api_key = os.environ.get('MEM0_API_KEY') |
|
|
| |
| llm = ChatOpenAI( |
| openai_api_base=openai_api_base, |
| openai_api_key=openai_api_key, |
| model="gpt-4o-mini", |
| streaming=False |
| ) |
|
|
| |
| embedding_model = OpenAIEmbeddings( |
| openai_api_base=openai_api_base, |
| openai_api_key=openai_api_key, |
| model='text-embedding-ada-002' |
| ) |
|
|
|
|
| |
| |
| |
| |
|
|
| |
| class AgentState(TypedDict): |
| query: str |
| expanded_query: str |
| context: List[Dict[str, Any]] |
| response: Any |
| precision_score: float |
| groundedness_score: float |
| groundedness_loop_count: int |
| precision_loop_count: int |
| feedback: str |
| query_feedback: str |
| loop_max_iter: int |
|
|
| |
| |
| vector_store = Chroma( |
| collection_name='nutritional_hypotheticals', |
| persist_directory="./nutritional_db", |
| embedding_function=embedding_model |
| ) |
| retriever = vector_store.as_retriever(search_type='similarity', search_kwargs={'k': 5}) |
|
|
| |
| def expand_query(state): |
| print("---------Expanding Query---------") |
| system_message = '''You are an expert at query expansion. Your goal is to rewrite the user's query to be more specific and comprehensive, making it ideal for a vector database search focused on nutritional disorders. |
| When expanding the query, consider the following: |
| - **Clarify Ambiguities**: Resolve any vague terms or phrases. |
| - **Add Synonyms and Related Terms**: Include alternative names for disorders, symptoms, or treatments. |
| - **Specify Context**: Frame the query within the context of nutritional health, deficiencies, symptoms, causes, and treatments. |
| - **Use Feedback**: Incorporate suggestions from previous refinement steps to improve the query. |
| Provide only the expanded query as a single, continuous string.''' |
| expand_prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_message), |
| ("user", "Expand this query: {query} using the feedback: {query_feedback}") |
| ]) |
| chain = expand_prompt | llm | StrOutputParser() |
| expanded_query = chain.invoke({"query": state['query'], "query_feedback":state["query_feedback"]}) |
| state["expanded_query"] = expanded_query |
| return state |
|
|
| def retrieve_context(state): |
| print("---------retrieve_context---------") |
| query = state['expanded_query'] |
| docs = retriever.invoke(query) |
| context = [{"content": doc.page_content, "metadata": doc.metadata} for doc in docs] |
| state['context'] = context |
| return state |
|
|
| def craft_response(state: Dict) -> Dict: |
| print("---------craft_response---------") |
| system_message = '''You are a knowledgeable and precise AI assistant specializing in nutritional disorders. Your task is to provide a clear and accurate answer to the user's query based *strictly* on the provided context. |
| Follow these guidelines: |
| 1. **Ground Your Answer**: Base your entire response on the information found in the context. Do not use any external knowledge. |
| 2. **Be Direct**: Address the user's query directly and concisely. |
| 3. **Acknowledge Limitations**: If the context does not contain the information needed to answer the query, clearly state that the information is not available in the provided documents. |
| 4. **Incorporate Feedback**: Use the provided feedback to refine your response and address any previous shortcomings.''' |
| response_prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_message), |
| ("user", "Query: {query}\nContext: {context}\n\nfeedback: {feedback}") |
| ]) |
| chain = response_prompt | llm |
| response = chain.invoke({ |
| "query": state['query'], |
| "context": "\n".join([doc["metadata"].get("original_content", "") for doc in state['context']]), |
| "feedback": state['feedback'] |
| }) |
| state['response'] = response |
| return state |
|
|
| def score_groundedness(state: Dict) -> Dict: |
| print("---------check_groundedness---------") |
| system_message = '''You are a groundedness scoring expert. Your role is to evaluate whether an AI-generated response is factually supported by the given context. |
| - **Score**: Provide a numerical score from 0.0 to 1.0. |
| - **1.0**: The response is fully and accurately supported by the context. |
| - **0.0**: The response is not supported by the context or contains fabricated information. |
| - **Crucial Rule**: If the provided context is empty or does not contain the information needed to answer the query, but the response still provides a factual answer, the score must be 0.0. |
| - **Output**: Return only the numerical score. Do not add any explanation or extra text.''' |
| groundedness_prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_message), |
| ("user", "Context: {context}\nResponse: {response}\n\nGroundedness score:") |
| ]) |
| chain = groundedness_prompt | llm | StrOutputParser() |
| groundedness_score = float(chain.invoke({ |
| "context": "\n".join([doc["metadata"].get("original_content", "") for doc in state['context']]), |
| "response": state['response'].content |
| })) |
| state['groundedness_loop_count'] += 1 |
| state['groundedness_score'] = groundedness_score |
| return state |
|
|
| def check_precision(state: Dict) -> Dict: |
| print("---------check_precision---------") |
| system_message = '''You are a precision scoring expert. Your role is to evaluate how well an AI-generated response addresses a specific user query. |
| - **Score**: Provide a numerical score from 0.0 to 1.0. |
| - **1.0**: The response is perfectly precise, comprehensive, and directly answers the user's query. |
| - **0.0**: The response is completely irrelevant or fails to answer the query. |
| - **Output**: Return only the numerical score. Do not add any explanation or extra text.''' |
| precision_prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_message), |
| ("user", "Query: {query}\nResponse: {response}\n\nPrecision score:") |
| ]) |
| chain = precision_prompt | llm | StrOutputParser() |
| precision_score = float(chain.invoke({ |
| "query": state['query'], |
| "response": state['response'].content |
| })) |
| state['precision_score'] = precision_score |
| state['precision_loop_count'] += 1 |
| return state |
|
|
| def refine_response(state: Dict) -> Dict: |
| print("---------refine_response---------") |
| system_message = '''You are a response refinement expert. Your task is to provide constructive feedback on an AI-generated response based on a user's query. |
| Analyze the response for: |
| - **Gaps**: Is any crucial information from the query missing? |
| - **Ambiguities**: Are there any unclear or vague statements? |
| - **Inaccuracies**: Does the response contradict the user's intent (even if it's based on the context)? |
| - **Completeness**: Could the response be more thorough while remaining concise? |
| **Do not rewrite the response.** Instead, provide specific, actionable suggestions for improvement.''' |
| refine_response_prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_message), |
| ("user", "Query: {query}\\nResponse: {response}\\n\\n" |
| "What improvements can be made to enhance accuracy and completeness?") |
| ]) |
| chain = refine_response_prompt | llm| StrOutputParser() |
| feedback = f"Previous Response: {state['response'].content}\\nSuggestions: {chain.invoke({'query': state['query'], 'response': state['response'].content})}" |
| state['feedback'] = feedback |
| return state |
|
|
| def refine_query(state: Dict) -> Dict: |
| print("---------refine_query---------") |
| system_message = '''You are a query refinement expert. Your task is to analyze an original user query and its expanded version to suggest improvements for a more effective vector database search. |
| Review the expanded query for: |
| - **Missing Keywords**: Are there essential terms or synonyms that should be added? |
| - **Lack of Specificity**: Could the query be narrowed down to a more precise topic? |
| - **Scope Refinements**: Is the query too broad or too narrow? |
| **Do not rewrite the query.** Instead, provide structured, actionable suggestions for improvement based on the original query's intent.''' |
| refine_query_prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_message), |
| ("user", "Original Query: {query}\\nExpanded Query: {expanded_query}\\n\\n" |
| "What improvements can be made for a better search?") |
| ]) |
| chain = refine_query_prompt | llm | StrOutputParser() |
| query_feedback = f"Previous Expanded Query: {state['expanded_query']}\\nSuggestions: {chain.invoke({'query': state['query'], 'expanded_query': state['expanded_query']})}" |
| state['query_feedback'] = query_feedback |
| return state |
|
|
| |
| def should_continue_groundedness(state): |
| if state['groundedness_score'] >= 0.7: |
| return "check_precision" |
| else: |
| return "max_iterations_reached" if state["groundedness_loop_count"] >= state['loop_max_iter'] else "refine_response" |
|
|
| def should_continue_precision(state: Dict) -> str: |
| if state['precision_score'] >= 0.7: |
| return "pass" |
| else: |
| return "max_iterations_reached" if state['precision_loop_count'] >= state['loop_max_iter'] else "refine_query" |
|
|
| def max_iterations_reached(state: Dict) -> Dict: |
| state['response'] = "I'm unable to refine the response further. Please provide more context or clarify your question." |
| return state |
|
|
| |
| def create_workflow() -> StateGraph: |
| workflow = StateGraph(AgentState) |
| workflow.add_node("expand_query", expand_query) |
| workflow.add_node("retrieve_context", retrieve_context) |
| workflow.add_node("craft_response", craft_response) |
| workflow.add_node("score_groundedness", score_groundedness) |
| workflow.add_node("refine_response", refine_response) |
| workflow.add_node("check_precision", check_precision) |
| workflow.add_node("refine_query", refine_query) |
| workflow.add_node("max_iterations_reached", max_iterations_reached) |
|
|
| workflow.add_edge(START, "expand_query") |
| workflow.add_edge("expand_query", "retrieve_context") |
| workflow.add_edge("retrieve_context", "craft_response") |
| workflow.add_edge("craft_response", "score_groundedness") |
| workflow.add_conditional_edges("score_groundedness", should_continue_groundedness, {"check_precision": "check_precision", "refine_response": "refine_response", "max_iterations_reached": "max_iterations_reached"}) |
| workflow.add_edge("refine_response", "craft_response") |
| workflow.add_conditional_edges("check_precision", should_continue_precision, {"pass": END, "refine_query": "refine_query", "max_iterations_reached": "max_iterations_reached"}) |
| workflow.add_edge("refine_query", "expand_query") |
| workflow.add_edge("max_iterations_reached", END) |
| return workflow |
|
|
| WORKFLOW_APP = create_workflow().compile() |
|
|
| |
| @tool |
| def agentic_rag(query: str): |
| """Runs the RAG-based agent for context-aware responses.""" |
| inputs = { |
| "query": query, "expanded_query": "", "context": [], "response": "", |
| "precision_score": 0.0, "groundedness_score": 0.0, |
| "groundedness_loop_count": 0, "precision_loop_count": 0, |
| "feedback": "", "query_feedback": "", "loop_max_iter": 3 |
| } |
| output = WORKFLOW_APP.invoke(inputs) |
| final_response = output.get('response') |
| if hasattr(final_response, 'content'): |
| return final_response.content |
| return str(final_response) |
|
|
|
|
| |
| |
| |
| |
| llama_guard_client = Groq( |
| api_key=groq_api_key, |
| http_client=httpx.Client() |
| ) |
| def filter_input_with_llama_guard(user_input, model="meta-llama/llama-guard-4-12b"): |
| try: |
| response = llama_guard_client.chat.completions.create( |
| messages=[{"role": "user", "content": user_input}], |
| model=model, |
| temperature=0.0 |
| ) |
| return response.choices[0].message.content.strip() |
| except Exception as e: |
| print(f"Error with Llama Guard (Groq): {e}") |
| return "safe" |
|
|
|
|
| |
| |
| |
| |
| class NutritionBot: |
| def __init__(self): |
| self.memory = MemoryClient(api_key=mem0_api_key) |
| self.client = ChatOpenAI( |
| model_name="gpt-4o-mini", |
| openai_api_key=openai_api_key, |
| openai_api_base=openai_api_base, |
| temperature=0 |
| ) |
| tools = [agentic_rag] |
| system_prompt = """You are a Medical Support Agent specializing ONLY in nutritional disorders...""" |
| prompt = ChatPromptTemplate.from_messages([ |
| ("system", system_prompt), |
| ("human", "{input}"), |
| ("placeholder", "{agent_scratchpad}") |
| ]) |
| agent = create_tool_calling_agent(self.client, tools, prompt) |
| self.agent_executor = AgentExecutor(agent=agent, tools=tools, verbose=True) |
|
|
| def store_customer_interaction(self, user_id: str, message: str, response: str, metadata: Dict = None): |
| if metadata is None: metadata = {} |
| metadata["timestamp"] = datetime.now().isoformat() |
| conversation = [{"role": "user", "content": message}, {"role": "assistant", "content": response}] |
| self.memory.add(messages=conversation, user_id=user_id, metadata=metadata) |
|
|
| def get_relevant_history(self, user_id: str, query: str) -> str: |
| memories = self.memory.search(query=query, user_id=user_id, limit=5) |
| context = "Previous relevant interactions:\\n" |
| if not memories: |
| return "No previous relevant interactions found.\\n" |
| for mem in memories: |
| context += f"Summary of past interaction: {mem.get('memory', 'N/A')}\\n---\\n" |
| return context |
|
|
| def handle_customer_query(self, user_id: str, query: str) -> str: |
| context = self.get_relevant_history(user_id, query) |
| prompt = f"Context:\\n{context}\\nCurrent customer query: {query}\\nProvide a helpful response that takes into account any relevant past interactions." |
| response = self.agent_executor.invoke({"input": prompt}) |
| self.store_customer_interaction(user_id=user_id, message=query, response=response["output"]) |
| return response['output'] |
|
|
|
|
| |
| |
| |
| |
| def nutrition_disorder_streamlit(): |
| st.title("Nutrition Disorder Specialist") |
| st.write("Ask me anything about nutrition disorders, symptoms, causes, treatments, and more.") |
|
|
| if 'chat_history' not in st.session_state: |
| st.session_state.chat_history = [] |
| if 'user_id' not in st.session_state: |
| st.session_state.user_id = None |
| if 'chatbot' not in st.session_state: |
| st.session_state.chatbot = None |
|
|
| if st.session_state.user_id is None: |
| with st.form("login_form"): |
| user_id = st.text_input("Please enter your name to begin:") |
| submit_button = st.form_submit_button("Login") |
| if submit_button and user_id: |
| st.session_state.user_id = user_id |
| st.session_state.chatbot = NutritionBot() |
| welcome_msg = f"Welcome, {user_id}! How can I help you with nutrition disorders today?" |
| st.session_state.chat_history.append({"role": "assistant", "content": welcome_msg}) |
| st.rerun() |
| else: |
| for message in st.session_state.chat_history: |
| with st.chat_message(message["role"]): |
| st.write(message["content"]) |
|
|
| if user_query := st.chat_input("Ask about a nutrition disorder..."): |
| st.session_state.chat_history.append({"role": "user", "content": user_query}) |
| with st.chat_message("user"): |
| st.write(user_query) |
|
|
| filtered_result = filter_input_with_llama_guard(user_query) |
| if "safe" in filtered_result: |
| try: |
| with st.spinner("Thinking..."): |
| response = st.session_state.chatbot.handle_customer_query(st.session_state.user_id, user_query) |
| st.session_state.chat_history.append({"role": "assistant", "content": response}) |
| with st.chat_message("assistant"): |
| st.write(response) |
| except Exception as e: |
| error_msg = f"Sorry, I encountered an error: {e}" |
| st.session_state.chat_history.append({"role": "assistant", "content": error_msg}) |
| with st.chat_message("assistant"): |
| st.write(error_msg) |
| else: |
| inappropriate_msg = "I apologize, but I cannot process that input as it may be inappropriate." |
| st.session_state.chat_history.append({"role": "assistant", "content": inappropriate_msg}) |
| with st.chat_message("assistant"): |
| st.write(inappropriate_msg) |
|
|
| if __name__ == "__main__": |
| nutrition_disorder_streamlit() |
|
|