Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 15 additions & 13 deletions pilottai_tools/config/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,21 @@ class KnowledgeSource(BaseModel):
retry_delay: int = 5
timeout: int = 30

class MemoryItem(BaseModel):
"""Enhanced memory item model"""
model_config = ConfigDict(arbitrary_types_allowed=True)
text: str
metadata: Dict[str, Any] = Field(default_factory=dict)
timestamp: datetime = Field(default_factory=datetime.now)
tags: Set[str] = Field(default_factory=set)
priority: int = Field(ge=0, default=0)
expires_at: Optional[datetime] = None
version: int = 1

def is_expired(self) -> bool:
return self.expires_at and datetime.now() > self.expires_at



class MemoryEntry(BaseModel):
"""Enhanced memory entry with job awareness"""
Expand All @@ -38,19 +53,6 @@ class CacheEntry(BaseModel):
last_access: datetime = Field(default_factory=datetime.now)


class MemoryItem(BaseModel):
"""Enhanced memory item model"""
model_config = ConfigDict(arbitrary_types_allowed=True)
text: str
metadata: Dict[str, Any] = Field(default_factory=dict)
timestamp: datetime = Field(default_factory=datetime.now)
tags: Set[str] = Field(default_factory=set)
priority: int = Field(ge=0, default=0)
expires_at: Optional[datetime] = None
version: int = 1

def is_expired(self) -> bool:
return self.expires_at and datetime.now() > self.expires_at



28 changes: 28 additions & 0 deletions pilottai_tools/memory/memory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
import asyncio
from typing import Dict
from datetime import datetime

from pilottai_tools.utils.logger import Logger


class MemoryHandler:
def __init__(self, cache_size: int = 1000, cache_ttl: int = 3600):
self.sources: Dict[str, KnowledgeSource] = {}
self.last_updated: Dict[str, datetime] = {}
self.source_locks: Dict[str, asyncio.Lock] = {}
self.cache_lock = asyncio.Lock()
self.MAX_CACHE_SIZE = max(100, cache_size)
self.DEFAULT_CACHE_TTL = max(60, cache_ttl)
self.logger = Logger("MemoryHandler")
self._setup_logging()
self._cleanup_task = None

def _setup_logging(self):
if not self.logger.handlers:
handler = self.logger.StreamHandler()
formatter = self.logger.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
handler.setFormatter(formatter)
self.logger.addHandler(handler)
self.logger.setLevel(self.logger.INFO)
1 change: 1 addition & 0 deletions pilottai_tools/memory/redis/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
from client import get_redis_client
6 changes: 4 additions & 2 deletions pilottai_tools/memory/redis/client.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
import redis
import logging

from pilottai_tools.memory.redis.config import RedisConfig
from pilottai_tools.utils.logger import Logger

logger = Logger("RedisClient")

logger = logging.getLogger(__name__)

def get_redis_client():
try:
Expand Down
62 changes: 52 additions & 10 deletions pilottai_tools/memory/redis/publisher.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,56 @@
import logging
import json
from redis.client import get_redis_client
from typing import Optional, List, Dict

logger = logging.getLogger(__name__)
from pilottai_tools.memory.redis.client import get_redis_client
from pilottai_tools.utils.logger import Logger


logger = Logger("Publisher")
client = get_redis_client()

def publish_to_redis(channel: str, data: dict):
try:
message = json.dumps(data)
client.publish(channel, message)
logger.info(f"Published to Redis channel '{channel}': {message}")
except Exception as e:
logger.error(f"Failed to publish to Redis: {e}")

def _key(chat_id: str) -> str:
return f"chat:{chat_id}"


def create_conversation(chat_id: str) -> bool:
"""Initialize empty conversation list"""
key = _key(chat_id)
if not client.exists(key):
return client.rpush(key, *[])
return False


def add_message(chat_id: str, role: str, content: str) -> int:
"""Append a message to the conversation"""
key = _key(chat_id)
message = {"role": role, "content": content}
return client.rpush(key, json.dumps(message))


def get_conversation(chat_id: str) -> List[Dict[str, str]]:
"""Get the full conversation as a list of messages"""
key = _key(chat_id)
messages = client.lrange(key, 0, -1)
return [json.loads(msg) for msg in messages]


def delete_conversation(chat_id: str) -> int:
"""Delete conversation from Redis"""
key = _key(chat_id)
return client.delete(key)


def conversation_exists(chat_id: str) -> bool:
"""Check if conversation exists in Redis"""
key = _key(chat_id)
return client.exists(key) == 1


def export_conversation(chat_id: str) -> Optional[List[Dict[str, str]]]:
"""Retrieve and remove conversation (for storing in DB)"""
if conversation_exists(chat_id):
conv = get_conversation(chat_id)
delete_conversation(chat_id)
return conv
return None
12 changes: 6 additions & 6 deletions pilottai_tools/source/base/base_input.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional
from datetime import datetime
import logging
from pydantic import BaseModel, Field, ConfigDict

from pilottai_tools.utils.logger import Logger
from pilottai.memory.storage.local import DataStorage


Expand Down Expand Up @@ -61,17 +61,17 @@ def __init__(
# Setup logging
self.logger = self._setup_logger()

def _setup_logger(self) -> logging.Logger:
def _setup_logger(self) ->Logger:
"""Setup a logger for this input base"""
logger = logging.getLogger(f"InputSource_{self.name}")
logger = Logger(f"InputSource_{self.name}")
if not logger.handlers:
handler = logging.StreamHandler()
formatter = logging.Formatter(
handler = logger.StreamHandler()
formatter = logger.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
handler.setFormatter(formatter)
logger.addHandler(handler)
logger.setLevel(logging.INFO)
logger.setLevel(logger.INFO)
return logger

@abstractmethod
Expand Down
11 changes: 5 additions & 6 deletions pilottai_tools/source/knowledge.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,10 @@
import asyncio
import json
import logging
from collections import OrderedDict
from datetime import datetime
from typing import Dict, List, Optional, Any


from pilottai_tools.utils.logger import Logger
from pilottai_tools.config.model import KnowledgeSource, CacheEntry


Expand All @@ -18,19 +17,19 @@ def __init__(self, cache_size: int = 1000, cache_ttl: int = 3600):
self.cache_lock = asyncio.Lock()
self.MAX_CACHE_SIZE = max(100, cache_size)
self.DEFAULT_CACHE_TTL = max(60, cache_ttl)
self.logger = logging.getLogger("KnowledgeManager")
self.logger = Logger("DataManager")
self._setup_logging()
self._cleanup_task = None

def _setup_logging(self):
if not self.logger.handlers:
handler = logging.StreamHandler()
formatter = logging.Formatter(
handler = self.logger.StreamHandler()
formatter = self.logger.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
handler.setFormatter(formatter)
self.logger.addHandler(handler)
self.logger.setLevel(logging.INFO)
self.logger.setLevel(self.logger.INFO)

async def add_source(self, source: KnowledgeSource):
try:
Expand Down
Empty file.
85 changes: 85 additions & 0 deletions pilottai_tools/utils/formatter.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
import logging
from datetime import datetime
import json
import traceback

class ColoredFormatter(logging.Formatter):
"""Custom formatter with colors for console output"""

# Color codes
COLORS = {
'DEBUG': '\033[36m', # Cyan
'INFO': '\033[32m', # Green
'WARNING': '\033[33m', # Yellow
'ERROR': '\033[31m', # Red
'CRITICAL': '\033[35m', # Magenta
'RESET': '\033[0m' # Reset
}

def format(self, record):
# Add color to levelname
if hasattr(record, 'levelname'):
color = self.COLORS.get(record.levelname, self.COLORS['RESET'])
record.levelname = f"{color}{record.levelname}{self.COLORS['RESET']}"

# Format the message
formatted = super().format(record)

# Add context information if present
if hasattr(record, 'context') and record.context:
context_str = json.dumps(record.context, indent=2)
formatted += f"\n📋 Context: {context_str}"

return formatted


class JsonFormatter(logging.Formatter):
"""JSON formatter for structured logging"""

def format(self, record):
log_entry = {
'timestamp': datetime.fromtimestamp(record.created).isoformat(),
'level': record.levelname,
'logger': record.name,
'message': record.getMessage(),
'module': record.module,
'function': record.funcName,
'line': record.lineno,
'thread': record.thread,
'thread_name': record.threadName
}

# Add context information
if hasattr(record, 'context'):
log_entry['context'] = record.context

if hasattr(record, 'user_id'):
log_entry['user_id'] = record.user_id

if hasattr(record, 'request_id'):
log_entry['request_id'] = record.request_id

if hasattr(record, 'ip_address'):
log_entry['ip_address'] = record.ip_address

if hasattr(record, 'endpoint'):
log_entry['endpoint'] = record.endpoint

if hasattr(record, 'method'):
log_entry['method'] = record.method

if hasattr(record, 'duration'):
log_entry['duration'] = record.duration

if hasattr(record, 'status_code'):
log_entry['status_code'] = record.status_code

# Add exception information
if record.exc_info:
log_entry['exception'] = {
'type': record.exc_info[0].__name__,
'message': str(record.exc_info[1]),
'traceback': traceback.format_exception(*record.exc_info)
}

return json.dumps(log_entry, ensure_ascii=False)
Loading
Loading