feat: Implemented memory persistent

This commit is contained in:
2026-03-02 22:55:24 -05:00
parent 390fd120f9
commit aab046b2ef
6 changed files with 90 additions and 18 deletions

View File

@@ -19,4 +19,4 @@ services:
- .env - .env
volumes: volumes:
data: data:

View File

@@ -1,8 +1,9 @@
from langgraph.graph import StateGraph, START, END from langgraph.graph import StateGraph, START, END
from langchain_core.messages import SystemMessage, ToolMessage, BaseMessage from langchain_core.messages import SystemMessage, ToolMessage, BaseMessage
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.memory import InMemorySaver from langgraph.checkpoint.memory import InMemorySaver
from langchain_core.runnables import RunnableConfig from langchain_core.runnables import RunnableConfig
from typing import Literal, Callable, Any from typing import Literal, Callable, Any, Union
from .schemas import MessagesState from .schemas import MessagesState
from dataclasses import dataclass from dataclasses import dataclass
import logging import logging
@@ -30,7 +31,7 @@ class Agent:
self, self,
model: Any, model: Any,
config: AgentConfig | None = None, config: AgentConfig | None = None,
checkpointer: InMemorySaver | None = None, checkpointer: Union[BaseCheckpointSaver, None] = None,
tools: list | None = None tools: list | None = None
): ):
""" """
@@ -72,6 +73,19 @@ class Agent:
"""Expone el grafo para inspección en tests.""" """Expone el grafo para inspección en tests."""
return self._build_agent() return self._build_agent()
async def ainvoke(self, state: dict, config: RunnableConfig | None = None) -> dict:
"""
Versión asíncrona de invoke, necesaria para checkpointers de DB.
"""
if not config:
config: RunnableConfig = {"configurable": {"thread_id": "default"}}
if "messages" not in state:
raise ValueError("State must contain 'messages' key")
# Llamamos al método ainvoke del grafo compilado
return await self._compiled_agent.ainvoke(state, config=config)
def invoke(self, state: dict, config: RunnableConfig | None = None) -> dict: def invoke(self, state: dict, config: RunnableConfig | None = None) -> dict:
""" """
Ejecuta el agente con el estado inicial. Ejecuta el agente con el estado inicial.

View File

@@ -4,22 +4,34 @@ from naliiabot.bot.agent.agent import Agent, AgentConfig
from naliiabot.bot.tools.naliia_tools import NaliiaTools from naliiabot.bot.tools.naliia_tools import NaliiaTools
from naliiabot.bot.factories.llm_factory import LLMFactory from naliiabot.bot.factories.llm_factory import LLMFactory
from naliiabot.bot.prompts.load_prompt import get_prompt_template from naliiabot.bot.prompts.load_prompt import get_prompt_template
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from psycopg_pool import AsyncConnectionPool
from psycopg.rows import dict_row
from ..logger import logger from ..logger import logger
from .settings import _settings
MODEL_NAME = os.getenv("LLM_MODEL", "anthropic")
def get_agent(): _connection_string = f"postgresql://{_settings.DB_USER}:{_settings.DB_PASSWORD}@localhost:5432/{_settings.DB_NAME}"
agent_instance: Agent | None = None
db_pool = None
async def get_agent() -> Agent:
""" """
Dependency function to get an instance of the Agent class. Dependencia simplificada que devuelve la instancia ya construida.
Returns:
Agent: An instance of the Agent class.
""" """
if agent_instance is None:
raise RuntimeError("El Agente no ha sido inicializado en el lifespan.")
return agent_instance
NALIIA_PROMPT = get_prompt_template(
"NALIIA_PROMPT")
naliia_tools = NaliiaTools().get_tools()
config = AgentConfig(system_prompt=NALIIA_PROMPT)
model = LLMFactory(MODEL_NAME).get_model()
return Agent(model=model, config=config, tools=naliia_tools) def get_async_connection_pool():
return AsyncConnectionPool(
conninfo=_connection_string,
max_size=20,
open=False,
kwargs={
"autocommit": True,
"prepare_threshold": 0,
"row_factory": dict_row
})

View File

@@ -0,0 +1,14 @@
import os
from dotenv import load_dotenv
class Settings:
DB_NAME = os.getenv("POSTGRES_DB", "db")
DB_USER = os.getenv("POSTGRES_USER", "postrges")
DB_PASSWORD = os.getenv("POSTGRES_PASSWORD", "password")
def __init__(self):
load_dotenv()
_settings = Settings()

View File

@@ -74,7 +74,7 @@ async def webhook_chat(request: Request, agent = Depends(get_agent)):
messages = [HumanMessage(content=user_message)] messages = [HumanMessage(content=user_message)]
agent_response = agent.invoke({"messages": messages}, config=config) agent_response = await agent.ainvoke({"messages": messages}, config=config)
agent_response_content = agent_response["messages"][-1].content agent_response_content = agent_response["messages"][-1].content
clean_jid = thread_id.split('@')[0] clean_jid = thread_id.split('@')[0]
@@ -93,4 +93,4 @@ async def webhook_chat(request: Request, agent = Depends(get_agent)):
logger.error(f"Error processing webhook: {e}") logger.error(f"Error processing webhook: {e}")
return {"error": str(e)} return {"error": str(e)}
return {"status": "sent", "reply": agent_response_content} return {"status": "sent", "reply": agent_response_content}

View File

@@ -1,11 +1,43 @@
from fastapi import FastAPI from fastapi import FastAPI
from naliiabotapi.api.v1.endpoints.chat import router as chat_router from naliiabotapi.api.v1.endpoints.chat import router as chat_router
from naliiabotapi.api.v1.webhooks.chat_hook import router as chat_hook_router from naliiabotapi.api.v1.webhooks.chat_hook import router as chat_hook_router
import naliiabotapi.api.dependencies as deps
from contextlib import asynccontextmanager
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
@asynccontextmanager
async def lifespan(app: FastAPI):
deps.db_pool = deps.get_async_connection_pool()
await deps.db_pool.open()
async with deps.db_pool.connection() as conn:
checkpointer = AsyncPostgresSaver(conn)
await checkpointer.setup()
model_name = "anthropic"
naliia_prompt = deps.get_prompt_template("NALIIA_PROMPT")
naliia_tools = deps.NaliiaTools().get_tools()
config = deps.AgentConfig(system_prompt=naliia_prompt)
model = deps.LLMFactory(model_name).get_model()
deps.agent_instance = deps.Agent(
model=model,
config=config,
checkpointer=checkpointer,
tools=naliia_tools
)
yield
await deps.db_pool.close()
app = FastAPI( app = FastAPI(
title="NaliiaBot API", title="NaliiaBot API",
description="API for NaliiaBot, a chatbot that provides customer service and related topics.", description="API for NaliiaBot, a chatbot that provides customer service and related topics.",
version="1.0.0" version="1.0.0",
lifespan=lifespan
) )
app.include_router(chat_router) app.include_router(chat_router)