feat: 重构 agentscope 缓存架构,新增消息和附件缓存
This commit is contained in:
@@ -1,11 +1,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from typing import Any, Callable, Protocol
|
||||
from uuid import UUID
|
||||
|
||||
from core.agentscope.caches.context_messages_cache import (
|
||||
create_context_messages_cache,
|
||||
)
|
||||
from core.agentscope.events.persistence import MessageRepository, SessionRepository
|
||||
from core.logging import get_logger
|
||||
from schemas.agent.forwarded_props import RuntimeMode
|
||||
from schemas.enums import AgentChatMessageRole, AgentChatSessionStatus
|
||||
from schemas.agent.system_agent import AgentType
|
||||
from schemas.agent.runtime_models import AgentOutput, RouterAgentOutput, ToolAgentOutput
|
||||
@@ -174,7 +179,10 @@ class SqlAlchemyEventStore:
|
||||
if locked_session is None:
|
||||
return
|
||||
seq = int(getattr(locked_session, "message_count", 0) or 0) + 1
|
||||
await message_repo.append_message(
|
||||
visibility_mask = self._resolve_stage_visibility_mask(
|
||||
event=event,
|
||||
)
|
||||
persisted = await message_repo.append_message(
|
||||
session_id=session_id,
|
||||
seq=seq,
|
||||
role=role,
|
||||
@@ -186,9 +194,16 @@ class SqlAlchemyEventStore:
|
||||
output_tokens=output_tokens,
|
||||
cost=cost,
|
||||
latency_ms=latency_ms,
|
||||
visibility_mask=self._resolve_stage_visibility_mask(
|
||||
event=event,
|
||||
),
|
||||
visibility_mask=visibility_mask,
|
||||
)
|
||||
await self._append_context_cache_message(
|
||||
session_id=session_id,
|
||||
event=event,
|
||||
visibility_mask=visibility_mask,
|
||||
role=role.value,
|
||||
content=content,
|
||||
metadata=metadata_model.model_dump(mode="json", exclude_none=True),
|
||||
timestamp=self._resolve_message_timestamp(persisted),
|
||||
)
|
||||
|
||||
current_status = getattr(chat_session, "status", AgentChatSessionStatus.RUNNING)
|
||||
@@ -339,16 +354,26 @@ class SqlAlchemyEventStore:
|
||||
if locked_session is None:
|
||||
return
|
||||
seq = int(getattr(locked_session, "message_count", 0) or 0) + 1
|
||||
await message_repo.append_message(
|
||||
visibility_mask = self._resolve_stage_visibility_mask(
|
||||
event=event,
|
||||
)
|
||||
persisted = await message_repo.append_message(
|
||||
session_id=session_id,
|
||||
seq=seq,
|
||||
role=AgentChatMessageRole.TOOL,
|
||||
content=content,
|
||||
tool_name=tool_output.tool_name,
|
||||
metadata=metadata_model.model_dump(mode="json", exclude_none=True),
|
||||
visibility_mask=self._resolve_stage_visibility_mask(
|
||||
event=event,
|
||||
),
|
||||
visibility_mask=visibility_mask,
|
||||
)
|
||||
await self._append_context_cache_message(
|
||||
session_id=session_id,
|
||||
event=event,
|
||||
visibility_mask=visibility_mask,
|
||||
role=AgentChatMessageRole.TOOL.value,
|
||||
content=content,
|
||||
metadata=metadata_model.model_dump(mode="json", exclude_none=True),
|
||||
timestamp=self._resolve_message_timestamp(persisted),
|
||||
)
|
||||
|
||||
current_status = getattr(chat_session, "status", AgentChatSessionStatus.RUNNING)
|
||||
@@ -377,6 +402,13 @@ class SqlAlchemyEventStore:
|
||||
*,
|
||||
event: dict[str, Any],
|
||||
) -> int:
|
||||
runtime_mode = self._event_value(event, "runtime_mode")
|
||||
if (
|
||||
isinstance(runtime_mode, str)
|
||||
and runtime_mode.strip().lower() == RuntimeMode.AUTOMATION.value
|
||||
):
|
||||
return bit_mask(bit=int(SystemVisibilityBit.UI_HISTORY))
|
||||
|
||||
raw_stage = self._event_value(event, "stage")
|
||||
if not isinstance(raw_stage, str):
|
||||
return bit_mask(bit=int(SystemVisibilityBit.UI_HISTORY))
|
||||
@@ -387,6 +419,61 @@ class SqlAlchemyEventStore:
|
||||
bit=int(SystemVisibilityBit.CONTEXT_ASSEMBLY)
|
||||
)
|
||||
|
||||
async def _append_context_cache_message(
|
||||
self,
|
||||
*,
|
||||
session_id: UUID,
|
||||
event: dict[str, Any],
|
||||
visibility_mask: int,
|
||||
role: str,
|
||||
content: str,
|
||||
metadata: dict[str, object] | None,
|
||||
timestamp: str,
|
||||
) -> None:
|
||||
message_payload: dict[str, object] = {
|
||||
"role": role,
|
||||
"content": content,
|
||||
"timestamp": timestamp,
|
||||
}
|
||||
if isinstance(metadata, dict):
|
||||
message_payload["metadata"] = metadata
|
||||
|
||||
try:
|
||||
context_cache = create_context_messages_cache()
|
||||
await context_cache.append_message(
|
||||
thread_id=str(session_id),
|
||||
runtime_mode=self._resolve_runtime_mode(event=event),
|
||||
visibility_mask=visibility_mask,
|
||||
message=message_payload,
|
||||
)
|
||||
except Exception as exc:
|
||||
self._logger.warning(
|
||||
"Failed to append context cache message from event",
|
||||
thread_id=str(session_id),
|
||||
error=str(exc),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_runtime_mode(*, event: dict[str, Any]) -> str:
|
||||
raw = event.get("runtime_mode")
|
||||
if isinstance(raw, str):
|
||||
normalized = raw.strip().lower()
|
||||
if normalized:
|
||||
return normalized
|
||||
return RuntimeMode.CHAT.value
|
||||
|
||||
@staticmethod
|
||||
def _resolve_message_timestamp(message: Any) -> str:
|
||||
created_at = getattr(message, "created_at", None)
|
||||
if isinstance(created_at, str) and created_at:
|
||||
return created_at
|
||||
if isinstance(created_at, datetime):
|
||||
try:
|
||||
return created_at.astimezone(timezone.utc).isoformat()
|
||||
except Exception:
|
||||
pass
|
||||
return datetime.now(timezone.utc).isoformat(timespec="seconds")
|
||||
|
||||
async def _update_session_state(
|
||||
self,
|
||||
*,
|
||||
|
||||
Reference in New Issue
Block a user