feat: 重构 memory 系统,支持 user memory 和 work memory 分离
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user