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

refactor: static type checking - ai agent, #8

parent 5d00e0b1
Loading
Loading
Loading
Loading

ai_agent/__init__.py

0 → 100644
+1 −0
Original line number Diff line number Diff line
+13 −6
Original line number Diff line number Diff line
@@ -2,7 +2,8 @@ import hashlib
import json
import logging
import os
from typing import Any
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
@@ -10,6 +11,12 @@ 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:
@@ -51,9 +58,9 @@ class RedisShortTermMemory:

        key = self._session_key(session_id)
        try:
            raw_messages = await client.lrange(key, 0, -1)
            raw_messages = cast(list[str], await _maybe_await(client.lrange(key, 0, -1)))
            if raw_messages:
                await client.expire(key, self.ttl_seconds)
                await _maybe_await(client.expire(key, self.ttl_seconds))
        except Exception:
            logger.exception("Failed reading memory key %s", key)
            return []
@@ -102,9 +109,9 @@ class RedisShortTermMemory:
        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)
            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)

+1 −0
Original line number Diff line number Diff line
+1 −1
Original line number Diff line number Diff line
import io
import sys
import objgraph
import objgraph  # type: ignore[import-untyped]
from fastapi import APIRouter
from fastapi.responses import PlainTextResponse

+1 −0
Original line number Diff line number Diff line
Loading