Skip to content

Commit 4d49ff6

Browse files
authored
Merge pull request #4 from pilottai/feat/logger
added customer logger to maintain consistency
2 parents 9e34805 + 19819f4 commit 4d49ff6

10 files changed

Lines changed: 426 additions & 37 deletions

File tree

pilottai_tools/config/model.py

Lines changed: 15 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,21 @@ class KnowledgeSource(BaseModel):
1717
retry_delay: int = 5
1818
timeout: int = 30
1919

20+
class MemoryItem(BaseModel):
21+
"""Enhanced memory item model"""
22+
model_config = ConfigDict(arbitrary_types_allowed=True)
23+
text: str
24+
metadata: Dict[str, Any] = Field(default_factory=dict)
25+
timestamp: datetime = Field(default_factory=datetime.now)
26+
tags: Set[str] = Field(default_factory=set)
27+
priority: int = Field(ge=0, default=0)
28+
expires_at: Optional[datetime] = None
29+
version: int = 1
30+
31+
def is_expired(self) -> bool:
32+
return self.expires_at and datetime.now() > self.expires_at
33+
34+
2035

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

4055

41-
class MemoryItem(BaseModel):
42-
"""Enhanced memory item model"""
43-
model_config = ConfigDict(arbitrary_types_allowed=True)
44-
text: str
45-
metadata: Dict[str, Any] = Field(default_factory=dict)
46-
timestamp: datetime = Field(default_factory=datetime.now)
47-
tags: Set[str] = Field(default_factory=set)
48-
priority: int = Field(ge=0, default=0)
49-
expires_at: Optional[datetime] = None
50-
version: int = 1
51-
52-
def is_expired(self) -> bool:
53-
return self.expires_at and datetime.now() > self.expires_at
5456

5557

5658

pilottai_tools/memory/memory.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,28 @@
1+
import asyncio
2+
from typing import Dict
3+
from datetime import datetime
4+
5+
from pilottai_tools.utils.logger import Logger
6+
7+
8+
class MemoryHandler:
9+
def __init__(self, cache_size: int = 1000, cache_ttl: int = 3600):
10+
self.sources: Dict[str, KnowledgeSource] = {}
11+
self.last_updated: Dict[str, datetime] = {}
12+
self.source_locks: Dict[str, asyncio.Lock] = {}
13+
self.cache_lock = asyncio.Lock()
14+
self.MAX_CACHE_SIZE = max(100, cache_size)
15+
self.DEFAULT_CACHE_TTL = max(60, cache_ttl)
16+
self.logger = Logger("MemoryHandler")
17+
self._setup_logging()
18+
self._cleanup_task = None
19+
20+
def _setup_logging(self):
21+
if not self.logger.handlers:
22+
handler = self.logger.StreamHandler()
23+
formatter = self.logger.Formatter(
24+
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
25+
)
26+
handler.setFormatter(formatter)
27+
self.logger.addHandler(handler)
28+
self.logger.setLevel(self.logger.INFO)
Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
from client import get_redis_client

pilottai_tools/memory/redis/client.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,10 @@
11
import redis
2-
import logging
2+
33
from pilottai_tools.memory.redis.config import RedisConfig
4+
from pilottai_tools.utils.logger import Logger
5+
6+
logger = Logger("RedisClient")
47

5-
logger = logging.getLogger(__name__)
68

79
def get_redis_client():
810
try:
Lines changed: 52 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,56 @@
1-
import logging
21
import json
3-
from redis.client import get_redis_client
2+
from typing import Optional, List, Dict
43

5-
logger = logging.getLogger(__name__)
4+
from pilottai_tools.memory.redis.client import get_redis_client
5+
from pilottai_tools.utils.logger import Logger
6+
7+
8+
logger = Logger("Publisher")
69
client = get_redis_client()
710

8-
def publish_to_redis(channel: str, data: dict):
9-
try:
10-
message = json.dumps(data)
11-
client.publish(channel, message)
12-
logger.info(f"Published to Redis channel '{channel}': {message}")
13-
except Exception as e:
14-
logger.error(f"Failed to publish to Redis: {e}")
11+
12+
def _key(chat_id: str) -> str:
13+
return f"chat:{chat_id}"
14+
15+
16+
def create_conversation(chat_id: str) -> bool:
17+
"""Initialize empty conversation list"""
18+
key = _key(chat_id)
19+
if not client.exists(key):
20+
return client.rpush(key, *[])
21+
return False
22+
23+
24+
def add_message(chat_id: str, role: str, content: str) -> int:
25+
"""Append a message to the conversation"""
26+
key = _key(chat_id)
27+
message = {"role": role, "content": content}
28+
return client.rpush(key, json.dumps(message))
29+
30+
31+
def get_conversation(chat_id: str) -> List[Dict[str, str]]:
32+
"""Get the full conversation as a list of messages"""
33+
key = _key(chat_id)
34+
messages = client.lrange(key, 0, -1)
35+
return [json.loads(msg) for msg in messages]
36+
37+
38+
def delete_conversation(chat_id: str) -> int:
39+
"""Delete conversation from Redis"""
40+
key = _key(chat_id)
41+
return client.delete(key)
42+
43+
44+
def conversation_exists(chat_id: str) -> bool:
45+
"""Check if conversation exists in Redis"""
46+
key = _key(chat_id)
47+
return client.exists(key) == 1
48+
49+
50+
def export_conversation(chat_id: str) -> Optional[List[Dict[str, str]]]:
51+
"""Retrieve and remove conversation (for storing in DB)"""
52+
if conversation_exists(chat_id):
53+
conv = get_conversation(chat_id)
54+
delete_conversation(chat_id)
55+
return conv
56+
return None

pilottai_tools/source/base/base_input.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,9 @@
11
from abc import ABC, abstractmethod
22
from typing import Any, Dict, List, Optional
33
from datetime import datetime
4-
import logging
54
from pydantic import BaseModel, Field, ConfigDict
65

6+
from pilottai_tools.utils.logger import Logger
77
from pilottai.memory.storage.local import DataStorage
88

99

@@ -61,17 +61,17 @@ def __init__(
6161
# Setup logging
6262
self.logger = self._setup_logger()
6363

64-
def _setup_logger(self) -> logging.Logger:
64+
def _setup_logger(self) ->Logger:
6565
"""Setup a logger for this input base"""
66-
logger = logging.getLogger(f"InputSource_{self.name}")
66+
logger = Logger(f"InputSource_{self.name}")
6767
if not logger.handlers:
68-
handler = logging.StreamHandler()
69-
formatter = logging.Formatter(
68+
handler = logger.StreamHandler()
69+
formatter = logger.Formatter(
7070
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
7171
)
7272
handler.setFormatter(formatter)
7373
logger.addHandler(handler)
74-
logger.setLevel(logging.INFO)
74+
logger.setLevel(logger.INFO)
7575
return logger
7676

7777
@abstractmethod

pilottai_tools/source/knowledge.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,10 @@
11
import asyncio
22
import json
3-
import logging
43
from collections import OrderedDict
54
from datetime import datetime
65
from typing import Dict, List, Optional, Any
76

8-
7+
from pilottai_tools.utils.logger import Logger
98
from pilottai_tools.config.model import KnowledgeSource, CacheEntry
109

1110

@@ -18,19 +17,19 @@ def __init__(self, cache_size: int = 1000, cache_ttl: int = 3600):
1817
self.cache_lock = asyncio.Lock()
1918
self.MAX_CACHE_SIZE = max(100, cache_size)
2019
self.DEFAULT_CACHE_TTL = max(60, cache_ttl)
21-
self.logger = logging.getLogger("KnowledgeManager")
20+
self.logger = Logger("DataManager")
2221
self._setup_logging()
2322
self._cleanup_task = None
2423

2524
def _setup_logging(self):
2625
if not self.logger.handlers:
27-
handler = logging.StreamHandler()
28-
formatter = logging.Formatter(
26+
handler = self.logger.StreamHandler()
27+
formatter = self.logger.Formatter(
2928
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
3029
)
3130
handler.setFormatter(formatter)
3231
self.logger.addHandler(handler)
33-
self.logger.setLevel(logging.INFO)
32+
self.logger.setLevel(self.logger.INFO)
3433

3534
async def add_source(self, source: KnowledgeSource):
3635
try:

pilottai_tools/utils/__init__.py

Whitespace-only changes.

pilottai_tools/utils/formatter.py

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
1+
import logging
2+
from datetime import datetime
3+
import json
4+
import traceback
5+
6+
class ColoredFormatter(logging.Formatter):
7+
"""Custom formatter with colors for console output"""
8+
9+
# Color codes
10+
COLORS = {
11+
'DEBUG': '\033[36m', # Cyan
12+
'INFO': '\033[32m', # Green
13+
'WARNING': '\033[33m', # Yellow
14+
'ERROR': '\033[31m', # Red
15+
'CRITICAL': '\033[35m', # Magenta
16+
'RESET': '\033[0m' # Reset
17+
}
18+
19+
def format(self, record):
20+
# Add color to levelname
21+
if hasattr(record, 'levelname'):
22+
color = self.COLORS.get(record.levelname, self.COLORS['RESET'])
23+
record.levelname = f"{color}{record.levelname}{self.COLORS['RESET']}"
24+
25+
# Format the message
26+
formatted = super().format(record)
27+
28+
# Add context information if present
29+
if hasattr(record, 'context') and record.context:
30+
context_str = json.dumps(record.context, indent=2)
31+
formatted += f"\n📋 Context: {context_str}"
32+
33+
return formatted
34+
35+
36+
class JsonFormatter(logging.Formatter):
37+
"""JSON formatter for structured logging"""
38+
39+
def format(self, record):
40+
log_entry = {
41+
'timestamp': datetime.fromtimestamp(record.created).isoformat(),
42+
'level': record.levelname,
43+
'logger': record.name,
44+
'message': record.getMessage(),
45+
'module': record.module,
46+
'function': record.funcName,
47+
'line': record.lineno,
48+
'thread': record.thread,
49+
'thread_name': record.threadName
50+
}
51+
52+
# Add context information
53+
if hasattr(record, 'context'):
54+
log_entry['context'] = record.context
55+
56+
if hasattr(record, 'user_id'):
57+
log_entry['user_id'] = record.user_id
58+
59+
if hasattr(record, 'request_id'):
60+
log_entry['request_id'] = record.request_id
61+
62+
if hasattr(record, 'ip_address'):
63+
log_entry['ip_address'] = record.ip_address
64+
65+
if hasattr(record, 'endpoint'):
66+
log_entry['endpoint'] = record.endpoint
67+
68+
if hasattr(record, 'method'):
69+
log_entry['method'] = record.method
70+
71+
if hasattr(record, 'duration'):
72+
log_entry['duration'] = record.duration
73+
74+
if hasattr(record, 'status_code'):
75+
log_entry['status_code'] = record.status_code
76+
77+
# Add exception information
78+
if record.exc_info:
79+
log_entry['exception'] = {
80+
'type': record.exc_info[0].__name__,
81+
'message': str(record.exc_info[1]),
82+
'traceback': traceback.format_exception(*record.exc_info)
83+
}
84+
85+
return json.dumps(log_entry, ensure_ascii=False)

0 commit comments

Comments
 (0)