Loading ai_agent/README.md +28 −1 Original line number Diff line number Diff line Loading @@ -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 Loading @@ -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 | --- Loading Loading @@ -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 ai_agent/requirements.txt +1 −0 Original line number Diff line number Diff line Loading @@ -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 ai_agent/routes/groq.py +17 −4 Original line number Diff line number Diff line Loading @@ -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)}) Loading @@ -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", Loading ai_agent/utils.py +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): Loading @@ -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 Loading
ai_agent/README.md +28 −1 Original line number Diff line number Diff line Loading @@ -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 Loading @@ -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 | --- Loading Loading @@ -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
ai_agent/requirements.txt +1 −0 Original line number Diff line number Diff line Loading @@ -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
ai_agent/routes/groq.py +17 −4 Original line number Diff line number Diff line Loading @@ -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)}) Loading @@ -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", Loading
ai_agent/utils.py +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): Loading @@ -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