adaptFast commited on
Commit
d8faa94
·
verified ·
1 Parent(s): d253655

Upload app.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. app.py +227 -52
app.py CHANGED
@@ -9,7 +9,7 @@ from dotenv import load_dotenv
9
  from typing import Dict, List, Any, TypedDict
10
  from datetime import datetime
11
  import streamlit as st
12
- import httpx
13
 
14
  from langchain_core.runnables import RunnablePassthrough
15
  from langchain.prompts import ChatPromptTemplate
@@ -28,7 +28,9 @@ from mem0 import MemoryClient
28
  # ==============================================================================
29
  # 2. SETUP & CONFIGURATION
30
  # Load secrets and initialize core models (LLM, Embeddings).
 
31
  # ==============================================================================
 
32
  load_dotenv()
33
 
34
  openai_api_key = os.environ.get("OPENAI_API_KEY")
@@ -36,6 +38,7 @@ openai_api_base = os.environ.get("OPENAI_API_BASE")
36
  groq_api_key = os.environ.get('GROQ_API_KEY')
37
  mem0_api_key = os.environ.get('MEM0_API_KEY')
38
 
 
39
  llm = ChatOpenAI(
40
  openai_api_base=openai_api_base,
41
  openai_api_key=openai_api_key,
@@ -43,175 +46,347 @@ llm = ChatOpenAI(
43
  streaming=False
44
  )
45
 
 
46
  embedding_model = OpenAIEmbeddings(
47
  openai_api_base=openai_api_base,
48
  openai_api_key=openai_api_key,
49
  model='text-embedding-ada-002'
50
  )
51
 
 
52
  # ==============================================================================
53
  # 3. ADVANCED RAG AGENT WORKFLOW
 
54
  # ==============================================================================
 
 
55
  class AgentState(TypedDict):
56
- query: str; expanded_query: str; context: List[Dict[str, Any]]; response: Any
57
- precision_score: float; groundedness_score: float; groundedness_loop_count: int
58
- precision_loop_count: int; feedback: str; query_feedback: str; loop_max_iter: int
 
 
 
 
 
 
 
 
59
 
 
 
60
  vector_store = Chroma(
61
  collection_name='nutritional_hypotheticals',
62
- persist_directory="./nutritional_db",
63
  embedding_function=embedding_model
64
  )
65
  retriever = vector_store.as_retriever(search_type='similarity', search_kwargs={'k': 5})
66
 
 
67
  def expand_query(state):
68
  print("---------Expanding Query---------")
69
- system_message = '''You are an expert at query expansion...''' # Your full prompt
70
- chain = ChatPromptTemplate.from_messages([("system", system_message), ("user", "Expand this query: {query} using the feedback: {query_feedback}")]) | llm | StrOutputParser()
71
- state["expanded_query"] = chain.invoke({"query": state['query'], "query_feedback":state["query_feedback"]})
 
 
 
 
 
 
 
 
 
 
 
72
  return state
73
 
74
  def retrieve_context(state):
75
  print("---------retrieve_context---------")
76
- docs = retriever.invoke(state['expanded_query'])
77
- state['context'] = [{"content": doc.page_content, "metadata": doc.metadata} for doc in docs]
 
 
78
  return state
79
 
80
  def craft_response(state: Dict) -> Dict:
81
  print("---------craft_response---------")
82
- system_message = '''You are a knowledgeable and precise AI assistant...''' # Your full prompt
83
- chain = ChatPromptTemplate.from_messages([("system", system_message), ("user", "Query: {query}\nContext: {context}\n\nfeedback: {feedback}")]) | llm
84
- state['response'] = chain.invoke({"query": state['query'], "context": "\n".join([doc["metadata"].get("original_content", "") for doc in state['context']]), "feedback": state['feedback']})
 
 
 
 
 
 
 
 
 
 
 
 
 
 
85
  return state
86
 
87
  def score_groundedness(state: Dict) -> Dict:
88
  print("---------check_groundedness---------")
89
- system_message = '''You are a groundedness scoring expert...''' # Your full prompt
90
- chain = ChatPromptTemplate.from_messages([("system", system_message), ("user", "Context: {context}\nResponse: {response}\n\nGroundedness score:")]) | llm | StrOutputParser()
91
- state['groundedness_score'] = float(chain.invoke({"context": "\n".join([doc["metadata"].get("original_content", "") for doc in state['context']]), "response": state['response'].content}))
 
 
 
 
 
 
 
 
 
 
 
 
92
  state['groundedness_loop_count'] += 1
 
93
  return state
94
 
95
  def check_precision(state: Dict) -> Dict:
96
  print("---------check_precision---------")
97
- system_message = '''You are a precision scoring expert...''' # Your full prompt
98
- chain = ChatPromptTemplate.from_messages([("system", system_message), ("user", "Query: {query}\nResponse: {response}\n\nPrecision score:")]) | llm | StrOutputParser()
99
- state['precision_score'] = float(chain.invoke({"query": state['query'], "response": state['response'].content}))
 
 
 
 
 
 
 
 
 
 
 
 
100
  state['precision_loop_count'] += 1
101
  return state
102
 
103
- def refine_response(state: Dict) -> Dict: return state # Placeholder for brevity
104
- def refine_query(state: Dict) -> Dict: return state # Placeholder for brevity
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105
 
 
106
  def should_continue_groundedness(state):
107
- return "check_precision" if state['groundedness_score'] >= 0.7 else ("max_iterations_reached" if state["groundedness_loop_count"] >= state['loop_max_iter'] else "refine_response")
 
 
 
108
 
109
  def should_continue_precision(state: Dict) -> str:
110
- return "pass" if state['precision_score'] >= 0.7 else ("max_iterations_reached" if state['precision_loop_count'] >= state['loop_max_iter'] else "refine_query")
 
 
 
111
 
112
  def max_iterations_reached(state: Dict) -> Dict:
113
- state['response'] = "I'm unable to refine the response further."
114
  return state
115
 
 
116
  def create_workflow() -> StateGraph:
117
  workflow = StateGraph(AgentState)
118
- nodes = ["expand_query", "retrieve_context", "craft_response", "score_groundedness", "refine_response", "check_precision", "refine_query", "max_iterations_reached"]
119
- for node in nodes: workflow.add_node(node, globals()[node])
120
- 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")
 
 
 
 
 
 
 
 
 
 
121
  workflow.add_conditional_edges("score_groundedness", should_continue_groundedness, {"check_precision": "check_precision", "refine_response": "refine_response", "max_iterations_reached": "max_iterations_reached"})
122
  workflow.add_edge("refine_response", "craft_response")
123
  workflow.add_conditional_edges("check_precision", should_continue_precision, {"pass": END, "refine_query": "refine_query", "max_iterations_reached": "max_iterations_reached"})
124
- workflow.add_edge("refine_query", "expand_query"); workflow.add_edge("max_iterations_reached", END)
 
125
  return workflow
126
 
127
  WORKFLOW_APP = create_workflow().compile()
128
 
 
129
  @tool
130
  def agentic_rag(query: str):
131
- 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}
 
 
 
 
 
 
132
  output = WORKFLOW_APP.invoke(inputs)
133
  final_response = output.get('response')
134
- return final_response.content if hasattr(final_response, 'content') else str(final_response)
 
 
 
135
 
136
  # ==============================================================================
137
  # 4. SAFETY GUARDRAIL
138
  # ==============================================================================
139
- llama_guard_client = Groq(api_key=groq_api_key, http_client=httpx.Client())
 
 
 
 
140
  def filter_input_with_llama_guard(user_input, model="meta-llama/llama-guard-4-12b"):
141
  try:
142
- response = llama_guard_client.chat.completions.create(messages=[{"role": "user", "content": user_input}], model=model, temperature=0.0)
 
 
 
 
143
  return response.choices[0].message.content.strip()
144
  except Exception as e:
145
- print(f"Error with Llama Guard (Groq): {e}"); return "safe"
 
 
146
 
147
  # ==============================================================================
148
  # 5. NUTRITION BOT CLASS (with Memory)
 
149
  # ==============================================================================
150
  class NutritionBot:
151
  def __init__(self):
152
  self.memory = MemoryClient(api_key=mem0_api_key)
153
- self.client = ChatOpenAI(model_name="gpt-4o-mini", openai_api_key=openai_api_key, openai_api_base=openai_api_base, temperature=0)
 
 
 
 
 
154
  tools = [agentic_rag]
155
- system_prompt = """You are a Medical Support Agent specializing ONLY in nutritional disorders...""" # Your full, robust prompt
156
- prompt = ChatPromptTemplate.from_messages([("system", system_prompt), ("human", "{input}"), ("placeholder", "{agent_scratchpad}")])
 
 
 
 
157
  agent = create_tool_calling_agent(self.client, tools, prompt)
158
  self.agent_executor = AgentExecutor(agent=agent, tools=tools, verbose=True)
159
 
160
  def store_customer_interaction(self, user_id: str, message: str, response: str, metadata: Dict = None):
161
- if metadata is None: metadata = {}; metadata["timestamp"] = datetime.now().isoformat()
162
- self.memory.add(messages=[{"role": "user", "content": message}, {"role": "assistant", "content": response}], user_id=user_id, metadata=metadata)
 
 
163
 
164
  def get_relevant_history(self, user_id: str, query: str) -> str:
165
  memories = self.memory.search(query=query, user_id=user_id, limit=5)
166
  context = "Previous relevant interactions:\\n"
167
- if not memories: return "No previous relevant interactions found.\\n"
168
- for mem in memories: context += f"Summary of past interaction: {mem.get('memory', 'N/A')}\\n---\\n"
 
 
169
  return context
170
 
171
  def handle_customer_query(self, user_id: str, query: str) -> str:
172
  context = self.get_relevant_history(user_id, query)
173
- prompt = f"Context:\\n{context}\\nCurrent customer query: {query}\\nProvide a helpful response..."
174
  response = self.agent_executor.invoke({"input": prompt})
175
  self.store_customer_interaction(user_id=user_id, message=query, response=response["output"])
176
  return response['output']
177
 
 
178
  # ==============================================================================
179
  # 6. STREAMLIT UI
 
180
  # ==============================================================================
181
  def nutrition_disorder_streamlit():
182
  st.title("Nutrition Disorder Specialist")
183
  st.write("Ask me anything about nutrition disorders, symptoms, causes, treatments, and more.")
184
- for key in ['chat_history', 'user_id', 'chatbot']:
185
- if key not in st.session_state: st.session_state[key] = [] if key == 'chat_history' else None
 
 
 
 
 
186
 
187
  if st.session_state.user_id is None:
188
  with st.form("login_form"):
189
  user_id = st.text_input("Please enter your name to begin:")
190
- if st.form_submit_button("Login") and user_id:
 
191
  st.session_state.user_id = user_id
192
  st.session_state.chatbot = NutritionBot()
193
- st.session_state.chat_history.append({"role": "assistant", "content": f"Welcome, {user_id}! How can I help you?"})
 
194
  st.rerun()
195
  else:
196
  for message in st.session_state.chat_history:
197
- with st.chat_message(message["role"]): st.write(message["content"])
 
 
198
  if user_query := st.chat_input("Ask about a nutrition disorder..."):
199
  st.session_state.chat_history.append({"role": "user", "content": user_query})
200
- with st.chat_message("user"): st.write(user_query)
201
- if "safe" in filter_input_with_llama_guard(user_query):
 
 
 
202
  try:
203
  with st.spinner("Thinking..."):
204
  response = st.session_state.chatbot.handle_customer_query(st.session_state.user_id, user_query)
205
  st.session_state.chat_history.append({"role": "assistant", "content": response})
206
- with st.chat_message("assistant"): st.write(response)
 
207
  except Exception as e:
208
  error_msg = f"Sorry, I encountered an error: {e}"
209
  st.session_state.chat_history.append({"role": "assistant", "content": error_msg})
210
- with st.chat_message("assistant"): st.write(error_msg)
 
211
  else:
212
- inappropriate_msg = "I apologize, but I cannot process that input."
213
  st.session_state.chat_history.append({"role": "assistant", "content": inappropriate_msg})
214
- with st.chat_message("assistant"): st.write(inappropriate_msg)
 
215
 
216
  if __name__ == "__main__":
217
  nutrition_disorder_streamlit()
 
9
  from typing import Dict, List, Any, TypedDict
10
  from datetime import datetime
11
  import streamlit as st
12
+ import httpx # ADDED THIS IMPORT
13
 
14
  from langchain_core.runnables import RunnablePassthrough
15
  from langchain.prompts import ChatPromptTemplate
 
28
  # ==============================================================================
29
  # 2. SETUP & CONFIGURATION
30
  # Load secrets and initialize core models (LLM, Embeddings).
31
+ # This section uses environment variables, which is correct for deployment.
32
  # ==============================================================================
33
+ # On Hugging Face Spaces, these will be set as secrets.
34
  load_dotenv()
35
 
36
  openai_api_key = os.environ.get("OPENAI_API_KEY")
 
38
  groq_api_key = os.environ.get('GROQ_API_KEY')
39
  mem0_api_key = os.environ.get('MEM0_API_KEY')
40
 
41
+ # Initialize the Chat OpenAI model
42
  llm = ChatOpenAI(
43
  openai_api_base=openai_api_base,
44
  openai_api_key=openai_api_key,
 
46
  streaming=False
47
  )
48
 
49
+ # Initialize the OpenAI Embeddings model
50
  embedding_model = OpenAIEmbeddings(
51
  openai_api_base=openai_api_base,
52
  openai_api_key=openai_api_key,
53
  model='text-embedding-ada-002'
54
  )
55
 
56
+
57
  # ==============================================================================
58
  # 3. ADVANCED RAG AGENT WORKFLOW
59
+ # This is the complete, self-contained logic for the LangGraph agent.
60
  # ==============================================================================
61
+
62
+ # 3.1. Define Agent State
63
  class AgentState(TypedDict):
64
+ query: str
65
+ expanded_query: str
66
+ context: List[Dict[str, Any]]
67
+ response: Any
68
+ precision_score: float
69
+ groundedness_score: float
70
+ groundedness_loop_count: int
71
+ precision_loop_count: int
72
+ feedback: str
73
+ query_feedback: str
74
+ loop_max_iter: int
75
 
76
+ # 3.2. Load the Vector Store
77
+ # This points to the pre-built database that will be in the Docker container.
78
  vector_store = Chroma(
79
  collection_name='nutritional_hypotheticals',
80
+ persist_directory="./nutritional_db", # The path inside the Docker container
81
  embedding_function=embedding_model
82
  )
83
  retriever = vector_store.as_retriever(search_type='similarity', search_kwargs={'k': 5})
84
 
85
+ # 3.3. Define All Workflow Nodes (Functions)
86
  def expand_query(state):
87
  print("---------Expanding Query---------")
88
+ 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.
89
+ When expanding the query, consider the following:
90
+ - **Clarify Ambiguities**: Resolve any vague terms or phrases.
91
+ - **Add Synonyms and Related Terms**: Include alternative names for disorders, symptoms, or treatments.
92
+ - **Specify Context**: Frame the query within the context of nutritional health, deficiencies, symptoms, causes, and treatments.
93
+ - **Use Feedback**: Incorporate suggestions from previous refinement steps to improve the query.
94
+ Provide only the expanded query as a single, continuous string.'''
95
+ expand_prompt = ChatPromptTemplate.from_messages([
96
+ ("system", system_message),
97
+ ("user", "Expand this query: {query} using the feedback: {query_feedback}")
98
+ ])
99
+ chain = expand_prompt | llm | StrOutputParser()
100
+ expanded_query = chain.invoke({"query": state['query'], "query_feedback":state["query_feedback"]})
101
+ state["expanded_query"] = expanded_query
102
  return state
103
 
104
  def retrieve_context(state):
105
  print("---------retrieve_context---------")
106
+ query = state['expanded_query']
107
+ docs = retriever.invoke(query)
108
+ context = [{"content": doc.page_content, "metadata": doc.metadata} for doc in docs]
109
+ state['context'] = context
110
  return state
111
 
112
  def craft_response(state: Dict) -> Dict:
113
  print("---------craft_response---------")
114
+ 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.
115
+ Follow these guidelines:
116
+ 1. **Ground Your Answer**: Base your entire response on the information found in the context. Do not use any external knowledge.
117
+ 2. **Be Direct**: Address the user's query directly and concisely.
118
+ 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.
119
+ 4. **Incorporate Feedback**: Use the provided feedback to refine your response and address any previous shortcomings.'''
120
+ response_prompt = ChatPromptTemplate.from_messages([
121
+ ("system", system_message),
122
+ ("user", "Query: {query}\nContext: {context}\n\nfeedback: {feedback}")
123
+ ])
124
+ chain = response_prompt | llm
125
+ response = chain.invoke({
126
+ "query": state['query'],
127
+ "context": "\n".join([doc["metadata"].get("original_content", "") for doc in state['context']]),
128
+ "feedback": state['feedback']
129
+ })
130
+ state['response'] = response
131
  return state
132
 
133
  def score_groundedness(state: Dict) -> Dict:
134
  print("---------check_groundedness---------")
135
+ 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.
136
+ - **Score**: Provide a numerical score from 0.0 to 1.0.
137
+ - **1.0**: The response is fully and accurately supported by the context.
138
+ - **0.0**: The response is not supported by the context or contains fabricated information.
139
+ - **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.
140
+ - **Output**: Return only the numerical score. Do not add any explanation or extra text.'''
141
+ groundedness_prompt = ChatPromptTemplate.from_messages([
142
+ ("system", system_message),
143
+ ("user", "Context: {context}\nResponse: {response}\n\nGroundedness score:")
144
+ ])
145
+ chain = groundedness_prompt | llm | StrOutputParser()
146
+ groundedness_score = float(chain.invoke({
147
+ "context": "\n".join([doc["metadata"].get("original_content", "") for doc in state['context']]),
148
+ "response": state['response'].content
149
+ }))
150
  state['groundedness_loop_count'] += 1
151
+ state['groundedness_score'] = groundedness_score
152
  return state
153
 
154
  def check_precision(state: Dict) -> Dict:
155
  print("---------check_precision---------")
156
+ system_message = '''You are a precision scoring expert. Your role is to evaluate how well an AI-generated response addresses a specific user query.
157
+ - **Score**: Provide a numerical score from 0.0 to 1.0.
158
+ - **1.0**: The response is perfectly precise, comprehensive, and directly answers the user's query.
159
+ - **0.0**: The response is completely irrelevant or fails to answer the query.
160
+ - **Output**: Return only the numerical score. Do not add any explanation or extra text.'''
161
+ precision_prompt = ChatPromptTemplate.from_messages([
162
+ ("system", system_message),
163
+ ("user", "Query: {query}\nResponse: {response}\n\nPrecision score:")
164
+ ])
165
+ chain = precision_prompt | llm | StrOutputParser()
166
+ precision_score = float(chain.invoke({
167
+ "query": state['query'],
168
+ "response": state['response'].content
169
+ }))
170
+ state['precision_score'] = precision_score
171
  state['precision_loop_count'] += 1
172
  return state
173
 
174
+ def refine_response(state: Dict) -> Dict:
175
+ print("---------refine_response---------")
176
+ 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.
177
+ Analyze the response for:
178
+ - **Gaps**: Is any crucial information from the query missing?
179
+ - **Ambiguities**: Are there any unclear or vague statements?
180
+ - **Inaccuracies**: Does the response contradict the user's intent (even if it's based on the context)?
181
+ - **Completeness**: Could the response be more thorough while remaining concise?
182
+ **Do not rewrite the response.** Instead, provide specific, actionable suggestions for improvement.'''
183
+ refine_response_prompt = ChatPromptTemplate.from_messages([
184
+ ("system", system_message),
185
+ ("user", "Query: {query}\\nResponse: {response}\\n\\n"
186
+ "What improvements can be made to enhance accuracy and completeness?")
187
+ ])
188
+ chain = refine_response_prompt | llm| StrOutputParser()
189
+ feedback = f"Previous Response: {state['response'].content}\\nSuggestions: {chain.invoke({'query': state['query'], 'response': state['response'].content})}"
190
+ state['feedback'] = feedback
191
+ return state
192
+
193
+ def refine_query(state: Dict) -> Dict:
194
+ print("---------refine_query---------")
195
+ 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.
196
+ Review the expanded query for:
197
+ - **Missing Keywords**: Are there essential terms or synonyms that should be added?
198
+ - **Lack of Specificity**: Could the query be narrowed down to a more precise topic?
199
+ - **Scope Refinements**: Is the query too broad or too narrow?
200
+ **Do not rewrite the query.** Instead, provide structured, actionable suggestions for improvement based on the original query's intent.'''
201
+ refine_query_prompt = ChatPromptTemplate.from_messages([
202
+ ("system", system_message),
203
+ ("user", "Original Query: {query}\\nExpanded Query: {expanded_query}\\n\\n"
204
+ "What improvements can be made for a better search?")
205
+ ])
206
+ chain = refine_query_prompt | llm | StrOutputParser()
207
+ query_feedback = f"Previous Expanded Query: {state['expanded_query']}\\nSuggestions: {chain.invoke({'query': state['query'], 'expanded_query': state['expanded_query']})}"
208
+ state['query_feedback'] = query_feedback
209
+ return state
210
 
211
+ # 3.4. Define Conditional Edges
212
  def should_continue_groundedness(state):
213
+ if state['groundedness_score'] >= 0.7:
214
+ return "check_precision"
215
+ else:
216
+ return "max_iterations_reached" if state["groundedness_loop_count"] >= state['loop_max_iter'] else "refine_response"
217
 
218
  def should_continue_precision(state: Dict) -> str:
219
+ if state['precision_score'] >= 0.7:
220
+ return "pass"
221
+ else:
222
+ return "max_iterations_reached" if state['precision_loop_count'] >= state['loop_max_iter'] else "refine_query"
223
 
224
  def max_iterations_reached(state: Dict) -> Dict:
225
+ state['response'] = "I'm unable to refine the response further. Please provide more context or clarify your question."
226
  return state
227
 
228
+ # 3.5. Assemble the Workflow Graph
229
  def create_workflow() -> StateGraph:
230
  workflow = StateGraph(AgentState)
231
+ workflow.add_node("expand_query", expand_query)
232
+ workflow.add_node("retrieve_context", retrieve_context)
233
+ workflow.add_node("craft_response", craft_response)
234
+ workflow.add_node("score_groundedness", score_groundedness)
235
+ workflow.add_node("refine_response", refine_response)
236
+ workflow.add_node("check_precision", check_precision)
237
+ workflow.add_node("refine_query", refine_query)
238
+ workflow.add_node("max_iterations_reached", max_iterations_reached)
239
+
240
+ workflow.add_edge(START, "expand_query")
241
+ workflow.add_edge("expand_query", "retrieve_context")
242
+ workflow.add_edge("retrieve_context", "craft_response")
243
+ workflow.add_edge("craft_response", "score_groundedness")
244
  workflow.add_conditional_edges("score_groundedness", should_continue_groundedness, {"check_precision": "check_precision", "refine_response": "refine_response", "max_iterations_reached": "max_iterations_reached"})
245
  workflow.add_edge("refine_response", "craft_response")
246
  workflow.add_conditional_edges("check_precision", should_continue_precision, {"pass": END, "refine_query": "refine_query", "max_iterations_reached": "max_iterations_reached"})
247
+ workflow.add_edge("refine_query", "expand_query")
248
+ workflow.add_edge("max_iterations_reached", END)
249
  return workflow
250
 
251
  WORKFLOW_APP = create_workflow().compile()
252
 
253
+ # 3.6. Create the Agentic RAG Tool
254
  @tool
255
  def agentic_rag(query: str):
256
+ """Runs the RAG-based agent for context-aware responses."""
257
+ inputs = {
258
+ "query": query, "expanded_query": "", "context": [], "response": "",
259
+ "precision_score": 0.0, "groundedness_score": 0.0,
260
+ "groundedness_loop_count": 0, "precision_loop_count": 0,
261
+ "feedback": "", "query_feedback": "", "loop_max_iter": 3
262
+ }
263
  output = WORKFLOW_APP.invoke(inputs)
264
  final_response = output.get('response')
265
+ if hasattr(final_response, 'content'):
266
+ return final_response.content
267
+ return str(final_response)
268
+
269
 
270
  # ==============================================================================
271
  # 4. SAFETY GUARDRAIL
272
  # ==============================================================================
273
+ # MODIFIED THIS SECTION TO FIX THE RUNTIME ERROR
274
+ llama_guard_client = Groq(
275
+ api_key=groq_api_key,
276
+ http_client=httpx.Client() # Manually pass a standard httpx client
277
+ )
278
  def filter_input_with_llama_guard(user_input, model="meta-llama/llama-guard-4-12b"):
279
  try:
280
+ response = llama_guard_client.chat.completions.create(
281
+ messages=[{"role": "user", "content": user_input}],
282
+ model=model,
283
+ temperature=0.0
284
+ )
285
  return response.choices[0].message.content.strip()
286
  except Exception as e:
287
+ print(f"Error with Llama Guard (Groq): {e}")
288
+ return "safe" # Fail-safe
289
+
290
 
291
  # ==============================================================================
292
  # 5. NUTRITION BOT CLASS (with Memory)
293
+ # This class encapsulates the agent, memory, and interaction logic.
294
  # ==============================================================================
295
  class NutritionBot:
296
  def __init__(self):
297
  self.memory = MemoryClient(api_key=mem0_api_key)
298
+ self.client = ChatOpenAI(
299
+ model_name="gpt-4o-mini",
300
+ openai_api_key=openai_api_key,
301
+ openai_api_base=openai_api_base,
302
+ temperature=0
303
+ )
304
  tools = [agentic_rag]
305
+ system_prompt = """You are a Medical Support Agent specializing ONLY in nutritional disorders...""" # (Your full, robust prompt here)
306
+ prompt = ChatPromptTemplate.from_messages([
307
+ ("system", system_prompt),
308
+ ("human", "{input}"),
309
+ ("placeholder", "{agent_scratchpad}")
310
+ ])
311
  agent = create_tool_calling_agent(self.client, tools, prompt)
312
  self.agent_executor = AgentExecutor(agent=agent, tools=tools, verbose=True)
313
 
314
  def store_customer_interaction(self, user_id: str, message: str, response: str, metadata: Dict = None):
315
+ if metadata is None: metadata = {}
316
+ metadata["timestamp"] = datetime.now().isoformat()
317
+ conversation = [{"role": "user", "content": message}, {"role": "assistant", "content": response}]
318
+ self.memory.add(messages=conversation, user_id=user_id, metadata=metadata)
319
 
320
  def get_relevant_history(self, user_id: str, query: str) -> str:
321
  memories = self.memory.search(query=query, user_id=user_id, limit=5)
322
  context = "Previous relevant interactions:\\n"
323
+ if not memories:
324
+ return "No previous relevant interactions found.\\n"
325
+ for mem in memories:
326
+ context += f"Summary of past interaction: {mem.get('memory', 'N/A')}\\n---\\n"
327
  return context
328
 
329
  def handle_customer_query(self, user_id: str, query: str) -> str:
330
  context = self.get_relevant_history(user_id, query)
331
+ prompt = f"Context:\\n{context}\\nCurrent customer query: {query}\\nProvide a helpful response that takes into account any relevant past interactions."
332
  response = self.agent_executor.invoke({"input": prompt})
333
  self.store_customer_interaction(user_id=user_id, message=query, response=response["output"])
334
  return response['output']
335
 
336
+
337
  # ==============================================================================
338
  # 6. STREAMLIT UI
339
+ # This is the entry point and user interface for the application.
340
  # ==============================================================================
341
  def nutrition_disorder_streamlit():
342
  st.title("Nutrition Disorder Specialist")
343
  st.write("Ask me anything about nutrition disorders, symptoms, causes, treatments, and more.")
344
+
345
+ if 'chat_history' not in st.session_state:
346
+ st.session_state.chat_history = []
347
+ if 'user_id' not in st.session_state:
348
+ st.session_state.user_id = None
349
+ if 'chatbot' not in st.session_state:
350
+ st.session_state.chatbot = None
351
 
352
  if st.session_state.user_id is None:
353
  with st.form("login_form"):
354
  user_id = st.text_input("Please enter your name to begin:")
355
+ submit_button = st.form_submit_button("Login")
356
+ if submit_button and user_id:
357
  st.session_state.user_id = user_id
358
  st.session_state.chatbot = NutritionBot()
359
+ welcome_msg = f"Welcome, {user_id}! How can I help you with nutrition disorders today?"
360
+ st.session_state.chat_history.append({"role": "assistant", "content": welcome_msg})
361
  st.rerun()
362
  else:
363
  for message in st.session_state.chat_history:
364
+ with st.chat_message(message["role"]):
365
+ st.write(message["content"])
366
+
367
  if user_query := st.chat_input("Ask about a nutrition disorder..."):
368
  st.session_state.chat_history.append({"role": "user", "content": user_query})
369
+ with st.chat_message("user"):
370
+ st.write(user_query)
371
+
372
+ filtered_result = filter_input_with_llama_guard(user_query)
373
+ if "safe" in filtered_result:
374
  try:
375
  with st.spinner("Thinking..."):
376
  response = st.session_state.chatbot.handle_customer_query(st.session_state.user_id, user_query)
377
  st.session_state.chat_history.append({"role": "assistant", "content": response})
378
+ with st.chat_message("assistant"):
379
+ st.write(response)
380
  except Exception as e:
381
  error_msg = f"Sorry, I encountered an error: {e}"
382
  st.session_state.chat_history.append({"role": "assistant", "content": error_msg})
383
+ with st.chat_message("assistant"):
384
+ st.write(error_msg)
385
  else:
386
+ inappropriate_msg = "I apologize, but I cannot process that input as it may be inappropriate."
387
  st.session_state.chat_history.append({"role": "assistant", "content": inappropriate_msg})
388
+ with st.chat_message("assistant"):
389
+ st.write(inappropriate_msg)
390
 
391
  if __name__ == "__main__":
392
  nutrition_disorder_streamlit()