Commit ba45cd9c authored by Kostas Chartsias's avatar Kostas Chartsias
Browse files

feature: agent memory, #6

parent 251c1737
Loading
Loading
Loading
Loading
+28 −1
Original line number Diff line number Diff line
@@ -36,6 +36,12 @@ MCP_SERVER_URL=http://127.0.0.1:8004/mcp
# --- Groq ---
GROQ_API_KEY=your_groq_api_key_here
GROQ_MODEL_NAME=qwen/qwen3-32b

# --- Optional short-term memory (Redis) ---
REDIS_URL=redis://127.0.0.1:6379/0
MEMORY_WINDOW_TURNS=6
MEMORY_TTL_SECONDS=3600
MEMORY_KEY_PREFIX=ai_agent:session
```

### Environment Variables
@@ -47,6 +53,10 @@ GROQ_MODEL_NAME=qwen/qwen3-32b
| `MCP_SERVER_URL`  | URL of the MCP server used by sub-agent |
| `GROQ_API_KEY`    | API key for Groq LLM access             |
| `GROQ_MODEL_NAME` | Groq model identifier                   |
| `REDIS_URL` | Redis DSN for short-term memory (`redis://host:6379/0`) |
| `MEMORY_WINDOW_TURNS` | Number of recent user/assistant turns to retain |
| `MEMORY_TTL_SECONDS` | Session memory TTL in seconds |
| `MEMORY_KEY_PREFIX` | Prefix for Redis session keys |

---

@@ -142,5 +152,22 @@ from ai_agent.routes.my_agent import router as my_agent_router
app.include_router(my_agent_router)
```

---
## 🧠 Session Memory

When `REDIS_URL` is configured, both Groq endpoints support short-term memory with a sliding window.

- `POST /groq-mcp`: pass `session_id` in the JSON body.
- `GET /groq-mcp/stream`: pass `session_id` as a query parameter.

Example:

```json
{
  "session_id": "user-123-session-a",
  "query": "What did I ask you before?"
}
```

If `REDIS_URL` is missing, the API remains stateless.

---

ai_agent/memory.py

0 → 100644
+114 −0
Original line number Diff line number Diff line
import hashlib
import json
import logging
import os
from typing import Any

from redis.asyncio import Redis
from redis.asyncio import from_url as redis_from_url

logger = logging.getLogger(__name__)


def _int_env(name: str, default: int) -> int:
    raw_value = os.getenv(name)
    if raw_value is None:
        return default
    try:
        value = int(raw_value)
    except ValueError:
        logger.warning("Invalid %s value '%s'. Using default %s", name, raw_value, default)
        return default
    return max(1, value)


class RedisShortTermMemory:
    def __init__(self) -> None:
        self.redis_url = os.getenv("REDIS_URL", "").strip()
        self.window_turns = _int_env("MEMORY_WINDOW_TURNS", 6)
        self.ttl_seconds = _int_env("MEMORY_TTL_SECONDS", 3600)
        self.key_prefix = os.getenv("MEMORY_KEY_PREFIX", "ai_agent:session").strip() or "ai_agent:session"
        self._client: Redis | None = None
        self.enabled = bool(self.redis_url)

    async def _get_client(self) -> Redis | None:
        if not self.enabled:
            return None
        if self._client is None:
            self._client = redis_from_url(self.redis_url, encoding="utf-8", decode_responses=True)
        return self._client

    def _session_key(self, session_id: str) -> str:
        session_hash = hashlib.sha256(session_id.encode("utf-8")).hexdigest()
        return f"{self.key_prefix}:{session_hash}"

    async def get_messages(self, session_id: str | None) -> list[dict[str, str]]:
        if not session_id:
            return []
        client = await self._get_client()
        if client is None:
            return []

        key = self._session_key(session_id)
        try:
            raw_messages = await client.lrange(key, 0, -1)
            if raw_messages:
                await client.expire(key, self.ttl_seconds)
        except Exception:
            logger.exception("Failed reading memory key %s", key)
            return []

        parsed: list[dict[str, str]] = []
        for entry in raw_messages:
            try:
                value: dict[str, Any] = json.loads(entry)
            except (json.JSONDecodeError, TypeError):
                logger.warning("Skipping malformed session message for key %s", key)
                continue
            role = str(value.get("role", "")).strip()
            content = str(value.get("content", "")).strip()
            if not role or not content:
                continue
            parsed.append({"role": role, "content": content})
        return parsed

    async def augment_query(self, query: str, session_id: str | None) -> str:
        messages = await self.get_messages(session_id)
        if not messages:
            return query

        history_lines = [
            "Use the recent conversation history below to answer consistently.",
            "History:",
        ]
        for message in messages:
            history_lines.append(f"{message['role'].upper()}: {message['content']}")
        history_lines.append("")
        history_lines.append(f"Current USER request: {query}")
        return "\n".join(history_lines)

    async def append_turn(self, session_id: str | None, user_query: str, assistant_response: str) -> None:
        if not session_id:
            return
        client = await self._get_client()
        if client is None:
            return

        key = self._session_key(session_id)
        entries = (
            json.dumps({"role": "user", "content": user_query}),
            json.dumps({"role": "assistant", "content": assistant_response}),
        )
        max_messages = self.window_turns * 2

        try:
            await client.rpush(key, *entries)
            await client.ltrim(key, -max_messages, -1)
            await client.expire(key, self.ttl_seconds)
        except Exception:
            logger.exception("Failed writing memory key %s", key)

    async def close(self) -> None:
        if self._client is not None:
            await self._client.close()
            self._client = None
+1 −0
Original line number Diff line number Diff line
@@ -4,3 +4,4 @@ objgraph==3.6.2
python-dotenv==1.2.1
fastapi==0.118.0
uvicorn==0.37.0
redis==5.2.1
+17 −4
Original line number Diff line number Diff line
@@ -2,23 +2,31 @@ import logging
from fastapi import APIRouter, Body, Query
from fastapi.responses import JSONResponse, StreamingResponse

from ai_agent.memory import RedisShortTermMemory
from ai_agent.sub_agents.groq_agent import create_groq_agent
from ai_agent.utils import stream_agent_response

router = APIRouter()
memory_store = RedisShortTermMemory()


@router.post("/groq-mcp")
async def groq_query(payload: dict | None = Body(default=None)):
    query = payload.get("query") if payload else None
    session_id = payload.get("session_id") if payload else None

    if not query:
        return JSONResponse(status_code=400, content={"error": "No query provided"})

    agent = await create_groq_agent()
    try:
        result = await agent.run(query)
        return {"response": result}
        effective_query = await memory_store.augment_query(query, session_id)
        result = await agent.run(effective_query)
        await memory_store.append_turn(session_id, query, result)
        response = {"response": result}
        if session_id:
            response["session_id"] = session_id
        return response
    except Exception as e:
        logging.error("Error in groq_query", exc_info=True)
        return JSONResponse(status_code=500, content={"error": str(e)})
@@ -27,14 +35,19 @@ async def groq_query(payload: dict | None = Body(default=None)):


@router.get("/groq-mcp/stream")
async def groq_stream(query: str | None = Query(default=None)):
async def groq_stream(query: str | None = Query(default=None), session_id: str | None = Query(default=None)):
    if not query:
        return JSONResponse(status_code=400, content={"error": "No query provided"})

    agent = await create_groq_agent()

    try:
        generator = await stream_agent_response(agent, query)
        effective_query = await memory_store.augment_query(query, session_id)

        async def on_complete(result: str):
            await memory_store.append_turn(session_id, query, result)

        generator = await stream_agent_response(agent, effective_query, on_complete=on_complete)
        return StreamingResponse(
            generator(),
            media_type="text/event-stream",
+11 −1
Original line number Diff line number Diff line
import asyncio
import json
import logging
from collections.abc import Awaitable, Callable

# --- SSE Helper ---
def sse_event(data, event="message", id=None, retry=None):
@@ -24,12 +25,21 @@ def sse_event(data, event="message", id=None, retry=None):


# --- Common SSE Stream Helper ---
async def stream_agent_response(agent, query):
async def stream_agent_response(
    agent,
    query,
    on_complete: Callable[[str], Awaitable[None] | None] | None = None,
):
    chunk_size = 50

    async def generator():
        event_id = 0
        try:
            result = await agent.run(query)
            if on_complete is not None:
                maybe_awaitable = on_complete(result)
                if maybe_awaitable is not None:
                    await maybe_awaitable
            for i in range(0, len(result), chunk_size):
                yield sse_event({"chunk": result[i:i + chunk_size]}, id=event_id)
                event_id += 1
Loading