feat: 重构 memory 系统,支持 user memory 和 work memory 分离
This commit is contained in:
@@ -6,10 +6,16 @@ from core.agentscope.tools.custom.calendar import (
|
||||
from core.agentscope.tools.custom.user_lookup import (
|
||||
user_lookup,
|
||||
)
|
||||
from core.agentscope.tools.custom.memory import (
|
||||
memory_forget,
|
||||
memory_write,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"calendar_read",
|
||||
"calendar_write",
|
||||
"calendar_share",
|
||||
"user_lookup",
|
||||
"memory_write",
|
||||
"memory_forget",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,330 @@
|
||||
from copy import deepcopy
|
||||
from typing import Annotated, Any, cast
|
||||
from uuid import UUID
|
||||
|
||||
from agentscope.tool import ToolResponse
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from core.agentscope.tools.tool_call_context import get_current_tool_call_id
|
||||
from core.agentscope.tools.utils.memory_domain import (
|
||||
create_memories_service,
|
||||
map_memory_exception,
|
||||
)
|
||||
from core.agentscope.tools.utils.tool_response_builder import (
|
||||
build_error_output,
|
||||
build_tool_response,
|
||||
)
|
||||
from models.memories import MemoryType
|
||||
from schemas.agent.runtime_models import ToolAgentOutput, ToolStatus
|
||||
from schemas.memories.memory_content import UserMemoryContent, WorkProfileContent
|
||||
|
||||
|
||||
class MemoryWriteArgs(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
memory_type: MemoryType = MemoryType.USER
|
||||
user_content: UserMemoryContent | None = None
|
||||
work_content: WorkProfileContent | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_content(self) -> "MemoryWriteArgs":
|
||||
if self.memory_type == MemoryType.USER:
|
||||
if self.user_content is None or self.work_content is not None:
|
||||
raise ValueError("memory_type=user requires user_content only")
|
||||
else:
|
||||
if self.work_content is None or self.user_content is not None:
|
||||
raise ValueError("memory_type=work requires work_content only")
|
||||
return self
|
||||
|
||||
|
||||
class MemoryForgetArgs(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
memory_type: MemoryType = MemoryType.USER
|
||||
forget_paths: list[str] = Field(min_length=1, max_length=100)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_forget_paths(self) -> "MemoryForgetArgs":
|
||||
allowed_roots = (
|
||||
set(UserMemoryContent.model_fields)
|
||||
if self.memory_type == MemoryType.USER
|
||||
else set(WorkProfileContent.model_fields)
|
||||
)
|
||||
normalized: list[str] = []
|
||||
for raw_path in self.forget_paths:
|
||||
path = raw_path.strip()
|
||||
if not path:
|
||||
continue
|
||||
parts = [part for part in path.split(".") if part]
|
||||
if not parts:
|
||||
continue
|
||||
if len(parts) > 5:
|
||||
raise ValueError("forget path depth exceeds limit")
|
||||
if parts[0] not in allowed_roots:
|
||||
raise ValueError("forget path root is not allowed")
|
||||
normalized.append(path)
|
||||
if not normalized:
|
||||
raise ValueError("forget_paths cannot be empty")
|
||||
self.forget_paths = normalized
|
||||
return self
|
||||
|
||||
|
||||
def _memory_error_output(
|
||||
*,
|
||||
tool_name: str,
|
||||
tool_call_args: dict[str, Any],
|
||||
code: str,
|
||||
message: str,
|
||||
retryable: bool,
|
||||
) -> ToolResponse:
|
||||
output = build_error_output(
|
||||
tool_name=tool_name,
|
||||
tool_call_id=get_current_tool_call_id(tool_name=tool_name),
|
||||
code=code,
|
||||
message=message,
|
||||
retryable=retryable,
|
||||
)
|
||||
output = output.model_copy(update={"tool_call_args": tool_call_args})
|
||||
return build_tool_response(output)
|
||||
|
||||
|
||||
def _validate_runtime_context(
|
||||
*,
|
||||
tool_name: str,
|
||||
tool_call_args: dict[str, Any],
|
||||
session: Any,
|
||||
owner_id: Any,
|
||||
) -> ToolResponse | None:
|
||||
if session is None or owner_id is None:
|
||||
return _memory_error_output(
|
||||
tool_name=tool_name,
|
||||
tool_call_args=tool_call_args,
|
||||
code="MISSING_RUNTIME_ARGS",
|
||||
message="记忆工具缺少运行时参数",
|
||||
retryable=False,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _deep_merge_dict(base: dict[str, Any], patch: dict[str, Any]) -> dict[str, Any]:
|
||||
merged = deepcopy(base)
|
||||
for key, value in patch.items():
|
||||
if isinstance(value, dict) and isinstance(merged.get(key), dict):
|
||||
merged[key] = _deep_merge_dict(cast(dict[str, Any], merged[key]), value)
|
||||
else:
|
||||
merged[key] = value
|
||||
return merged
|
||||
|
||||
|
||||
def _remove_content_paths(
|
||||
base_payload: dict[str, Any],
|
||||
paths: list[str],
|
||||
) -> tuple[dict[str, Any], list[str]]:
|
||||
result = deepcopy(base_payload)
|
||||
removed: list[str] = []
|
||||
for raw_path in paths:
|
||||
path = raw_path.strip()
|
||||
if not path:
|
||||
continue
|
||||
keys = [part for part in path.split(".") if part]
|
||||
if not keys:
|
||||
continue
|
||||
if _delete_nested_path(result, keys):
|
||||
removed.append(path)
|
||||
return result, removed
|
||||
|
||||
|
||||
def _delete_nested_path(payload: dict[str, Any], keys: list[str]) -> bool:
|
||||
current: dict[str, Any] = payload
|
||||
for key in keys[:-1]:
|
||||
next_value = current.get(key)
|
||||
if not isinstance(next_value, dict):
|
||||
return False
|
||||
current = next_value
|
||||
leaf = keys[-1]
|
||||
if leaf in current:
|
||||
del current[leaf]
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
async def memory_write(
|
||||
memory_type: Annotated[
|
||||
str,
|
||||
Field(description="Memory type: user or work."),
|
||||
] = "user",
|
||||
user_content: Annotated[
|
||||
UserMemoryContent | None,
|
||||
Field(description="Patch payload for user memory content."),
|
||||
] = None,
|
||||
work_content: Annotated[
|
||||
WorkProfileContent | None,
|
||||
Field(description="Patch payload for work memory content."),
|
||||
] = None,
|
||||
session: Any = None,
|
||||
owner_id: Any = None,
|
||||
) -> ToolResponse:
|
||||
tool_name = "memory_write"
|
||||
tool_call_args: dict[str, Any] = {
|
||||
"memory_type": memory_type,
|
||||
"user_content": user_content,
|
||||
"work_content": work_content,
|
||||
}
|
||||
runtime_error = _validate_runtime_context(
|
||||
tool_name=tool_name,
|
||||
tool_call_args=tool_call_args,
|
||||
session=session,
|
||||
owner_id=owner_id,
|
||||
)
|
||||
if runtime_error is not None:
|
||||
return runtime_error
|
||||
|
||||
try:
|
||||
parsed_args = MemoryWriteArgs.model_validate(tool_call_args)
|
||||
service = create_memories_service(
|
||||
session=cast(AsyncSession, session),
|
||||
owner_id=cast(UUID, owner_id),
|
||||
)
|
||||
existing = await service.get_memory_model(memory_type=parsed_args.memory_type)
|
||||
|
||||
if parsed_args.memory_type == MemoryType.USER:
|
||||
base_model = (
|
||||
UserMemoryContent.model_validate(existing.content)
|
||||
if existing is not None
|
||||
else UserMemoryContent()
|
||||
)
|
||||
patch_model = cast(UserMemoryContent, parsed_args.user_content)
|
||||
merged = _deep_merge_dict(
|
||||
base_model.model_dump(),
|
||||
patch_model.model_dump(exclude_unset=True),
|
||||
)
|
||||
validated = UserMemoryContent.model_validate(merged)
|
||||
await service.update_user_memory(
|
||||
content=validated,
|
||||
)
|
||||
else:
|
||||
base_model = (
|
||||
WorkProfileContent.model_validate(existing.content)
|
||||
if existing is not None
|
||||
else WorkProfileContent()
|
||||
)
|
||||
patch_model = cast(WorkProfileContent, parsed_args.work_content)
|
||||
merged = _deep_merge_dict(
|
||||
base_model.model_dump(),
|
||||
patch_model.model_dump(exclude_unset=True),
|
||||
)
|
||||
validated = WorkProfileContent.model_validate(merged)
|
||||
await service.update_work_memory(
|
||||
content=validated,
|
||||
)
|
||||
|
||||
summary = f"status=success memory_type={parsed_args.memory_type.value}"
|
||||
return build_tool_response(
|
||||
ToolAgentOutput(
|
||||
tool_name=tool_name,
|
||||
tool_call_id=get_current_tool_call_id(tool_name=tool_name),
|
||||
tool_call_args=tool_call_args,
|
||||
status=ToolStatus.SUCCESS,
|
||||
result=summary,
|
||||
)
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
code, message, retryable = map_memory_exception(exc)
|
||||
return _memory_error_output(
|
||||
tool_name=tool_name,
|
||||
tool_call_args=tool_call_args,
|
||||
code=code,
|
||||
message=message,
|
||||
retryable=retryable,
|
||||
)
|
||||
|
||||
|
||||
async def memory_forget(
|
||||
memory_type: Annotated[
|
||||
str,
|
||||
Field(description="Memory type: user or work."),
|
||||
] = "user",
|
||||
forget_paths: Annotated[
|
||||
list[str] | None,
|
||||
Field(description="Dot paths to remove from content."),
|
||||
] = None,
|
||||
session: Any = None,
|
||||
owner_id: Any = None,
|
||||
) -> ToolResponse:
|
||||
tool_name = "memory_forget"
|
||||
tool_call_args: dict[str, Any] = {
|
||||
"memory_type": memory_type,
|
||||
"forget_paths": forget_paths or [],
|
||||
}
|
||||
runtime_error = _validate_runtime_context(
|
||||
tool_name=tool_name,
|
||||
tool_call_args=tool_call_args,
|
||||
session=session,
|
||||
owner_id=owner_id,
|
||||
)
|
||||
if runtime_error is not None:
|
||||
return runtime_error
|
||||
|
||||
try:
|
||||
parsed_args = MemoryForgetArgs.model_validate(tool_call_args)
|
||||
service = create_memories_service(
|
||||
session=cast(AsyncSession, session),
|
||||
owner_id=cast(UUID, owner_id),
|
||||
)
|
||||
existing = await service.get_memory_model(memory_type=parsed_args.memory_type)
|
||||
if existing is None:
|
||||
summary = f"status=success memory_type={parsed_args.memory_type.value} forgotten=0"
|
||||
return build_tool_response(
|
||||
ToolAgentOutput(
|
||||
tool_name=tool_name,
|
||||
tool_call_id=get_current_tool_call_id(tool_name=tool_name),
|
||||
tool_call_args=tool_call_args,
|
||||
status=ToolStatus.SUCCESS,
|
||||
result=summary,
|
||||
)
|
||||
)
|
||||
|
||||
if parsed_args.memory_type == MemoryType.USER:
|
||||
base_model = UserMemoryContent.model_validate(existing.content)
|
||||
updated_dict, removed_paths = _remove_content_paths(
|
||||
base_model.model_dump(),
|
||||
parsed_args.forget_paths,
|
||||
)
|
||||
validated = UserMemoryContent.model_validate(updated_dict)
|
||||
await service.update_user_memory(
|
||||
content=validated,
|
||||
)
|
||||
else:
|
||||
base_model = WorkProfileContent.model_validate(existing.content)
|
||||
updated_dict, removed_paths = _remove_content_paths(
|
||||
base_model.model_dump(),
|
||||
parsed_args.forget_paths,
|
||||
)
|
||||
validated = WorkProfileContent.model_validate(updated_dict)
|
||||
await service.update_work_memory(
|
||||
content=validated,
|
||||
)
|
||||
|
||||
summary = (
|
||||
f"status=success memory_type={parsed_args.memory_type.value} forgotten={len(removed_paths)} "
|
||||
f"skipped=0"
|
||||
)
|
||||
return build_tool_response(
|
||||
ToolAgentOutput(
|
||||
tool_name=tool_name,
|
||||
tool_call_id=get_current_tool_call_id(tool_name=tool_name),
|
||||
tool_call_args=tool_call_args,
|
||||
status=ToolStatus.SUCCESS,
|
||||
result=summary,
|
||||
)
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
code, message, retryable = map_memory_exception(exc)
|
||||
return _memory_error_output(
|
||||
tool_name=tool_name,
|
||||
tool_call_args=tool_call_args,
|
||||
code=code,
|
||||
message=message,
|
||||
retryable=retryable,
|
||||
)
|
||||
@@ -4,17 +4,13 @@ from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ToolGroup(str, Enum):
|
||||
READ = "read"
|
||||
EXECUTE = "execute"
|
||||
MEMORY = "memory"
|
||||
|
||||
|
||||
class AgentTool(str, Enum):
|
||||
CALENDAR_READ = "calendar.read"
|
||||
CALENDAR_WRITE = "calendar.write"
|
||||
CALENDAR_SHARE = "calendar.share"
|
||||
USER_LOOKUP = "user.lookup"
|
||||
MEMORY_WRITE = "memory.write"
|
||||
MEMORY_FORGET = "memory.forget"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -25,29 +21,32 @@ class ToolApprovalConfig:
|
||||
@dataclass(frozen=True)
|
||||
class ToolConfig:
|
||||
name: str
|
||||
group: ToolGroup
|
||||
approval: ToolApprovalConfig
|
||||
|
||||
|
||||
TOOL_CONFIGS: dict[str, ToolConfig] = {
|
||||
"calendar_read": ToolConfig(
|
||||
name="calendar_read",
|
||||
group=ToolGroup.READ,
|
||||
approval=ToolApprovalConfig(required=False),
|
||||
),
|
||||
"user_lookup": ToolConfig(
|
||||
name="user_lookup",
|
||||
group=ToolGroup.MEMORY,
|
||||
approval=ToolApprovalConfig(required=False),
|
||||
),
|
||||
"calendar_write": ToolConfig(
|
||||
name="calendar_write",
|
||||
group=ToolGroup.EXECUTE,
|
||||
approval=ToolApprovalConfig(required=False),
|
||||
),
|
||||
"calendar_share": ToolConfig(
|
||||
name="calendar_share",
|
||||
group=ToolGroup.EXECUTE,
|
||||
approval=ToolApprovalConfig(required=False),
|
||||
),
|
||||
"memory_write": ToolConfig(
|
||||
name="memory_write",
|
||||
approval=ToolApprovalConfig(required=False),
|
||||
),
|
||||
"memory_forget": ToolConfig(
|
||||
name="memory_forget",
|
||||
approval=ToolApprovalConfig(required=False),
|
||||
),
|
||||
}
|
||||
@@ -57,6 +56,8 @@ AGENT_TOOL_TO_FUNCTION_NAME: dict[AgentTool, str] = {
|
||||
AgentTool.CALENDAR_WRITE: "calendar_write",
|
||||
AgentTool.CALENDAR_SHARE: "calendar_share",
|
||||
AgentTool.USER_LOOKUP: "user_lookup",
|
||||
AgentTool.MEMORY_WRITE: "memory_write",
|
||||
AgentTool.MEMORY_FORGET: "memory_forget",
|
||||
}
|
||||
|
||||
TOOL_NAME_ALIASES: dict[str, AgentTool] = {
|
||||
@@ -68,6 +69,10 @@ TOOL_NAME_ALIASES: dict[str, AgentTool] = {
|
||||
"calendar_share": AgentTool.CALENDAR_SHARE,
|
||||
AgentTool.USER_LOOKUP.value: AgentTool.USER_LOOKUP,
|
||||
"user_lookup": AgentTool.USER_LOOKUP,
|
||||
AgentTool.MEMORY_WRITE.value: AgentTool.MEMORY_WRITE,
|
||||
"memory_write": AgentTool.MEMORY_WRITE,
|
||||
AgentTool.MEMORY_FORGET.value: AgentTool.MEMORY_FORGET,
|
||||
"memory_forget": AgentTool.MEMORY_FORGET,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,10 @@ from core.agentscope.tools.custom.calendar import (
|
||||
calendar_share,
|
||||
calendar_write,
|
||||
)
|
||||
from core.agentscope.tools.custom.memory import (
|
||||
memory_forget,
|
||||
memory_write,
|
||||
)
|
||||
from core.agentscope.tools.custom.user_lookup import user_lookup
|
||||
from core.agentscope.tools.tool_config import (
|
||||
TOOL_CONFIGS,
|
||||
@@ -23,6 +27,8 @@ TOOL_FUNCTIONS: dict[str, Any] = {
|
||||
"calendar_write": calendar_write,
|
||||
"calendar_share": calendar_share,
|
||||
"user_lookup": user_lookup,
|
||||
"memory_write": memory_write,
|
||||
"memory_forget": memory_forget,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from core.auth.models import CurrentUser
|
||||
from v1.memories.repository import SQLAlchemyMemoriesRepository
|
||||
from v1.memories.service import MemoriesService
|
||||
|
||||
|
||||
def create_memories_service(
|
||||
session: AsyncSession,
|
||||
owner_id: UUID,
|
||||
) -> MemoriesService:
|
||||
return MemoriesService(
|
||||
repository=SQLAlchemyMemoriesRepository(session),
|
||||
session=session,
|
||||
current_user=CurrentUser(id=owner_id),
|
||||
)
|
||||
|
||||
|
||||
def map_memory_exception(exc: Exception) -> tuple[str, str, bool]:
|
||||
if isinstance(exc, HTTPException):
|
||||
detail = exc.detail
|
||||
if isinstance(detail, str) and detail.strip():
|
||||
return "OPERATION_FAILED", detail.strip(), exc.status_code >= 500
|
||||
return "OPERATION_FAILED", "记忆操作失败", exc.status_code >= 500
|
||||
if isinstance(exc, ValueError):
|
||||
return "INVALID_ARGUMENT", "请求参数无效", False
|
||||
return "INTERNAL_ERROR", "记忆操作失败", True
|
||||
Reference in New Issue
Block a user