feat: 重构 memory 系统,支持 user memory 和 work memory 分离

This commit is contained in:
qzl
2026-03-23 14:25:47 +08:00
parent 3aacc756db
commit 6be616f108
70 changed files with 7031 additions and 431 deletions
@@ -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