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

refactor: linting and formatting, #8

parent c4e929a1
Loading
Loading
Loading
Loading
+16 −5
Original line number Diff line number Diff line
@@ -24,7 +24,9 @@ def _int_env(name: str, default: int) -> int:
    try:
        value = int(raw_value)
    except ValueError:
        logger.warning("Invalid %s value '%s'. Using default %s", name, raw_value, default)
        logger.warning(
            "Invalid %s value '%s'. Using default %s", name, raw_value, default
        )
        return default
    return max(1, value)

@@ -34,7 +36,10 @@ class RedisShortTermMemory:
        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.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)

@@ -42,7 +47,9 @@ class RedisShortTermMemory:
        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)
            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:
@@ -58,7 +65,9 @@ class RedisShortTermMemory:

        key = self._session_key(session_id)
        try:
            raw_messages = cast(list[str], await _maybe_await(client.lrange(key, 0, -1)))
            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:
@@ -94,7 +103,9 @@ class RedisShortTermMemory:
        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:
    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()
+7 −2
Original line number Diff line number Diff line
@@ -35,7 +35,10 @@ 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), session_id: 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"})

@@ -47,7 +50,9 @@ async def groq_stream(query: str | None = Query(default=None), session_id: str |
        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)
        generator = await stream_agent_response(
            agent, effective_query, on_complete=on_complete
        )
        return StreamingResponse(
            generator(),
            media_type="text/event-stream",
+13 −6
Original line number Diff line number Diff line
@@ -10,12 +10,14 @@ logging.getLogger("mcp_use").setLevel(logging.ERROR)

logger = logging.getLogger(__name__)


def require_env(name: str) -> str:
    value = os.getenv(name)
    if value is None or not value.strip():
        raise EnvironmentError(f"{name} environment variable is not set")
    return value.strip()


async def create_groq_agent():
    """Create an MCPAgent using Groq backend."""
    try:
@@ -25,13 +27,18 @@ async def create_groq_agent():
        logger.error("Groq agent configuration error: %s", e)
        raise

    model_name = os.getenv("GROQ_MODEL_NAME", "openai/gpt-oss-20b").strip() or "openai/gpt-oss-20b"
    model_name = (
        os.getenv("GROQ_MODEL_NAME", "openai/gpt-oss-20b").strip()
        or "openai/gpt-oss-20b"
    )

    client = MCPClient({
    client = MCPClient(
        {
            "mcpServers": {"http": {"url": mcp_server_url}},
            "use-oauth2": False,
        "use-oidc": False
    })
            "use-oidc": False,
        }
    )
    await client.create_session("http")

    llm = ChatGroq(
+3 −1
Original line number Diff line number Diff line
@@ -3,6 +3,7 @@ import json
import logging
from collections.abc import Awaitable, Callable


# --- SSE Helper ---
def sse_event(data, event="message", id=None, retry=None):
    lines = []
@@ -50,4 +51,5 @@ async def stream_agent_response(
        finally:
            await agent.close()
            logging.info("Agent closed after streaming")

    return generator
+10 −3
Original line number Diff line number Diff line
@@ -9,9 +9,15 @@ from fastapi import FastAPI
from fastapi.responses import PlainTextResponse
from dotenv import load_dotenv
from mcp_module.tools.edge_application import get_app_definitions
from mcp_module.tools.qod import create_qod_session_oai, get_qod_session_oai, delete_qod_session_oai
from mcp_module.tools.qod import (
    create_qod_session_oai,
    get_qod_session_oai,
    delete_qod_session_oai,
)

logging.basicConfig(level=logging.DEBUG, format="%(asctime)s [%(levelname)s] %(message)s")
logging.basicConfig(
    level=logging.DEBUG, format="%(asctime)s [%(levelname)s] %(message)s"
)
logger = logging.getLogger("MCP server")

load_dotenv()
@@ -39,16 +45,17 @@ app = FastAPI(
    redirect_slashes=False,
)


# Health check endpoint
@app.get("/health", response_class=PlainTextResponse)
async def health_check():
    return "healthy"


# Mount MCP server
app.mount("", mcp_app)

if __name__ == "__main__":

    logger.info("Starting MCP Server with FastAPI...")
    host = os.getenv("MCP_HOST", "127.0.0.1")
    port = int(os.getenv("MCP_PORT", "8004"))
Loading