forked from MemTensor/MemOS
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbase_handler.py
More file actions
68 lines (52 loc) · 2.37 KB
/
Copy pathbase_handler.py
File metadata and controls
68 lines (52 loc) · 2.37 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
from __future__ import annotations
from abc import abstractmethod
from typing import TYPE_CHECKING
from memos.log import get_logger
from memos.mem_scheduler.utils.misc_utils import group_messages_by_user_and_mem_cube
if TYPE_CHECKING:
from collections.abc import Callable
from memos.mem_scheduler.schemas.message_schemas import ScheduleMessageItem
from memos.mem_scheduler.task_schedule_modules.context import SchedulerHandlerContext
logger = get_logger(__name__)
class BaseSchedulerHandler:
def __init__(self, scheduler_context: SchedulerHandlerContext) -> None:
self.scheduler_context = scheduler_context
@property
@abstractmethod
def expected_task_label(self) -> str:
"""The expected task label for this handler."""
...
def validate_and_log_messages(self, messages: list[ScheduleMessageItem], label: str) -> None:
logger.info(f"Messages {messages} assigned to {label} handler.")
self.scheduler_context.services.validate_messages(messages=messages, label=label)
def handle_exception(self, e: Exception, message: str = "Error processing messages") -> None:
logger.error(f"{message}: {e}", exc_info=True)
def process_grouped_messages(
self,
messages: list[ScheduleMessageItem],
message_handler: Callable[[str, str, list[ScheduleMessageItem]], None],
) -> None:
grouped_messages = group_messages_by_user_and_mem_cube(messages=messages)
for user_id, user_batches in grouped_messages.items():
for mem_cube_id, batch in user_batches.items():
if not batch:
continue
try:
message_handler(user_id, mem_cube_id, batch)
except Exception as e:
self.handle_exception(
e, f"Error processing batch for user {user_id}, mem_cube {mem_cube_id}"
)
@abstractmethod
def batch_handler(
self, user_id: str, mem_cube_id: str, batch: list[ScheduleMessageItem]
) -> None: ...
def __call__(self, messages: list[ScheduleMessageItem]) -> None:
"""
Process the messages.
"""
self.validate_and_log_messages(messages=messages, label=self.expected_task_label)
self.process_grouped_messages(
messages=messages,
message_handler=self.batch_handler,
)