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

Merge branch 'feature/edge_apps' into 'develop'

Feature/edge apps

See merge request !3
parents 251c1737 8324752d
Loading
Loading
Loading
Loading
Loading
+72 −2
Original line number Diff line number Diff line
default:
  image: python:3.12-slim
  cache:
    paths:
      - .cache/pip
  before_script:
    - pip install -r "$CI_PROJECT_DIR/mcp_module/requirements.txt"
    - pip install -r "$CI_PROJECT_DIR/ai_agent/requirements.txt"
    - pip install mypy pytest ruff import-linter
    - export PYTHONPATH="$CI_PROJECT_DIR:$PYTHONPATH"

stages:
  - type
  - architecture
  - lint
  - format
  - test
  - build

variables:
  PIP_CACHE_DIR: "$CI_PROJECT_DIR/.cache/pip"
  SOURCE_DIRS: "mcp_module/ ai_agent/"
  TEST_DIRS: "tests/"

include:
  - template: Security/SAST.gitlab-ci.yml

sast:
  stage: test

type:
  stage: type
  tags:
    - docker
  script:
    - mypy $SOURCE_DIRS

architecture:
  stage: architecture
  tags:
    - docker
  script:
    - |
      if [ -f "$CI_PROJECT_DIR/.importlinter" ] || \
         [ -f "$CI_PROJECT_DIR/.importlinter.ini" ] || \
         [ -f "$CI_PROJECT_DIR/pyproject.toml" ] || \
         [ -f "$CI_PROJECT_DIR/setup.cfg" ]; then
        lint-imports
      else
        echo "No import-linter config found; skipping architecture check."
      fi

lint:
  stage: lint
  tags:
    - docker
  script:
    - ruff check $SOURCE_DIRS $TEST_DIRS

format:
  stage: format
  tags:
    - docker
  script:
    - ruff format --check $SOURCE_DIRS $TEST_DIRS

test:
  stage: test
  tags:
    - docker
  script:
    - pytest

.variables-template:
  tags:
    - "shell"
    - shell
  before_script:
    - docker login -u "$CI_REGISTRY_USER" -p "$CI_REGISTRY_PASSWORD" "$CI_REGISTRY"
    - |
+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/__init__.py

0 → 100644
+1 −0
Original line number Diff line number Diff line

ai_agent/memory.py

0 → 100644
+132 −0
Original line number Diff line number Diff line
import hashlib
import json
import logging
import os
from collections.abc import Awaitable
from typing import Any, cast

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

logger = logging.getLogger(__name__)


async def _maybe_await(value: Awaitable[Any] | Any) -> Any:
    if isinstance(value, Awaitable):
        return await value
    return value


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 = cast(
                list[str], await _maybe_await(client.lrange(key, 0, -1))
            )
            if raw_messages:
                await _maybe_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 _maybe_await(client.rpush(key, *entries))
            await _maybe_await(client.ltrim(key, -max_messages, -1))
            await _maybe_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
+7 −0
Original line number Diff line number Diff line
from ai_agent.prompts.memory import CONVERSATION_MEMORY_PROMPT
from ai_agent.prompts.tool_agent import TOOL_USING_ASSISTANT_PROMPT

__all__ = [
    "CONVERSATION_MEMORY_PROMPT",
    "TOOL_USING_ASSISTANT_PROMPT",
]
Loading