NightmareFox12 commited on
Commit
d8caf2b
·
verified ·
1 Parent(s): da7a2a6

Update rag_api.py

Browse files
Files changed (1) hide show
  1. rag_api.py +125 -125
rag_api.py CHANGED
@@ -1,7 +1,7 @@
1
  import os
2
  import requests
3
- import shutil
4
- from langchain_community.vectorstores import FAISS
5
  from fastapi import FastAPI
6
  from pydantic import BaseModel
7
  from langchain_huggingface import HuggingFaceEmbeddings
@@ -9,183 +9,183 @@ from langchain_core.runnables import RunnablePassthrough
9
  from langchain_core.prompts import PromptTemplate
10
  from langchain_groq import ChatGroq
11
 
12
-
13
  # --------------------------------------------------------
14
- # ESTABLECER LA RUTA DEL CACHÉ A /tmp, donde sí hay permisos
15
- # Esto resuelve el PermissionError: [Errno 13] Permission denied: '/.cache'
16
  # --------------------------------------------------------
17
  TEMP_CACHE_DIR = '/tmp/huggingface_cache'
18
- os.environ['TRANSFORMERS_CACHE'] = TEMP_CACHE_DIR
19
  os.environ['HF_HOME'] = TEMP_CACHE_DIR
20
  os.environ['SENTENCE_TRANSFORMERS_HOME'] = TEMP_CACHE_DIR
21
- os.makedirs(TEMP_CACHE_DIR, exist_ok=True) # Aseguramos que el directorio exista
22
 
23
  # --------------------------------------------------------
24
  # 1. CONFIGURACIÓN
25
  # --------------------------------------------------------
26
- # URLs DE DESCARGA DIRECTA DE GOOGLE DRIVE (construidas con las IDs)
27
  URL_FAISS = "https://drive.google.com/uc?export=download&id=1XqImFIKiuRDhSDK6Rm6dAZbHm03NdzQa"
28
- URL_PKL = "https://drive.google.com/uc?export=download&id=156BWHHGi-JuD9EM2Nek1mNcyitivQWAH"
29
-
30
- # Directorio donde guardaremos los archivos de la BD FAISS dentro del contenedor Docker
31
- DOWNLOAD_DIR = "/tmp/db_faiss"
32
- DB_FAISS_PATH = DOWNLOAD_DIR
33
 
34
  # --------------------------------------------------------
35
- # 2. FUNCIONES DE DESCARGA Y CARGA DEL RAG CORE
36
  # --------------------------------------------------------
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
38
  class QueryRequest(BaseModel):
39
- """Define el formato de la pregunta que recibirá el endpoint /query."""
40
  query: str
41
 
42
  def download_file(url, local_path):
43
- """Descarga un archivo desde una URL y lo guarda localmente en /tmp."""
44
  file_name = os.path.basename(local_path)
45
- print(f"Descargando: {file_name} desde la nube...")
46
-
47
- # Usar headers de agente de usuario para evitar bloqueos 403 de Google Drive
48
- headers = {'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36'}
49
-
50
  try:
51
  response = requests.get(url, stream=True, headers=headers, timeout=30)
52
-
53
- # Manejar errores de descarga (ej. si el archivo no es público)
54
  if response.status_code == 403:
55
- raise PermissionError(f"Error 403: El archivo {file_name} no es público en Google Drive.")
56
- response.raise_for_status()
57
-
58
- # Asegurar que el directorio de destino exista
59
  os.makedirs(os.path.dirname(local_path), exist_ok=True)
60
-
61
  with open(local_path, 'wb') as f:
62
  shutil.copyfileobj(response.raw, f)
63
- print(f"Descarga de {file_name} completada.")
64
-
65
  except requests.exceptions.RequestException as e:
66
- print(f"Error en la descarga de {file_name}: {e}")
67
- # Si la descarga falla, volvemos a lanzar un error para detener la inicialización.
68
- raise RuntimeError(f"Fallo de conexión o timeout al descargar {file_name}: {e}")
69
 
70
  def load_and_configure_rag():
71
- """
72
- Descarga la base de datos FAISS de la nube, carga el modelo de embeddings
73
- y configura la cadena de RAG.
74
- """
75
  try:
76
- # 1. Descargar los archivos y guardarlos en /tmp/db_faiss
77
  download_file(URL_FAISS, os.path.join(DOWNLOAD_DIR, 'index.faiss'))
78
- download_file(URL_PKL, os.path.join(DOWNLOAD_DIR, 'index.pkl'))
79
-
80
- # 2. Cargar Embeddings (el mismo modelo que usaste)
81
- # Se pasa 'cache_folder' para asegurar que el modelo se guarde en /tmp si se descarga
82
- print("Cargando modelo de embeddings (all-MiniLM-L6-v2)...")
83
  embeddings = HuggingFaceEmbeddings(
84
  model_name="sentence-transformers/all-MiniLM-L6-v2",
85
- model_kwargs={'device': 'cpu'}, # Forzar a usar CPU si no hay GPU, estándar en Spaces
86
- cache_folder=TEMP_CACHE_DIR # Se pasa la ruta de caché directamente
87
  )
88
- print("Modelo de embeddings cargado con éxito.")
89
-
90
- # 3. Cargar Vector Store (desde la ruta temporal /tmp/db_faiss)
91
- print(f"Cargando FAISS desde: {DB_FAISS_PATH}...")
92
  vectorstore = FAISS.load_local(
93
- DB_FAISS_PATH,
94
- embeddings,
95
- allow_dangerous_deserialization=True
96
  )
97
- print("Vector Store FAISS cargada con éxito.")
98
-
99
- # 4. Configurar LLM y Prompt
100
- # Nota: Asegúrate de tener la variable de entorno GROQ_API_KEY configurada en tu Space.
101
- llm_groq = ChatGroq(temperature=0.0, model_name="llama-3.1-8b-instant")
102
-
103
- custom_prompt = """
104
- Eres un asistente de preguntas y respuestas experto en la documentación de NutriActive.
105
- Tu tarea es responder a la pregunta del usuario basándote EXCLUSIVAMENTE en el contexto proporcionado.
106
- Si la respuesta no se encuentra en el contexto, indica amablemente: "Lo siento, la información que buscas no se encuentra en la documentación de NutriActive."
107
-
108
- Contexto: {context}
109
- Pregunta: {question}
110
-
111
- Respuesta concisa:
112
- """
113
- RAG_PROMPT = PromptTemplate(template=custom_prompt, input_variables=["context", "question"])
114
-
115
-
116
- # 5. Crear la cadena de RAG
117
- # qa_chain = RetrievalQA.from_chain_type(
118
- # llm=llm_groq,
119
- # chain_type="stuff",
120
- # retriever=vectorstore.as_retriever(search_kwargs={"k": 3}),
121
- # return_source_documents=True,
122
- # chain_type_kwargs={"prompt": RAG_PROMPT}
123
- # )
124
-
125
- # Con LCEL defines el pipeline
126
- retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
127
-
128
- qa_chain = (
129
  {"context": retriever, "question": RunnablePassthrough()}
130
  | RAG_PROMPT
131
- | llm_groq
132
  )
133
- return qa_chain,retriever
134
-
 
135
  except Exception as e:
136
- # Imprime un mensaje de error claro en los logs
137
- print("-------------------------------------------------------------------------")
138
- print(f"Error CRÍTICO al inicializar el RAG: {type(e).__name__}: {e}")
139
- print("Asegúrate de que los archivos de Google Drive sean públicos y que tengas 'langchain-huggingface' en requirements.txt.")
140
- print("-------------------------------------------------------------------------")
141
- # Esta excepción detendrá la carga de FastAPI
142
- raise RuntimeError(f"Falla al cargar FAISS o Embeddings: {e}")
143
 
144
  # --------------------------------------------------------
145
- # 3. CONFIGURACIÓN DE FASTAPI Y ENDPOINTS
146
  # --------------------------------------------------------
147
-
148
- # Iniciar servidor y RAG (manejo de errores de carga)
149
  app = FastAPI(title="NutriActive RAG API")
150
- qa_chain = None # Inicializamos a None
 
151
 
152
  try:
153
- qa_chain,retriever = load_and_configure_rag()
154
  except RuntimeError:
155
- # Si load,_and_configure_rag falla con RuntimeError, qa_chain se queda en None,
156
- # y el endpoint / lo reportará. No es necesario hacer nada más aquí.
157
  pass
158
 
159
  @app.get("/")
160
  def home():
161
- """Verifica que el servidor está corriendo."""
162
  if qa_chain is None:
163
- return {"error": "El servidor está activo, pero el RAG no se pudo inicializar. Revisa los logs de inicio para ver el error de descarga o carga de embeddings."}
164
- return {"message": "API de NutriActive RAG operativa. Usa el endpoint /query."}
165
 
166
  @app.post("/query")
167
  async def process_query(request: QueryRequest):
168
- """Endpoint principal para recibir la pregunta y devolver la respuesta."""
169
  if qa_chain is None:
170
- return {"error": "El sistema RAG no se pudo cargar. Revisa los logs de inicio para ver el error específico."}
171
-
172
  try:
173
- # Aquí usamos .invoke() para LangChain Expression Language (LCEL)
174
- # result = qa_chain.invoke({"query": request.query})
175
- # sources = [doc.metadata.get('source', 'N/A') for doc in result['source_documents']]
176
-
177
- # 1. Ejecutar el pipeline LCEL → devuelve solo la respuesta
178
- response = qa_chain.invoke(request.query)
179
-
180
- # 2. Obtener las fuentes directamente del retriever
181
- docs = retriever.invoke(request.query)
182
- sources = [doc.metadata.get("source", "N/A") for doc in docs]
183
-
184
- return {
185
- "query": request.query,
186
- "response": response,
187
- # "response": result['result'],
188
- "sources": sources
189
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
190
  except Exception as e:
191
- return {"error": f"Error al procesar la consulta: {e}"}
 
1
  import os
2
  import requests
3
+ import shutil
4
+ from langchain_community.vectorstores import FAISS
5
  from fastapi import FastAPI
6
  from pydantic import BaseModel
7
  from langchain_huggingface import HuggingFaceEmbeddings
 
9
  from langchain_core.prompts import PromptTemplate
10
  from langchain_groq import ChatGroq
11
 
 
12
  # --------------------------------------------------------
13
+ # CACHÉ EN /tmp
 
14
  # --------------------------------------------------------
15
  TEMP_CACHE_DIR = '/tmp/huggingface_cache'
16
+ os.environ['TRANSFORMERS_CACHE'] = TEMP_CACHE_DIR
17
  os.environ['HF_HOME'] = TEMP_CACHE_DIR
18
  os.environ['SENTENCE_TRANSFORMERS_HOME'] = TEMP_CACHE_DIR
19
+ os.makedirs(TEMP_CACHE_DIR, exist_ok=True)
20
 
21
  # --------------------------------------------------------
22
  # 1. CONFIGURACIÓN
23
  # --------------------------------------------------------
 
24
  URL_FAISS = "https://drive.google.com/uc?export=download&id=1XqImFIKiuRDhSDK6Rm6dAZbHm03NdzQa"
25
+ URL_PKL = "https://drive.google.com/uc?export=download&id=156BWHHGi-JuD9EM2Nek1mNcyitivQWAH"
26
+ DOWNLOAD_DIR = "/tmp/db_faiss"
27
+ DB_FAISS_PATH = DOWNLOAD_DIR
 
 
28
 
29
  # --------------------------------------------------------
30
+ # 2. CLASIFICADOR DE INTENCIÓN ← NUEVO
31
  # --------------------------------------------------------
32
+ INTENT_PROMPT = PromptTemplate(
33
+ template="""Eres un clasificador de intenciones para un asistente de nutrición llamado NutriActive.
34
+
35
+ Analiza el mensaje del usuario y clasifícalo en UNA de estas categorías:
36
+ - SALUDO: saludos, despedidas, conversación casual ("hola", "gracias", "adiós", "¿cómo estás?")
37
+ - NUTRICION: preguntas sobre nutrición, dieta, salud, cálculos como IMC, calorías, macros, alimentos, etc.
38
+ - OTRO: preguntas no relacionadas con nutrición ni saludos
39
+
40
+ Responde SOLO con la categoría, sin explicación.
41
+
42
+ Mensaje: {query}
43
+ Categoría:""",
44
+ input_variables=["query"]
45
+ )
46
+
47
+ SALUDO_PROMPT = PromptTemplate(
48
+ template="""Eres NutriActive, un asistente amigable especializado en nutrición y salud.
49
+ Responde de forma natural y cálida al siguiente mensaje casual del usuario.
50
+ Si el usuario se despide o agradece, invítalo a preguntar sobre nutrición.
51
 
52
+ Mensaje: {query}
53
+ Respuesta:""",
54
+ input_variables=["query"]
55
+ )
56
+
57
+ RAG_PROMPT = PromptTemplate(
58
+ template="""Eres NutriActive, un asistente experto en nutrición y salud.
59
+ Tu tarea es responder basándote en el contexto proporcionado.
60
+ Si el contexto no tiene suficiente información, usa tu conocimiento general sobre nutrición para dar una respuesta útil.
61
+ Sé amigable, claro y conciso.
62
+
63
+ Contexto de la base de datos: {context}
64
+ Pregunta del usuario: {question}
65
+
66
+ Respuesta:""",
67
+ input_variables=["context", "question"]
68
+ )
69
+
70
+ # --------------------------------------------------------
71
+ # 3. FUNCIONES DE DESCARGA Y CARGA
72
+ # --------------------------------------------------------
73
  class QueryRequest(BaseModel):
 
74
  query: str
75
 
76
  def download_file(url, local_path):
 
77
  file_name = os.path.basename(local_path)
78
+ print(f"Descargando: {file_name}...")
79
+ headers = {'User-Agent': 'Mozilla/5.0'}
 
 
 
80
  try:
81
  response = requests.get(url, stream=True, headers=headers, timeout=30)
 
 
82
  if response.status_code == 403:
83
+ raise PermissionError(f"Error 403: {file_name} no es público.")
84
+ response.raise_for_status()
 
 
85
  os.makedirs(os.path.dirname(local_path), exist_ok=True)
 
86
  with open(local_path, 'wb') as f:
87
  shutil.copyfileobj(response.raw, f)
88
+ print(f"�� {file_name} descargado.")
 
89
  except requests.exceptions.RequestException as e:
90
+ raise RuntimeError(f"Fallo al descargar {file_name}: {e}")
 
 
91
 
92
  def load_and_configure_rag():
 
 
 
 
93
  try:
 
94
  download_file(URL_FAISS, os.path.join(DOWNLOAD_DIR, 'index.faiss'))
95
+ download_file(URL_PKL, os.path.join(DOWNLOAD_DIR, 'index.pkl'))
96
+
97
+ print("Cargando embeddings...")
 
 
98
  embeddings = HuggingFaceEmbeddings(
99
  model_name="sentence-transformers/all-MiniLM-L6-v2",
100
+ model_kwargs={'device': 'cpu'},
101
+ cache_folder=TEMP_CACHE_DIR
102
  )
103
+
104
+ print("Cargando FAISS...")
 
 
105
  vectorstore = FAISS.load_local(
106
+ DB_FAISS_PATH, embeddings, allow_dangerous_deserialization=True
 
 
107
  )
108
+
109
+ llm = ChatGroq(temperature=0.3, model_name="llama-3.3-70b-versatile")
110
+
111
+ # Cadena clasificadora de intención
112
+ intent_chain = INTENT_PROMPT | llm
113
+
114
+ # Cadena para saludos
115
+ saludo_chain = SALUDO_PROMPT | llm
116
+
117
+ # Cadena RAG principal
118
+ retriever = vectorstore.as_retriever(search_kwargs={"k": 4})
119
+ rag_chain = (
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
  {"context": retriever, "question": RunnablePassthrough()}
121
  | RAG_PROMPT
122
+ | llm
123
  )
124
+
125
+ return intent_chain, saludo_chain, rag_chain, retriever
126
+
127
  except Exception as e:
128
+ print(f"Error CRÍTICO al inicializar: {type(e).__name__}: {e}")
129
+ raise RuntimeError(f"Falla al cargar RAG: {e}")
 
 
 
 
 
130
 
131
  # --------------------------------------------------------
132
+ # 4. FASTAPI
133
  # --------------------------------------------------------
 
 
134
  app = FastAPI(title="NutriActive RAG API")
135
+
136
+ intent_chain = saludo_chain = qa_chain = retriever = None
137
 
138
  try:
139
+ intent_chain, saludo_chain, qa_chain, retriever = load_and_configure_rag()
140
  except RuntimeError:
 
 
141
  pass
142
 
143
  @app.get("/")
144
  def home():
 
145
  if qa_chain is None:
146
+ return {"error": "RAG no inicializado. Revisa los logs."}
147
+ return {"message": "API NutriActive operativa. Usa /query."}
148
 
149
  @app.post("/query")
150
  async def process_query(request: QueryRequest):
 
151
  if qa_chain is None:
152
+ return {"error": "El sistema RAG no se pudo cargar."}
153
+
154
  try:
155
+ # ── 1. Clasificar intención ──────────────────────────────
156
+ intent_result = intent_chain.invoke({"query": request.query})
157
+ intent = intent_result.content.strip().upper()
158
+ print(f"[Intent] '{request.query}' → {intent}")
159
+
160
+ # ── 2. Ruta según intención ──────────────────────────────
161
+ if "SALUDO" in intent:
162
+ respuesta = saludo_chain.invoke({"query": request.query})
163
+ return {
164
+ "query": request.query,
165
+ "response": respuesta.content,
166
+ "intent": "SALUDO",
167
+ "sources": []
168
+ }
169
+
170
+ elif "OTRO" in intent:
171
+ return {
172
+ "query": request.query,
173
+ "response": "Soy NutriActive, especializado en nutrición y salud. ¿Tienes alguna pregunta sobre alimentación, dietas o bienestar? 🥗",
174
+ "intent": "OTRO",
175
+ "sources": []
176
+ }
177
+
178
+ else:
179
+ # NUTRICION → RAG completo
180
+ respuesta = qa_chain.invoke(request.query)
181
+ docs = retriever.invoke(request.query)
182
+ sources = [doc.metadata.get("source", "N/A") for doc in docs]
183
+ return {
184
+ "query": request.query,
185
+ "response": respuesta.content,
186
+ "intent": "NUTRICION",
187
+ "sources": sources
188
+ }
189
+
190
  except Exception as e:
191
+ return {"error": f"Error al procesar la consulta: {e}"}