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
@@ -1,11 +1,15 @@
from core.agentscope.prompts.agent_prompt import build_agent_prompt
from core.agentscope.prompts.memory_prompt import build_memory_prompt
from core.agentscope.prompts.memory_prompt import (
build_user_memory_prompt,
build_work_memory_prompt,
)
from core.agentscope.prompts.system_prompt import build_system_prompt
from core.agentscope.prompts.tool_prompt import build_tools_prompt
__all__ = [
"build_agent_prompt",
"build_memory_prompt",
"build_user_memory_prompt",
"build_work_memory_prompt",
"build_system_prompt",
"build_tools_prompt",
]
@@ -1,52 +1,59 @@
from __future__ import annotations
import json
from typing import Any
from schemas.memories import MemoryContext, MemoryListResponse
from schemas.memories.memory_content import UserMemoryContent, WorkProfileContent
def _wrap_section(section: str, content: str) -> str:
marker_map = {
"memory": ("<!-- MEMORY_START -->", "<!-- MEMORY_END -->"),
"user_memory": ("<!-- USER_MEMORY_START -->", "<!-- USER_MEMORY_END -->"),
"work_memory": ("<!-- WORK_MEMORY_START -->", "<!-- WORK_MEMORY_END -->"),
}
start, end = marker_map[section]
body = content.strip()
return f"{start}\n{body}\n{end}" if body else f"{start}\n{end}"
def _format_memory_content(content: dict[str, Any]) -> str:
def _format_content(content: UserMemoryContent | WorkProfileContent) -> str:
if isinstance(content, dict):
return json.dumps(content, ensure_ascii=True, separators=(",", ":"))
return str(content)
return json.dumps(
content.model_dump(mode="json"), ensure_ascii=True, separators=(",", ":")
)
def _format_memory(ctx: MemoryContext) -> str:
parts = [
f"[{ctx.memory_type.value.upper()}] {ctx.title or 'Untitled'}",
f" source: {ctx.source.value}",
f" content: {_format_memory_content(ctx.content)}",
]
if ctx.created_at:
parts.append(f" created_at: {ctx.created_at.isoformat()}")
return "\n".join(parts)
def build_memory_prompt(
def build_user_memory_prompt(
*,
memories: MemoryListResponse,
user_memory: UserMemoryContent | None,
) -> str | None:
if not memories.memories:
if user_memory is None:
return None
lines: list[str] = [
"[User Memories]",
"- Memories are persistent context from previous sessions.",
"- Use them to ground responses in known user facts and preferences.",
"- Do not invent facts not present in memories.",
"[User Memory]",
"- User memory contains personal preferences, habits, people, and places.",
"- Use this to understand the user's personal context and preferences.",
"- Do not invent facts not present here.",
f"content: {_format_content(user_memory)}",
]
for ctx in memories.memories:
lines.append(_format_memory(ctx))
return _wrap_section("user_memory", "\n".join(lines))
return _wrap_section("memory", "\n".join(lines))
def build_work_memory_prompt(
*,
work_memory: WorkProfileContent | None,
) -> str | None:
if work_memory is None:
return None
lines: list[str] = [
"[Work Memory]",
"- Work memory contains projects, team members, habits, and milestones.",
"- Use this to understand the user's work context and ongoing tasks.",
"- Do not invent facts not present here.",
f"content: {_format_content(work_memory)}",
]
return _wrap_section("work_memory", "\n".join(lines))
@@ -9,12 +9,15 @@ from ag_ui.core.types import Tool
from core.agentscope.prompts.agent_prompt import (
build_agent_prompt,
)
from core.agentscope.prompts.memory_prompt import build_memory_prompt
from core.agentscope.prompts.memory_prompt import (
build_user_memory_prompt,
build_work_memory_prompt,
)
from core.agentscope.prompts.route_prompt import build_frontend_route_prompt
from core.agentscope.prompts.tool_prompt import build_tools_prompt
from schemas.agent.system_agent import AgentType, SystemAgentLLMConfig
from schemas.agent.forwarded_props import ClientTimeContext
from schemas.memories import MemoryListResponse
from schemas.memories.memory_content import UserMemoryContent, WorkProfileContent
from schemas.user.context import UserContext
@@ -210,9 +213,16 @@ def build_system_prompt(
runtime_client_time: ClientTimeContext | None = None,
extra_context: str | None = None,
tools: Sequence[Tool | dict[str, Any]] | None = None,
memories: MemoryListResponse | None = None,
user_memory: UserMemoryContent | None = None,
work_memory: WorkProfileContent | None = None,
) -> str:
include_route_section = agent_type == AgentType.WORKER
if agent_type == AgentType.ROUTER:
memory_prompt = build_user_memory_prompt(user_memory=user_memory)
else:
memory_prompt = build_work_memory_prompt(work_memory=work_memory)
sections: list[str | None] = [
_build_identity_section(),
_build_env_section(
@@ -228,7 +238,7 @@ def build_system_prompt(
llm_config=llm_config,
),
build_tools_prompt(tools=tools) if tools else None,
build_memory_prompt(memories=memories) if memories else None,
memory_prompt,
_build_output_rules(),
]
return "\n\n".join(item for item in sections if item).strip()
@@ -7,7 +7,7 @@ from agentscope.message import Msg
from core.agentscope.runtime.runner import AgentScopeRunner
from core.logging import get_logger
from schemas.automation import RuntimeConfig
from schemas.memories import MemoryListResponse
from schemas.memories.memory_content import UserMemoryContent, WorkProfileContent
from schemas.user import UserContext
logger = get_logger("core.agentscope.runtime.orchestrator")
@@ -26,7 +26,8 @@ class RunnerLike(Protocol):
pipeline: PipelineLike,
run_input: RunAgentInput,
runtime_config: RuntimeConfig,
memories: MemoryListResponse | None,
user_memory: UserMemoryContent | None,
work_memory: WorkProfileContent | None,
) -> dict[str, Any]: ...
@@ -50,7 +51,8 @@ class AgentScopeRuntimeOrchestrator:
context_messages: list[Msg],
user_context: UserContext,
runtime_config: RuntimeConfig,
memories: MemoryListResponse | None = None,
user_memory: UserMemoryContent | None = None,
work_memory: WorkProfileContent | None = None,
) -> dict[str, Any]:
thread_id = run_input.thread_id
run_id = run_input.run_id
@@ -70,7 +72,8 @@ class AgentScopeRuntimeOrchestrator:
pipeline=self._pipeline,
run_input=run_input,
runtime_config=runtime_config,
memories=memories,
user_memory=user_memory,
work_memory=work_memory,
)
await self._pipeline.emit(
+13 -12
View File
@@ -41,7 +41,7 @@ from schemas.agent.system_agent import (
SystemAgentLLMConfig,
)
from schemas.automation import RuntimeConfig
from schemas.memories import MemoryListResponse
from schemas.memories.memory_content import UserMemoryContent, WorkProfileContent
from schemas.user import UserContext
from services.litellm.service import LiteLLMService
from sqlalchemy import select
@@ -71,7 +71,8 @@ class AgentScopeRunner:
pipeline: PipelineLike,
run_input: RunAgentInput,
runtime_config: RuntimeConfig,
memories: MemoryListResponse | None = None,
user_memory: UserMemoryContent | None = None,
work_memory: WorkProfileContent | None = None,
) -> dict[str, Any]:
owner_id = UUID(user_context.id)
runtime_client_time = self._resolve_runtime_client_time(run_input=run_input)
@@ -98,7 +99,7 @@ class AgentScopeRunner:
context_messages=context_messages,
stage_config=router_config,
runtime_client_time=runtime_client_time,
memories=memories,
user_memory=user_memory,
)
worker_output = await self._execute_worker_step(
pipeline=pipeline,
@@ -108,7 +109,7 @@ class AgentScopeRunner:
toolkit=worker_toolkit,
stage_config=worker_config,
runtime_client_time=runtime_client_time,
memories=memories,
work_memory=work_memory,
)
return {
"router": router_output.model_dump(mode="json", exclude_none=True),
@@ -166,7 +167,7 @@ class AgentScopeRunner:
context_messages: list[Msg],
stage_config: SystemAgentRuntimeConfig,
runtime_client_time: ClientTimeContext | None,
memories: MemoryListResponse | None,
user_memory: UserMemoryContent | None,
) -> RouterAgentOutput:
await self._emit_step_event(
pipeline=pipeline,
@@ -179,7 +180,7 @@ class AgentScopeRunner:
context_messages=context_messages,
stage_config=stage_config,
runtime_client_time=runtime_client_time,
memories=memories,
user_memory=user_memory,
run_input=run_input,
)
router_output = RouterAgentOutput.model_validate(router_result.payload)
@@ -201,7 +202,7 @@ class AgentScopeRunner:
toolkit: Any,
stage_config: SystemAgentRuntimeConfig,
runtime_client_time: ClientTimeContext | None,
memories: MemoryListResponse | None,
work_memory: WorkProfileContent | None,
) -> WorkerAgentOutputLite:
worker_output_model = resolve_worker_output_model(router_output.ui.ui_mode)
await self._emit_step_event(
@@ -221,7 +222,7 @@ class AgentScopeRunner:
worker_output_model=worker_output_model,
pipeline=pipeline,
runtime_client_time=runtime_client_time,
memories=memories,
work_memory=work_memory,
)
worker_output = worker_output_model.model_validate(worker_result.payload)
await self._emit_step_event(
@@ -239,7 +240,7 @@ class AgentScopeRunner:
context_messages: list[Msg],
stage_config: SystemAgentRuntimeConfig,
runtime_client_time: ClientTimeContext | None,
memories: MemoryListResponse | None,
user_memory: UserMemoryContent | None,
run_input: RunAgentInput,
) -> StageExecutionResult:
messages_for_router = self._build_router_messages(
@@ -260,7 +261,7 @@ class AgentScopeRunner:
now_utc=datetime.now(timezone.utc),
runtime_client_time=runtime_client_time,
tools=None,
memories=memories,
user_memory=user_memory,
),
"system",
),
@@ -319,7 +320,7 @@ class AgentScopeRunner:
worker_output_model: type[WorkerAgentOutputLite],
pipeline: PipelineLike,
runtime_client_time: ClientTimeContext | None,
memories: MemoryListResponse | None,
work_memory: WorkProfileContent | None,
) -> StageExecutionResult:
tracking_model = self._build_model(stage_config=stage_config)
emitter = PipelineStageEmitter(
@@ -340,7 +341,7 @@ class AgentScopeRunner:
runtime_client_time=runtime_client_time,
extra_context=stage_config.extra_context,
tools=None,
memories=memories,
work_memory=work_memory,
),
toolkit=toolkit,
model=tracking_model,
+14 -8
View File
@@ -20,8 +20,8 @@ from core.config.settings import config
from core.db.session import AsyncSessionLocal
from core.logging import get_logger
from core.taskiq.app import worker_agent_broker, worker_automation_broker
from schemas.automation import MemoryContextConfig, RuntimeConfig
from schemas.memories import MemoryListResponse
from schemas.automation import MessageContextConfig, RuntimeConfig
from schemas.memories.memory_content import UserMemoryContent, WorkProfileContent
from schemas.messages.chat_message import (
AgentChatMessageMetadata,
extract_user_message_attachments,
@@ -30,7 +30,7 @@ from schemas.user import UserContext
from services.base.redis import get_or_init_redis_client
from services.base.supabase import supabase_service
from v1.agent.repository import AgentRepository
from v1.memories.repository import MemoriesRepository
from v1.memories.repository import SQLAlchemyMemoriesRepository
from v1.memories.service import MemoriesService
from v1.users.dependencies import get_user_service
@@ -83,7 +83,7 @@ async def _build_recent_context_messages(
*,
session: Any,
thread_id: str,
context_config: "MemoryContextConfig",
context_config: "MessageContextConfig",
) -> list[Msg]:
context_service = AgentContextService(repository=AgentRepository(session))
result = await context_service.load_context_messages(
@@ -194,11 +194,16 @@ async def run_agentscope_task(command: dict[str, Any]) -> dict[str, object]:
orchestrator = _load_runtime()
async with AsyncSessionLocal() as session:
current_user = CurrentUser(id=owner_id)
user_context = await _build_user_context(owner_id=owner_id, session=session)
memories_service = MemoriesService(MemoriesRepository(session))
memories: MemoryListResponse = await memories_service.get_all_memories(
owner_id=owner_id
memories_service = MemoriesService(
repository=SQLAlchemyMemoriesRepository(session),
session=session,
current_user=current_user,
)
memories_result = await memories_service.get_all_memories()
user_memory: UserMemoryContent | None = memories_result.get("user_memory")
work_memory: WorkProfileContent | None = memories_result.get("work_memory")
redis_client = await get_or_init_redis_client()
bus = RedisStreamBus(
@@ -229,7 +234,8 @@ async def run_agentscope_task(command: dict[str, Any]) -> dict[str, object]:
context_messages=context_messages,
user_context=user_context,
runtime_config=runtime_config,
memories=memories,
user_memory=user_memory,
work_memory=work_memory,
)
logger.info(
"agentscope runtime task completed",
@@ -6,7 +6,7 @@ from typing import Any, Protocol
from schemas.agent.visibility import SystemVisibilityBit, bit_mask
from schemas.automation import ContextWindowMode, MemoryContextConfig
from schemas.automation import ContextWindowMode, MessageContextConfig
_DEFAULT_CONTEXT_WINDOW_USER_MESSAGES = 20
@@ -86,7 +86,7 @@ class AgentContextService:
self,
*,
thread_id: str,
context_config: MemoryContextConfig,
context_config: MessageContextConfig,
) -> dict[str, object] | None:
visibility_mask = bit_mask(bit=int(SystemVisibilityBit.CONTEXT_ASSEMBLY))
context_loader = CONTEXT_LOADER_REGISTRY.resolve(
@@ -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