NightmareFox12 commited on
Commit
2c7b13b
·
verified ·
1 Parent(s): c6f5f47

Update rag_api.py

Browse files
Files changed (1) hide show
  1. rag_api.py +64 -26
rag_api.py CHANGED
@@ -3,15 +3,23 @@ import requests
3
  import shutil
4
  from fastapi import FastAPI
5
  from pydantic import BaseModel
 
 
6
  from langchain_community.vectorstores import FAISS
7
- from langchain_community.embeddings import HuggingFaceEmbeddings
8
  from langchain.chains import RetrievalQA
9
  from langchain.prompts import PromptTemplate
10
  from langchain_groq import ChatGroq
11
 
 
12
  # ESTABLECER LA RUTA DEL CACHÉ A /tmp, donde sí hay permisos
13
- os.environ['TRANSFORMERS_CACHE'] = '/tmp/huggingface_cache'
14
- os.environ['HF_HOME'] = '/tmp/huggingface_cache'
 
 
 
 
 
15
 
16
  # --------------------------------------------------------
17
  # 1. CONFIGURACIÓN
@@ -20,7 +28,7 @@ os.environ['HF_HOME'] = '/tmp/huggingface_cache'
20
  URL_FAISS = "https://drive.google.com/uc?export=download&id=1bFLDqk0fEsdJlxjYPnxIqOUxagO7qcy3"
21
  URL_PKL = "https://drive.google.com/uc?export=download&id=1D0JGeRft3798x-rsTam1s_2lVhjC9yTR"
22
 
23
- # Directorio donde guardaremos los archivos dentro del contenedor Docker
24
  DOWNLOAD_DIR = "/tmp/db_faiss"
25
  DB_FAISS_PATH = DOWNLOAD_DIR
26
 
@@ -34,26 +42,36 @@ class QueryRequest(BaseModel):
34
 
35
  def download_file(url, local_path):
36
  """Descarga un archivo desde una URL y lo guarda localmente en /tmp."""
37
- print(f"Descargando: {os.path.basename(local_path)} desde la nube...")
38
- # Usar headers de agente de usuario para evitar bloqueos 403 de Google Drive
39
- headers = {'User-Agent': 'Mozilla/5.0'}
40
- response = requests.get(url, stream=True, headers=headers)
41
 
42
- # Manejar errores de descarga (ej. si el archivo no es público)
43
- if response.status_code == 403:
44
- raise PermissionError(f"Error 403: El archivo {os.path.basename(local_path)} no es público en Google Drive.")
45
- response.raise_for_status()
46
-
47
- # Asegurar que el directorio de destino exista
48
- os.makedirs(os.path.dirname(local_path), exist_ok=True)
49
 
50
- with open(local_path, 'wb') as f:
51
- shutil.copyfileobj(response.raw, f)
52
- print("Descarga completada.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
53
 
54
  def load_and_configure_rag():
55
  """
56
- Descarga la base de datos FAISS de la nube y la configura.
 
57
  """
58
  try:
59
  # 1. Descargar los archivos y guardarlos en /tmp/db_faiss
@@ -61,16 +79,26 @@ def load_and_configure_rag():
61
  download_file(URL_PKL, os.path.join(DOWNLOAD_DIR, 'index.pkl'))
62
 
63
  # 2. Cargar Embeddings (el mismo modelo que usaste)
64
- embeddings = HuggingFaceEmbeddings(model_name="sentence-transformers/all-MiniLM-L6-v2")
 
 
 
 
 
 
 
65
 
66
  # 3. Cargar Vector Store (desde la ruta temporal /tmp/db_faiss)
 
67
  vectorstore = FAISS.load_local(
68
  DB_FAISS_PATH,
69
  embeddings,
70
  allow_dangerous_deserialization=True
71
  )
 
72
 
73
  # 4. Configurar LLM y Prompt
 
74
  llm_groq = ChatGroq(temperature=0.0, model_name="llama-3.1-8b-instant")
75
 
76
  custom_prompt = """
@@ -96,9 +124,13 @@ def load_and_configure_rag():
96
  return qa_chain
97
 
98
  except Exception as e:
99
- print(f"Error CRÍTICO al descargar/cargar FAISS desde la nube: {e}")
100
- # Esta es la excepción que verás en los logs si falla la descarga
101
- raise RuntimeError(f"Falla al cargar FAISS: {e}")
 
 
 
 
102
 
103
  # --------------------------------------------------------
104
  # 3. CONFIGURACIÓN DE FASTAPI Y ENDPOINTS
@@ -106,16 +138,20 @@ def load_and_configure_rag():
106
 
107
  # Iniciar servidor y RAG (manejo de errores de carga)
108
  app = FastAPI(title="NutriActive RAG API")
 
 
109
  try:
110
  qa_chain = load_and_configure_rag()
111
  except RuntimeError:
112
- qa_chain = None
 
 
113
 
114
  @app.get("/")
115
  def home():
116
  """Verifica que el servidor está corriendo."""
117
  if qa_chain is None:
118
- return {"error": "El servidor está activo, pero el RAG no se pudo inicializar. Revisa los logs de inicio para ver el error de descarga."}
119
  return {"message": "API de NutriActive RAG operativa. Usa el endpoint /query."}
120
 
121
  @app.post("/query")
@@ -125,6 +161,7 @@ async def process_query(request: QueryRequest):
125
  return {"error": "El sistema RAG no se pudo cargar. Revisa los logs de inicio para ver el error específico."}
126
 
127
  try:
 
128
  result = qa_chain.invoke({"query": request.query})
129
  sources = [doc.metadata.get('source', 'N/A') for doc in result['source_documents']]
130
 
@@ -134,4 +171,5 @@ async def process_query(request: QueryRequest):
134
  "sources": sources
135
  }
136
  except Exception as e:
137
- return {"error": f"Error al procesar la consulta: {e}"}
 
 
3
  import shutil
4
  from fastapi import FastAPI
5
  from pydantic import BaseModel
6
+ # Nota: La importación de LangChain ha sido actualizada para usar los paquetes comunitarios
7
+ # y el nuevo paquete langchain_huggingface para evitar la advertencia de deprecación.
8
  from langchain_community.vectorstores import FAISS
9
+ from langchain_huggingface import HuggingFaceEmbeddings # Importación recomendada
10
  from langchain.chains import RetrievalQA
11
  from langchain.prompts import PromptTemplate
12
  from langchain_groq import ChatGroq
13
 
14
+ # --------------------------------------------------------
15
  # ESTABLECER LA RUTA DEL CACHÉ A /tmp, donde sí hay permisos
16
+ # Esto resuelve el PermissionError: [Errno 13] Permission denied: '/.cache'
17
+ # --------------------------------------------------------
18
+ TEMP_CACHE_DIR = '/tmp/huggingface_cache'
19
+ os.environ['TRANSFORMERS_CACHE'] = TEMP_CACHE_DIR
20
+ os.environ['HF_HOME'] = TEMP_CACHE_DIR
21
+ os.environ['SENTENCE_TRANSFORMERS_HOME'] = TEMP_CACHE_DIR
22
+ os.makedirs(TEMP_CACHE_DIR, exist_ok=True) # Aseguramos que el directorio exista
23
 
24
  # --------------------------------------------------------
25
  # 1. CONFIGURACIÓN
 
28
  URL_FAISS = "https://drive.google.com/uc?export=download&id=1bFLDqk0fEsdJlxjYPnxIqOUxagO7qcy3"
29
  URL_PKL = "https://drive.google.com/uc?export=download&id=1D0JGeRft3798x-rsTam1s_2lVhjC9yTR"
30
 
31
+ # Directorio donde guardaremos los archivos de la BD FAISS dentro del contenedor Docker
32
  DOWNLOAD_DIR = "/tmp/db_faiss"
33
  DB_FAISS_PATH = DOWNLOAD_DIR
34
 
 
42
 
43
  def download_file(url, local_path):
44
  """Descarga un archivo desde una URL y lo guarda localmente en /tmp."""
45
+ file_name = os.path.basename(local_path)
46
+ print(f"Descargando: {file_name} desde la nube...")
 
 
47
 
48
+ # Usar headers de agente de usuario para evitar bloqueos 403 de Google Drive
49
+ 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'}
 
 
 
 
 
50
 
51
+ try:
52
+ response = requests.get(url, stream=True, headers=headers, timeout=30)
53
+
54
+ # Manejar errores de descarga (ej. si el archivo no es público)
55
+ if response.status_code == 403:
56
+ raise PermissionError(f"Error 403: El archivo {file_name} no es público en Google Drive.")
57
+ response.raise_for_status()
58
+
59
+ # Asegurar que el directorio de destino exista
60
+ os.makedirs(os.path.dirname(local_path), exist_ok=True)
61
+
62
+ with open(local_path, 'wb') as f:
63
+ shutil.copyfileobj(response.raw, f)
64
+ print(f"Descarga de {file_name} completada.")
65
+
66
+ except requests.exceptions.RequestException as e:
67
+ print(f"Error en la descarga de {file_name}: {e}")
68
+ # Si la descarga falla, volvemos a lanzar un error para detener la inicialización.
69
+ raise RuntimeError(f"Fallo de conexión o timeout al descargar {file_name}: {e}")
70
 
71
  def load_and_configure_rag():
72
  """
73
+ Descarga la base de datos FAISS de la nube, carga el modelo de embeddings
74
+ y configura la cadena de RAG.
75
  """
76
  try:
77
  # 1. Descargar los archivos y guardarlos en /tmp/db_faiss
 
79
  download_file(URL_PKL, os.path.join(DOWNLOAD_DIR, 'index.pkl'))
80
 
81
  # 2. Cargar Embeddings (el mismo modelo que usaste)
82
+ # Se pasa 'cache_folder' para asegurar que el modelo se guarde en /tmp si se descarga
83
+ print("Cargando modelo de embeddings (all-MiniLM-L6-v2)...")
84
+ embeddings = HuggingFaceEmbeddings(
85
+ model_name="sentence-transformers/all-MiniLM-L6-v2",
86
+ model_kwargs={'device': 'cpu'}, # Forzar a usar CPU si no hay GPU, estándar en Spaces
87
+ cache_folder=TEMP_CACHE_DIR # Se pasa la ruta de caché directamente
88
+ )
89
+ print("Modelo de embeddings cargado con éxito.")
90
 
91
  # 3. Cargar Vector Store (desde la ruta temporal /tmp/db_faiss)
92
+ print(f"Cargando FAISS desde: {DB_FAISS_PATH}...")
93
  vectorstore = FAISS.load_local(
94
  DB_FAISS_PATH,
95
  embeddings,
96
  allow_dangerous_deserialization=True
97
  )
98
+ print("Vector Store FAISS cargada con éxito.")
99
 
100
  # 4. Configurar LLM y Prompt
101
+ # Nota: Asegúrate de tener la variable de entorno GROQ_API_KEY configurada en tu Space.
102
  llm_groq = ChatGroq(temperature=0.0, model_name="llama-3.1-8b-instant")
103
 
104
  custom_prompt = """
 
124
  return qa_chain
125
 
126
  except Exception as e:
127
+ # Imprime un mensaje de error claro en los logs
128
+ print("-------------------------------------------------------------------------")
129
+ print(f"Error CRÍTICO al inicializar el RAG: {type(e).__name__}: {e}")
130
+ print("Asegúrate de que los archivos de Google Drive sean públicos y que tengas 'langchain-huggingface' en requirements.txt.")
131
+ print("-------------------------------------------------------------------------")
132
+ # Esta excepción detendrá la carga de FastAPI
133
+ raise RuntimeError(f"Falla al cargar FAISS o Embeddings: {e}")
134
 
135
  # --------------------------------------------------------
136
  # 3. CONFIGURACIÓN DE FASTAPI Y ENDPOINTS
 
138
 
139
  # Iniciar servidor y RAG (manejo de errores de carga)
140
  app = FastAPI(title="NutriActive RAG API")
141
+ qa_chain = None # Inicializamos a None
142
+
143
  try:
144
  qa_chain = load_and_configure_rag()
145
  except RuntimeError:
146
+ # Si load_and_configure_rag falla con RuntimeError, qa_chain se queda en None,
147
+ # y el endpoint / lo reportará. No es necesario hacer nada más aquí.
148
+ pass
149
 
150
  @app.get("/")
151
  def home():
152
  """Verifica que el servidor está corriendo."""
153
  if qa_chain is None:
154
+ 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."}
155
  return {"message": "API de NutriActive RAG operativa. Usa el endpoint /query."}
156
 
157
  @app.post("/query")
 
161
  return {"error": "El sistema RAG no se pudo cargar. Revisa los logs de inicio para ver el error específico."}
162
 
163
  try:
164
+ # Aquí usamos .invoke() para LangChain Expression Language (LCEL)
165
  result = qa_chain.invoke({"query": request.query})
166
  sources = [doc.metadata.get('source', 'N/A') for doc in result['source_documents']]
167
 
 
171
  "sources": sources
172
  }
173
  except Exception as e:
174
+ # En caso de error durante la ejecución (no durante el setup)
175
+ return {"error": f"Error al procesar la consulta: {e}"}