diff --git a/docs/en/api/01-overview.md b/docs/en/api/01-overview.md index d1bbeef6fa..9cefab8e17 100644 --- a/docs/en/api/01-overview.md +++ b/docs/en/api/01-overview.md @@ -542,6 +542,8 @@ This catalog follows the routes actually mounted by the server. Each group headi | DELETE | `/api/v1/admin/accounts/{account_id}` | Delete an account | | POST | `/api/v1/admin/accounts/{account_id}/users` | Register a user | | GET | `/api/v1/admin/accounts/{account_id}/users` | List users | +| GET | `/api/v1/admin/accounts/{account_id}/users/{user_id}/settings` | Get a user's memory policy | +| PATCH | `/api/v1/admin/accounts/{account_id}/users/{user_id}/settings` | Update a user's memory policy | | DELETE | `/api/v1/admin/accounts/{account_id}/users/{user_id}` | Remove a user | | PUT | `/api/v1/admin/accounts/{account_id}/users/{user_id}/role` | Promote a user to ADMIN | | POST | `/api/v1/admin/accounts/{account_id}/users/{user_id}/key` | Regenerate a user key | diff --git a/docs/en/api/08-admin.md b/docs/en/api/08-admin.md index dc9499ec5e..e0f23d08b6 100644 --- a/docs/en/api/08-admin.md +++ b/docs/en/api/08-admin.md @@ -117,6 +117,31 @@ Content-Type: application/json Before an existing setting is replaced, it is backed up to `/local/{account_id}/_system/setting.backup.json`. +### user_settings + +ROOT can manage any User and ADMIN can manage Users in its own account. The +User settings endpoint currently allowlists only `memory_policy`. Each target +has its own `enabled` switch and `memory_types` filter. Agent memory types are +self-only; configuring them for `peer` is rejected. + +```http +GET /api/v1/admin/accounts/{account_id}/users/{user_id}/settings +PATCH /api/v1/admin/accounts/{account_id}/users/{user_id}/settings +Content-Type: application/json + +{ + "memory_policy": { + "self": {"enabled": true, "memory_types": ["profile", "experiences"]}, + "peer": {"enabled": false, "memory_types": []} + } +} +``` + +The response contains the explicit `overrides`, the effective `settings`, and +the account-level `agent_evolution_enabled` switch. Updates are backed up to +the User's `settings/user_config.backup.json` before replacement. A Session +without an explicit policy reads the latest User policy when it is committed. + --- ### create_account diff --git a/docs/zh/api/01-overview.md b/docs/zh/api/01-overview.md index 432433e317..0dce92717b 100644 --- a/docs/zh/api/01-overview.md +++ b/docs/zh/api/01-overview.md @@ -537,6 +537,8 @@ JSON 输出 - 错误: | DELETE | `/api/v1/admin/accounts/{account_id}` | 删除账号 | | POST | `/api/v1/admin/accounts/{account_id}/users` | 注册用户 | | GET | `/api/v1/admin/accounts/{account_id}/users` | 列出用户 | +| GET | `/api/v1/admin/accounts/{account_id}/users/{user_id}/settings` | 获取用户记忆策略 | +| PATCH | `/api/v1/admin/accounts/{account_id}/users/{user_id}/settings` | 更新用户记忆策略 | | DELETE | `/api/v1/admin/accounts/{account_id}/users/{user_id}` | 移除用户 | | PUT | `/api/v1/admin/accounts/{account_id}/users/{user_id}/role` | 将用户提升为 ADMIN | | POST | `/api/v1/admin/accounts/{account_id}/users/{user_id}/key` | 重新生成用户 Key | diff --git a/docs/zh/api/08-admin.md b/docs/zh/api/08-admin.md index 52967bce0e..12cad597fa 100644 --- a/docs/zh/api/08-admin.md +++ b/docs/zh/api/08-admin.md @@ -115,6 +115,30 @@ Content-Type: application/json 覆盖已有配置前,内核会先备份到 `/local/{account_id}/_system/setting.backup.json`。 +### user_settings + +ROOT 可管理任意 User,ADMIN 仅可管理所属 account 内的 User。User 配置接口当前 +仅允许修改 `memory_policy`。`self` 和 `peer` 分别配置 `enabled` 与 +`memory_types`;Agent 记忆只允许写入 self,peer 配置 Agent 记忆类型会被拒绝。 + +```http +GET /api/v1/admin/accounts/{account_id}/users/{user_id}/settings +PATCH /api/v1/admin/accounts/{account_id}/users/{user_id}/settings +Content-Type: application/json + +{ + "memory_policy": { + "self": {"enabled": true, "memory_types": ["profile", "experiences"]}, + "peer": {"enabled": false, "memory_types": []} + } +} +``` + +响应同时返回显式 `overrides`、生效 `settings` 和 account 级 +`agent_evolution_enabled`。更新前会备份到该 User 的 +`settings/user_config.backup.json`。未显式配置策略的 Session 在 commit 时读取 +该 User 最新策略。 + --- ### create_account diff --git a/openviking/server/config.py b/openviking/server/config.py index 3eb6dfd70e..e3545b5237 100644 --- a/openviking/server/config.py +++ b/openviking/server/config.py @@ -3,7 +3,7 @@ """Server configuration for OpenViking HTTP Server.""" import sys -from typing import Dict, List, Literal, Optional +from typing import Any, Dict, List, Literal, Optional from pydantic import BaseModel, Field, ValidationError, field_validator @@ -103,6 +103,7 @@ class UserConfig(BaseModel): """User configuration values that can be defaulted or initialized.""" add_targets: AddTargetsConfig = Field(default_factory=AddTargetsConfig) + memory_policy: Optional[Dict[str, Any]] = None agent_evolution: DeprecatedUserAgentEvolutionConfig = Field( default_factory=DeprecatedUserAgentEvolutionConfig, exclude=True, @@ -110,6 +111,15 @@ class UserConfig(BaseModel): model_config = {"extra": "forbid"} + @field_validator("memory_policy", mode="before") + @classmethod + def validate_memory_policy(cls, value: Any) -> Optional[Dict[str, Any]]: + if value is None: + return None + from openviking.session.memory_policy import MemoryPolicy + + return MemoryPolicy.from_dict(value).to_dict() + class MetricsAccountDimensionConfig(BaseModel): """Account-dimension configuration for metrics label injection.""" diff --git a/openviking/server/routers/admin.py b/openviking/server/routers/admin.py index e31cbf6a10..770dfde7d6 100644 --- a/openviking/server/routers/admin.py +++ b/openviking/server/routers/admin.py @@ -25,7 +25,14 @@ from openviking.server.dependencies import get_service from openviking.server.identity import RequestContext, Role from openviking.server.models import Response -from openviking.server.user_config import validate_add_targets, write_user_config +from openviking.server.user_config import ( + delete_user_config, + read_user_config, + validate_add_targets, + validate_user_memory_policy, + write_user_config, + write_user_memory_policy, +) from openviking.service.legacy_migration import LegacyDataMigration from openviking.service.task_store import ( SYSTEM_TASK_ACCOUNT_ID, @@ -34,6 +41,8 @@ from openviking.service.task_tracker import ( get_task_tracker, ) +from openviking.session.memory.memory_type_registry import MemoryTypeRegistry +from openviking.session.memory_policy import MemoryPolicy from openviking.storage.viking_fs import get_viking_fs from openviking_cli.exceptions import ( FailedPreconditionError, @@ -80,6 +89,12 @@ class SetAgentEvolutionRequest(BaseModel): enabled: bool +class UserSettingsPatch(BaseModel): + memory_policy: dict + + model_config = {"extra": "forbid"} + + def _agent_evolution_account_id(ctx: RequestContext) -> str: if ctx.role == Role.ROOT: return get_openviking_config().default_account @@ -177,7 +192,7 @@ def _has_add_targets(user_config: UserConfig | None) -> bool: def _has_initial_user_config(user_config: UserConfig | None) -> bool: - return _has_add_targets(user_config) + return bool(_has_add_targets(user_config) or (user_config and user_config.memory_policy)) def _validate_initial_user_config( @@ -185,15 +200,18 @@ def _validate_initial_user_config( user_ctx: RequestContext, user_config: UserConfig | None, ) -> None: - if not _has_add_targets(user_config): + if not _has_initial_user_config(user_config): return if service.viking_fs is None: raise FailedPreconditionError("OpenViking service is not initialized.") - validate_add_targets( - user_config.add_targets, - ctx=user_ctx, - viking_fs=service.viking_fs, - ) + if _has_add_targets(user_config): + validate_add_targets( + user_config.add_targets, + ctx=user_ctx, + viking_fs=service.viking_fs, + ) + if user_config is not None: + validate_user_memory_policy(user_config.memory_policy) async def _write_initial_user_config( @@ -206,6 +224,34 @@ async def _write_initial_user_config( await write_user_config(service.viking_fs, user_ctx, user_config) +def _check_user_exists(request: Request, account_id: str, user_id: str) -> None: + manager = _get_api_key_manager(request) + if not manager.has_user(account_id, user_id): + raise NotFoundError(user_id, "user") + + +async def _user_settings_result( + account_id: str, + user_id: str, + user_config: UserConfig, +) -> dict: + configured = user_config.memory_policy + policy = MemoryPolicy.from_dict(configured) + policy.validate_memory_types(set(MemoryTypeRegistry().list_names(include_disabled=False))) + enabled = await get_service().sessions.get_agent_evolution_enabled(account_id) + effective = policy.resolve( + set(MemoryTypeRegistry().list_names(include_disabled=False)), + agent_evolution_enabled=enabled, + ) + return { + "account_id": account_id, + "user_id": user_id, + "settings": {"memory_policy": effective.to_dict()}, + "overrides": {"memory_policy": configured} if configured is not None else {}, + "agent_evolution_enabled": enabled, + } + + async def _run_legacy_migration_task( task_id: str, migration: LegacyDataMigration, @@ -435,8 +481,35 @@ async def register_user( str(resolved_role), seed=body.seed, ) - await service.initialize_user_directories(user_ctx) - await _write_initial_user_config(service, user_ctx, body.user_config) + try: + await service.initialize_user_directories(user_ctx) + await _write_initial_user_config(service, user_ctx, body.user_config) + except Exception: + logger.exception( + "Failed to initialize user %s in account %s; rolling back registration", + body.user_id, + account_id, + ) + if service.viking_fs is not None: + try: + await delete_user_config(service.viking_fs, user_ctx) + except Exception: + logger.warning( + "Failed to remove user config while rolling back %s/%s", + account_id, + body.user_id, + exc_info=True, + ) + try: + await manager.remove_user(account_id, body.user_id) + except Exception: + logger.warning( + "Failed to remove user registration while rolling back %s/%s", + account_id, + body.user_id, + exc_info=True, + ) + raise result = { "account_id": account_id, "user_id": body.user_id, @@ -466,6 +539,60 @@ async def list_users( return Response(status="ok", result=users) +@router.get("/accounts/{account_id}/users/{user_id}/settings") +@require_auth_root_or_admin +async def get_user_settings( + request: Request, + account_id: str = Path(..., description="Account ID"), + user_id: str = Path(..., description="User ID"), + ctx: RequestContext = Depends(get_request_context), +): + """Return the configured and effective memory policy for one User.""" + _check_account_access(ctx, account_id) + _check_account_exists(request, account_id) + _check_user_exists(request, account_id, user_id) + service = get_service() + if service.viking_fs is None: + raise FailedPreconditionError("OpenViking service is not initialized.") + user_ctx = RequestContext( + user=UserIdentifier(account_id, user_id), + role=Role.USER, + ) + user_config = await read_user_config(service.viking_fs, user_ctx) + return Response( + status="ok", + result=await _user_settings_result(account_id, user_id, user_config), + ) + + +@router.patch("/accounts/{account_id}/users/{user_id}/settings") +@require_auth_root_or_admin +async def patch_user_settings( + body: UserSettingsPatch, + request: Request, + account_id: str = Path(..., description="Account ID"), + user_id: str = Path(..., description="User ID"), + ctx: RequestContext = Depends(get_request_context), +): + """Update the allowlisted User memory policy without restarting the server.""" + _check_account_access(ctx, account_id) + _check_account_exists(request, account_id) + _check_user_exists(request, account_id, user_id) + service = get_service() + if service.viking_fs is None: + raise FailedPreconditionError("OpenViking service is not initialized.") + user_ctx = RequestContext( + user=UserIdentifier(account_id, user_id), + role=Role.USER, + ) + await write_user_memory_policy(service.viking_fs, user_ctx, body.memory_policy) + user_config = await read_user_config(service.viking_fs, user_ctx) + return Response( + status="ok", + result=await _user_settings_result(account_id, user_id, user_config), + ) + + @router.delete("/accounts/{account_id}/users/{user_id}") @require_auth_root_or_admin async def remove_user( diff --git a/openviking/server/user_config.py b/openviking/server/user_config.py index 908d5a8587..02380ebffd 100644 --- a/openviking/server/user_config.py +++ b/openviking/server/user_config.py @@ -33,6 +33,10 @@ def user_config_uri(ctx: RequestContext) -> str: return f"{canonical_user_root(ctx)}/settings/user_config.json" +def user_config_backup_uri(ctx: RequestContext) -> str: + return f"{canonical_user_root(ctx)}/settings/user_config.backup.json" + + @asynccontextmanager async def _user_config_lock( viking_fs: VikingFS, @@ -69,6 +73,17 @@ def _ensure_mutable(viking_fs: VikingFS, uri: str, ctx: RequestContext) -> None: ensure(uri, ctx) +def validate_user_memory_policy(memory_policy: Optional[dict[str, Any]]) -> None: + if memory_policy is None: + return + from openviking.session.memory.memory_type_registry import MemoryTypeRegistry + from openviking.session.memory_policy import MemoryPolicy + + MemoryPolicy.from_dict(memory_policy).validate_memory_types( + set(MemoryTypeRegistry().list_names(include_disabled=False)) + ) + + def validate_resource_add_target( uri: str, *, @@ -151,17 +166,33 @@ async def update_user_config( before = current.model_dump() result = updater(current) validate_add_targets(current.add_targets, ctx=ctx, viking_fs=viking_fs) + validate_user_memory_policy(current.memory_policy) if current.model_dump() != before: + before_content = json.dumps(before, ensure_ascii=False, sort_keys=True) await viking_fs.write_file( - uri, - json.dumps( - current.model_dump(exclude_none=True), - ensure_ascii=False, - sort_keys=True, - ), + user_config_backup_uri(ctx), + before_content, ctx=ctx, - lease_ref=handle, ) + try: + await viking_fs.write_file( + uri, + json.dumps( + current.model_dump(exclude_none=True), + ensure_ascii=False, + sort_keys=True, + ), + ctx=ctx, + lease_ref=handle, + ) + except Exception: + await viking_fs.write_file( + uri, + before_content, + ctx=ctx, + lease_ref=handle, + ) + raise return result @@ -171,6 +202,7 @@ async def write_user_config( user_config: UserConfig, ) -> ResolvedAddTargets: runtime = validate_add_targets(user_config.add_targets, ctx=ctx, viking_fs=viking_fs) + validate_user_memory_policy(user_config.memory_policy) uri = user_config_uri(ctx) async with _user_config_lock(viking_fs, uri, ctx) as handle: await viking_fs.write_file( @@ -221,6 +253,30 @@ def _clear(user_config: UserConfig) -> None: await update_user_config(viking_fs, ctx, _clear) +async def read_user_memory_policy( + viking_fs: VikingFS, + ctx: RequestContext, +) -> Optional[dict[str, Any]]: + return (await read_user_config(viking_fs, ctx)).memory_policy + + +async def write_user_memory_policy( + viking_fs: VikingFS, + ctx: RequestContext, + memory_policy: dict[str, Any], +) -> dict[str, Any]: + normalized = UserConfig(memory_policy=memory_policy).memory_policy + if normalized is None: + raise InvalidArgumentError("memory_policy must be an object") + validate_user_memory_policy(normalized) + + def _set(user_config: UserConfig) -> None: + user_config.memory_policy = normalized + + await update_user_config(viking_fs, ctx, _set) + return normalized + + async def effective_resource_add_target( *, viking_fs: VikingFS, diff --git a/openviking/service/session_service.py b/openviking/service/session_service.py index fbd870eb09..7f6a0132b0 100644 --- a/openviking/service/session_service.py +++ b/openviking/service/session_service.py @@ -15,6 +15,7 @@ from openviking.server.agent_evolution_config import AgentEvolutionConfigProvider from openviking.server.config import AgentEvolutionConfig, ToolOutputExternalizationConfig from openviking.server.identity import RequestContext +from openviking.server.user_config import read_user_config from openviking.service.session_auto_commit import ( compute_next_check_at, get_idle_timeout_seconds, @@ -203,9 +204,15 @@ def session( agent_evolution_enabled_provider=lambda: self.get_agent_evolution_enabled( ctx.account_id ), + memory_policy_provider=lambda: self._get_user_memory_policy(ctx), usage_reporter=self._usage_reporter, ) + async def _get_user_memory_policy(self, ctx: RequestContext) -> Optional[Dict[str, Any]]: + """Read the latest persisted User policy at commit time.""" + config = await read_user_config(self._viking_fs, ctx) + return config.memory_policy + async def create( self, ctx: RequestContext, diff --git a/openviking/session/compressor_v3.py b/openviking/session/compressor_v3.py index ba93068743..935c64f70a 100644 --- a/openviking/session/compressor_v3.py +++ b/openviking/session/compressor_v3.py @@ -331,6 +331,8 @@ async def extract_long_term_memories( latest_archive_overview: str = "", archive_uri: Optional[str] = None, allowed_memory_types: Optional[set[str]] = None, + allowed_self_memory_types: Optional[set[str]] = None, + allowed_peer_memory_types: Optional[set[str]] = None, agent_evolution_enabled: bool = True, allow_self_memory: bool = True, allowed_peer_ids: Optional[set[str]] = None, @@ -343,9 +345,22 @@ async def extract_long_term_memories( else set(allowed_memory_types) ) allowed_memory_types = effective_types - AGENT_EVOLUTION_MEMORY_TYPES + if allowed_self_memory_types is not None: + allowed_self_memory_types = ( + set(allowed_self_memory_types) - AGENT_EVOLUTION_MEMORY_TYPES + ) + if allowed_peer_memory_types is not None: + allowed_peer_memory_types = ( + set(allowed_peer_memory_types) - AGENT_EVOLUTION_MEMORY_TYPES + ) message_list = list(messages) - fast_path_case = _training_case_from_first_message(message_list, allowed_memory_types) + self_types = ( + allowed_self_memory_types + if allowed_self_memory_types is not None + else allowed_memory_types + ) + fast_path_case = _training_case_from_first_message(message_list, self_types) if fast_path_case is not None: return await self._commit_training_case_fast_path( case=fast_path_case, @@ -355,7 +370,7 @@ async def extract_long_term_memories( archive_uri=archive_uri or "", strict_extract_errors=strict_extract_errors, agent_evolution_enabled=agent_evolution_enabled, - allowed_memory_types=allowed_memory_types, + allowed_memory_types=self_types, ) result = await self._extract_user_memories( @@ -367,12 +382,14 @@ async def extract_long_term_memories( latest_archive_overview=latest_archive_overview, archive_uri=archive_uri, allowed_memory_types=allowed_memory_types, + allowed_self_memory_types=allowed_self_memory_types, + allowed_peer_memory_types=allowed_peer_memory_types, allow_self_memory=allow_self_memory, allowed_peer_ids=allowed_peer_ids, event_search_tags=event_search_tags, ) - agent_memory_types = _allowed_agent_memory_types(allowed_memory_types) - cases_allowed = allowed_memory_types is None or _CASES_MEMORY_TYPE in allowed_memory_types + agent_memory_types = _allowed_agent_memory_types(self_types) + cases_allowed = self_types is None or _CASES_MEMORY_TYPE in self_types session_skills_enabled = self._session_skill_extraction_enabled() if ( agent_evolution_enabled @@ -561,6 +578,8 @@ async def _extract_user_memories( latest_archive_overview: str = "", archive_uri: Optional[str] = None, allowed_memory_types: Optional[set[str]] = None, + allowed_self_memory_types: Optional[set[str]] = None, + allowed_peer_memory_types: Optional[set[str]] = None, allow_self_memory: bool = True, allowed_peer_ids: Optional[set[str]] = None, event_search_tags: Optional[List[str]] = None, @@ -582,7 +601,11 @@ async def _extract_user_memories( if allow_self_memory: await registry.initialize_memory_files( ctx, - allowed_memory_types=allowed_memory_types, + allowed_memory_types=( + allowed_self_memory_types + if allowed_self_memory_types is not None + else allowed_memory_types + ), ) extract_context = ExtractContext(messages) @@ -592,6 +615,8 @@ async def _extract_user_memories( allowed_memory_types=allowed_memory_types, allow_self=allow_self_memory, allowed_peer_ids=allowed_peer_ids, + allowed_self_memory_types=allowed_self_memory_types, + allowed_peer_memory_types=allowed_peer_memory_types, ) isolation_handler.prepare_messages() @@ -630,6 +655,8 @@ async def _extract_user_memories( "allowed_memory_types": allowed_memory_types, "allow_self": allow_self_memory, "allowed_peer_ids": allowed_peer_ids, + "allowed_self_memory_types": allowed_self_memory_types, + "allowed_peer_memory_types": allowed_peer_memory_types, }, metadata={ "source_extraction_id": extraction_id, diff --git a/openviking/session/memory/memory_isolation_handler.py b/openviking/session/memory/memory_isolation_handler.py index f9f4fe1bbf..6b5cdf9d46 100644 --- a/openviking/session/memory/memory_isolation_handler.py +++ b/openviking/session/memory/memory_isolation_handler.py @@ -43,6 +43,8 @@ def __init__( allowed_memory_types: Optional[Set[str]] = None, allow_self: bool = True, allowed_peer_ids: Optional[Set[str]] = None, + allowed_self_memory_types: Optional[Set[str]] = None, + allowed_peer_memory_types: Optional[Set[str]] = None, ): self.ctx = ctx self._extract_context = extract_context @@ -51,6 +53,16 @@ def __init__( if allowed_memory_types is not None else None ) + self.allowed_self_memory_types = ( + {str(item) for item in allowed_self_memory_types} + if allowed_self_memory_types is not None + else self.allowed_memory_types + ) + self.allowed_peer_memory_types = ( + {str(item) for item in allowed_peer_memory_types} + if allowed_peer_memory_types is not None + else self.allowed_memory_types + ) peer_ids = { item for item in (safe_peer_id(item) for item in allowed_peer_ids or set()) @@ -123,11 +135,23 @@ def allows_schema(self, memory_type_schema: MemoryTypeSchema) -> bool: memory_type = getattr(memory_type_schema, "memory_type", "") if memory_type in _INTERNAL_MEMORY_TYPES: return True - if self.allowed_memory_types is not None and memory_type not in self.allowed_memory_types: - return False - if not self.allow_self and not getattr(memory_type_schema, "peer_enabled", True): - return False - return True + self_allowed = self.allow_self and self._allows_self_type(memory_type) + peer_allowed = ( + self.allow_peer + and getattr(memory_type_schema, "peer_enabled", True) + and self._allows_peer_type(memory_type) + ) + return self_allowed or peer_allowed + + def _allows_self_type(self, memory_type: str) -> bool: + return ( + self.allowed_self_memory_types is None or memory_type in self.allowed_self_memory_types + ) + + def _allows_peer_type(self, memory_type: str) -> bool: + return ( + self.allowed_peer_memory_types is None or memory_type in self.allowed_peer_memory_types + ) def _can_write_peer(self, peer_id: str) -> bool: return self.allow_peer and peer_id in self.allowed_peer_ids @@ -147,9 +171,14 @@ def render_schema_directories(self, memory_type_schema: MemoryTypeSchema) -> Lis user_id = self.ctx.user.user_id if self.ctx and self.ctx.user else "default" user_space = user_id user_spaces: List[str] = [] - if self.allow_self: + memory_type = getattr(memory_type_schema, "memory_type", "") + if self.allow_self and self._allows_self_type(memory_type): user_spaces.append(user_space) - if self.allow_peer and getattr(memory_type_schema, "peer_enabled", True): + if ( + self.allow_peer + and getattr(memory_type_schema, "peer_enabled", True) + and self._allows_peer_type(memory_type) + ): for peer_id in sorted(self.allowed_peer_ids): user_spaces.append(peer_user_space(user_space, peer_id)) @@ -207,21 +236,26 @@ def calculate_memory_uris( user_id = self.ctx.user.user_id operation.memory_fields["user_id"] = user_id + memory_type = getattr(memory_type_schema, "memory_type", "") target_ids: List[str] = [] has_ranges = operation.memory_fields.get("ranges") is not None if not getattr(memory_type_schema, "peer_enabled", True): operation.memory_fields.pop("peer_id", None) - target_ids = [_SELF_PEER_ID] if self.allow_self else [] + target_ids = ( + [_SELF_PEER_ID] if self.allow_self and self._allows_self_type(memory_type) else [] + ) elif operation.memory_fields.get("ranges") is not None: target_ids = self._range_targets( operation.memory_fields.get("ranges"), ) operation.memory_fields.pop("peer_id", None) else: - target_id = self._resolve_operation_target_id( - operation.memory_fields.get("peer_id"), - ) + raw_peer_id = operation.memory_fields.get("peer_id") + if raw_peer_id in (None, "") and not self._allows_self_type(memory_type): + target_id = self._unique_peer_target_id_in_messages() + else: + target_id = self._resolve_operation_target_id(raw_peer_id) if target_id: target_ids = [target_id] if target_id == _SELF_PEER_ID: @@ -234,6 +268,15 @@ def calculate_memory_uris( if not target_ids: return [] + target_ids = [ + target_id + for target_id in target_ids + if (target_id == _SELF_PEER_ID and self._allows_self_type(memory_type)) + or (target_id != _SELF_PEER_ID and self._allows_peer_type(memory_type)) + ] + if not target_ids: + return [] + # 文件 uris = set() user_space = user_id diff --git a/openviking/session/memory/streaming_memory_updater.py b/openviking/session/memory/streaming_memory_updater.py index 4453463f0e..ce8be07499 100644 --- a/openviking/session/memory/streaming_memory_updater.py +++ b/openviking/session/memory/streaming_memory_updater.py @@ -1798,6 +1798,8 @@ def _make_isolation_handler( allowed_memory_types=options.get("allowed_memory_types"), allow_self=options.get("allow_self", True), allowed_peer_ids=options.get("allowed_peer_ids"), + allowed_self_memory_types=options.get("allowed_self_memory_types"), + allowed_peer_memory_types=options.get("allowed_peer_memory_types"), ) diff --git a/openviking/session/memory_policy.py b/openviking/session/memory_policy.py index f5c2bf3e6f..61b0207b29 100644 --- a/openviking/session/memory_policy.py +++ b/openviking/session/memory_policy.py @@ -11,10 +11,14 @@ from openviking_cli.exceptions import InvalidArgumentError _POLICY_KEYS = {"self", "peer", "memory_types", "working_memory"} -_TARGET_KEYS = {"enabled"} +_TARGET_KEYS = {"enabled", "memory_types"} +_WORKING_MEMORY_KEYS = {"enabled"} _TRUE_STRINGS = {"1", "true", "yes", "on"} _FALSE_STRINGS = {"", "0", "false", "no", "off"} +AGENT_MEMORY_TYPES = {"cases", "trajectories", "experiences"} +_EXPERIENCE_DEPENDENCIES = AGENT_MEMORY_TYPES + def _memory_policy_shape_error(key: str) -> str: if key == "working_memory": @@ -25,7 +29,7 @@ def _memory_policy_shape_error(key: str) -> str: def _memory_policy_keys_error(key: str) -> str: if key == "working_memory": return "memory_policy.working_memory supports only: enabled" - return "memory_policy target supports only: enabled" + return "memory_policy target supports only: enabled, memory_types" def _parse_enabled(value: Any, *, key: str) -> bool: @@ -44,47 +48,68 @@ def _parse_enabled(value: Any, *, key: str) -> bool: if normalized in _FALSE_STRINGS: return False - # Preserve the permissive pre-v0.5 behavior for uncommon legacy values. - # This avoids breaking persisted policy dictionaries while callers migrate - # to genuine JSON booleans. return bool(value) -def _target_enabled(data: Any, *, default_enabled: bool, key: str = "target") -> bool: - if data is None: - return default_enabled - if not isinstance(data, dict): - raise InvalidArgumentError(_memory_policy_shape_error(key)) - extra_keys = set(data) - _TARGET_KEYS - if extra_keys: - raise InvalidArgumentError(_memory_policy_keys_error(key)) - if "enabled" not in data: - return default_enabled - return _parse_enabled(data["enabled"], key=key) - - -def _parse_memory_types(data: Any) -> Optional[set[str]]: +def _parse_memory_types(data: Any, *, field: str) -> Optional[set[str]]: if data is None: return None if not isinstance(data, list): - raise InvalidArgumentError("memory_policy.memory_types must be a list") + raise InvalidArgumentError(f"{field} must be a list") memory_types = set() for item in data: if not isinstance(item, str) or not item: - raise InvalidArgumentError("memory_policy.memory_types must contain non-empty strings") + raise InvalidArgumentError(f"{field} must contain non-empty strings") memory_types.add(item) return memory_types +def _parse_target( + data: Any, + *, + default_enabled: bool, + key: str, +) -> tuple[bool, Optional[set[str]], bool]: + if data is None: + return default_enabled, None, False + if not isinstance(data, dict): + raise InvalidArgumentError(_memory_policy_shape_error(key)) + allowed_keys = _WORKING_MEMORY_KEYS if key == "working_memory" else _TARGET_KEYS + extra_keys = set(data) - allowed_keys + if extra_keys: + raise InvalidArgumentError(_memory_policy_keys_error(key)) + enabled = _parse_enabled(data["enabled"], key=key) if "enabled" in data else default_enabled + if key == "working_memory": + return enabled, None, False + has_memory_types = "memory_types" in data + memory_types = _parse_memory_types( + data.get("memory_types"), + field=f"memory_policy.{key}.memory_types", + ) + return enabled, memory_types, has_memory_types + + @dataclass class MemoryPolicy: - """Effective memory policy for one commit.""" + """Effective memory policy for one commit. + + ``self`` and ``peer`` have independent type filters. The legacy top-level + ``memory_types`` input is still accepted and normalized into this shape. + """ self_enabled: bool = True peer_enabled: bool = True - memory_types: Optional[set[str]] = None + self_memory_types: Optional[set[str]] = None + peer_memory_types: Optional[set[str]] = None working_memory_enabled: bool = True + @property + def memory_types(self) -> Optional[set[str]]: + """Compatibility view used by callers that only need the union.""" + if self.self_memory_types is None or self.peer_memory_types is None: + return None + return set(self.self_memory_types) | set(self.peer_memory_types) + @classmethod def default(cls) -> "MemoryPolicy": return cls() @@ -102,31 +127,91 @@ def from_dict(cls, data: Any) -> "MemoryPolicy": raise InvalidArgumentError( "memory_policy supports only: " + ", ".join(sorted(_POLICY_KEYS)) ) + + self_enabled, self_types, self_types_explicit = _parse_target( + data.get("self"), default_enabled=True, key="self" + ) + peer_enabled, peer_types, peer_types_explicit = _parse_target( + data.get("peer"), default_enabled=True, key="peer" + ) + working_enabled, _, _ = _parse_target( + data.get("working_memory"), default_enabled=True, key="working_memory" + ) + + legacy_types = _parse_memory_types( + data.get("memory_types"), field="memory_policy.memory_types" + ) + if not self_types_explicit: + self_types = None if legacy_types is None else set(legacy_types) + if not peer_types_explicit: + peer_types = None if legacy_types is None else set(legacy_types) - AGENT_MEMORY_TYPES + + if peer_types and peer_types & AGENT_MEMORY_TYPES: + raise InvalidArgumentError( + "memory_policy.peer.memory_types does not support agent memory types: " + + ", ".join(sorted(peer_types & AGENT_MEMORY_TYPES)) + ) + return cls( - self_enabled=_target_enabled(data.get("self"), default_enabled=True, key="self"), - peer_enabled=_target_enabled(data.get("peer"), default_enabled=True, key="peer"), - memory_types=_parse_memory_types(data.get("memory_types")), - working_memory_enabled=_target_enabled( - data.get("working_memory"), default_enabled=True, key="working_memory" - ), + self_enabled=self_enabled, + peer_enabled=peer_enabled, + self_memory_types=self_types, + peer_memory_types=peer_types, + working_memory_enabled=working_enabled, ) def validate_memory_types(self, known_memory_types: set[str]) -> None: - if self.memory_types is None: - return - unknown = self.memory_types - known_memory_types + configured = set() + if self.self_memory_types is not None: + configured.update(self.self_memory_types) + if self.peer_memory_types is not None: + configured.update(self.peer_memory_types) + unknown = configured - known_memory_types if unknown: raise InvalidArgumentError( - "Unknown memory_policy.memory_types: " + ", ".join(sorted(unknown)) + "Unknown memory_policy memory types: " + ", ".join(sorted(unknown)) ) + def resolve( + self, + known_memory_types: set[str], + *, + agent_evolution_enabled: bool, + ) -> "MemoryPolicy": + """Expand defaults and apply the account-scoped Agent switch.""" + self_types = ( + set(known_memory_types) + if self.self_memory_types is None + else set(self.self_memory_types) + ) + if "experiences" in self_types: + self_types.update(_EXPERIENCE_DEPENDENCIES) + peer_types = ( + set(known_memory_types) - AGENT_MEMORY_TYPES + if self.peer_memory_types is None + else set(self.peer_memory_types) + ) + if not agent_evolution_enabled: + self_types -= AGENT_MEMORY_TYPES + return MemoryPolicy( + self_enabled=self.self_enabled, + peer_enabled=self.peer_enabled, + self_memory_types=self_types, + peer_memory_types=peer_types, + working_memory_enabled=self.working_memory_enabled, + ) + def to_dict(self) -> dict[str, Any]: + self_target: dict[str, Any] = {"enabled": self.self_enabled} + peer_target: dict[str, Any] = {"enabled": self.peer_enabled} + if self.self_memory_types is not None: + self_target["memory_types"] = sorted(self.self_memory_types) + if self.peer_memory_types is not None: + peer_target["memory_types"] = sorted(self.peer_memory_types) data: dict[str, Any] = { - "self": {"enabled": self.self_enabled}, - "peer": {"enabled": self.peer_enabled}, + "self": self_target, + "peer": peer_target, } if not self.working_memory_enabled: data["working_memory"] = {"enabled": False} - if self.memory_types is not None: - data["memory_types"] = sorted(self.memory_types) return data diff --git a/openviking/session/session.py b/openviking/session/session.py index cd92be3737..55deb8b1c5 100644 --- a/openviking/session/session.py +++ b/openviking/session/session.py @@ -107,8 +107,6 @@ def _enabled_memory_types() -> set[str]: def _validate_memory_policy_types(policy: MemoryPolicy) -> None: - if policy.memory_types is None: - return policy.validate_memory_types(_enabled_memory_types()) @@ -117,24 +115,29 @@ def _apply_agent_evolution_setting( *, agent_evolution_enabled: bool, ) -> MemoryPolicy: - if agent_evolution_enabled: - return policy - effective_types = ( - _enabled_memory_types() if policy.memory_types is None else set(policy.memory_types) - ) - effective_types -= AGENT_EVOLUTION_MEMORY_TYPES - return MemoryPolicy( - self_enabled=policy.self_enabled, - peer_enabled=policy.peer_enabled, - memory_types=effective_types, - working_memory_enabled=policy.working_memory_enabled, + return policy.resolve( + _enabled_memory_types(), + agent_evolution_enabled=agent_evolution_enabled, ) def _effective_memory_types(policy: MemoryPolicy) -> set[str]: - if policy.memory_types is None: + enabled_types = _enabled_memory_types() + self_types = ( + enabled_types if policy.self_memory_types is None else set(policy.self_memory_types) + ) + peer_types = ( + enabled_types - AGENT_EVOLUTION_MEMORY_TYPES + if policy.peer_memory_types is None + else set(policy.peer_memory_types) + ) + return self_types | peer_types + + +def _effective_self_memory_types(policy: MemoryPolicy) -> set[str]: + if policy.self_memory_types is None: return _enabled_memory_types() - return set(policy.memory_types) + return set(policy.self_memory_types) def _agent_memory_skip_reason( @@ -196,6 +199,8 @@ class _MemoryExtractionScope: allow_self_memory: bool allowed_peer_ids: set[str] include_session_skills: bool + self_memory_types: set[str] + peer_memory_types: set[str] memory_types: Optional[set[str]] @@ -208,12 +213,35 @@ def _resolve_memory_extraction_scope( ) -> _MemoryExtractionScope: allow_self_memory = policy.self_enabled allowed_peer_ids = _message_peer_ids(messages) if policy.peer_enabled else set() + configured_scope_types: list[Optional[set[str]]] = [] + if allow_self_memory: + configured_scope_types.append(policy.self_memory_types) + if allowed_peer_ids: + configured_scope_types.append(policy.peer_memory_types) + if not configured_scope_types: + combined_memory_types: Optional[set[str]] = set() + elif any(memory_types is None for memory_types in configured_scope_types): + combined_memory_types = None + else: + combined_memory_types = set().union( + *(memory_types or set() for memory_types in configured_scope_types) + ) return _MemoryExtractionScope( allow_self_memory=allow_self_memory, allowed_peer_ids=allowed_peer_ids, include_session_skills=config_session_skill_extraction_enabled and allow_self_memory, - memory_types=policy.memory_types, + self_memory_types=(_effective_self_memory_types(policy) if allow_self_memory else set()), + peer_memory_types=( + ( + _enabled_memory_types() - AGENT_EVOLUTION_MEMORY_TYPES + if policy.peer_memory_types is None + else set(policy.peer_memory_types) + ) + if allowed_peer_ids + else set() + ), + memory_types=combined_memory_types, ) @@ -605,6 +633,7 @@ def __init__( agent_evolution_enabled: bool = True, usage_reporter: Optional["UsageReporter"] = None, agent_evolution_enabled_provider: Optional[Callable[[], bool | Awaitable[bool]]] = None, + memory_policy_provider: Optional[Callable[[], Any | Awaitable[Any]]] = None, ): self._viking_fs = viking_fs self._vikingdb_manager = vikingdb_manager @@ -637,10 +666,17 @@ def __init__( ) self._agent_evolution_enabled = agent_evolution_enabled self._agent_evolution_enabled_provider = agent_evolution_enabled_provider + self._memory_policy_provider = memory_policy_provider self._usage_reporter = usage_reporter logger.info(f"Session created: {self.session_id} for user {self.user}") + async def _provided_memory_policy(self) -> Any: + if self._memory_policy_provider is None: + return None + provided = self._memory_policy_provider() + return await provided if inspect.isawaitable(provided) else provided + async def load(self): """Load session data from storage.""" if self._loaded: @@ -1883,9 +1919,12 @@ async def commit_async( if turn_mode and effective_token_budget <= 0: raise ValueError("retained_message_token_budget must be greater than 0") in_memory_default_memory_policy = self._meta.memory_policy - effective_policy = MemoryPolicy.from_dict( - memory_policy if memory_policy is not None else self._meta.memory_policy - ) + initial_policy = memory_policy + if initial_policy is None: + initial_policy = self._meta.memory_policy + if initial_policy is None: + initial_policy = await self._provided_memory_policy() + effective_policy = MemoryPolicy.from_dict(initial_policy) _validate_memory_policy_types(effective_policy) agent_evolution_enabled = self._agent_evolution_enabled if self._agent_evolution_enabled_provider is not None: @@ -1900,10 +1939,9 @@ async def commit_async( agent_evolution_enabled=agent_evolution_enabled, ) effective_memory_policy = effective_policy.to_dict() - effective_memory_types = sorted(_effective_memory_types(effective_policy)) agent_memory_skip_reason = _agent_memory_skip_reason( agent_evolution_enabled=agent_evolution_enabled, - effective_memory_types=set(effective_memory_types), + effective_memory_types=_effective_self_memory_types(effective_policy), ) logger.info( f"[TRACER] session_commit started, trace_id={trace_id}, " @@ -1961,17 +1999,19 @@ async def commit_async( # messages being archived, unless this commit supplied an explicit # override. if memory_policy is None: - effective_policy = MemoryPolicy.from_dict(self._meta.memory_policy) + resolved_policy = self._meta.memory_policy + if resolved_policy is None: + resolved_policy = await self._provided_memory_policy() + effective_policy = MemoryPolicy.from_dict(resolved_policy) _validate_memory_policy_types(effective_policy) effective_policy = _apply_agent_evolution_setting( effective_policy, agent_evolution_enabled=agent_evolution_enabled, ) effective_memory_policy = effective_policy.to_dict() - effective_memory_types = sorted(_effective_memory_types(effective_policy)) agent_memory_skip_reason = _agent_memory_skip_reason( agent_evolution_enabled=agent_evolution_enabled, - effective_memory_types=set(effective_memory_types), + effective_memory_types=_effective_self_memory_types(effective_policy), ) archive_refs = await self._list_archive_refs() @@ -2561,17 +2601,25 @@ async def _run_recorded_memory_step( self_memory_enabled = extraction_scope.allow_self_memory allowed_peer_ids = extraction_scope.allowed_peer_ids session_skill_extraction_enabled = extraction_scope.include_session_skills + self_memory_type_filter = extraction_scope.self_memory_types + peer_memory_type_filter = extraction_scope.peer_memory_types memory_type_filter = extraction_scope.memory_types has_execution_memory = hasattr( self._session_compressor, "extract_execution_memories" ) if has_execution_memory: - long_term_memory_types, execution_memory_types = _split_policy_memory_types( - memory_type_filter + self_long_term_memory_types, execution_memory_types = ( + _split_policy_memory_types(self_memory_type_filter) ) + peer_long_term_memory_types, _ = _split_policy_memory_types( + peer_memory_type_filter + ) + long_term_memory_types, _ = _split_policy_memory_types(memory_type_filter) else: - long_term_memory_types = memory_type_filter + self_long_term_memory_types = self_memory_type_filter + peer_long_term_memory_types = peer_memory_type_filter execution_memory_types = set() + long_term_memory_types = memory_type_filter long_term_messages = [ message @@ -2630,6 +2678,8 @@ async def _run_long_term_memory_extraction() -> Any: latest_archive_overview=latest_archive_overview, archive_uri=archive_uri, allowed_memory_types=long_term_memory_types, + allowed_self_memory_types=self_long_term_memory_types, + allowed_peer_memory_types=peer_long_term_memory_types, agent_evolution_enabled=agent_evolution_enabled, allow_self_memory=self_memory_enabled, allowed_peer_ids=allowed_peer_ids, diff --git a/tests/server/test_admin_api.py b/tests/server/test_admin_api.py index 47221d9b0b..1a6fea6b22 100644 --- a/tests/server/test_admin_api.py +++ b/tests/server/test_admin_api.py @@ -96,6 +96,11 @@ async def rm(self, uri, **_kwargs): class _FakeService: def __init__(self): self.viking_fs = _FakeVikingFS() + self.sessions = self + + async def get_agent_evolution_enabled(self, account_id): + del account_id + return True async def initialize_account_directories(self, ctx): return None @@ -318,6 +323,72 @@ async def test_create_user_paths_accept_initial_user_config( assert bob_settings.resource_uri == "viking://user/resources/bob" +async def test_user_memory_policy_can_be_initialized_and_hot_updated( + lightweight_admin_client: httpx.AsyncClient, + lightweight_admin_app: FastAPI, +): + acct = _uid() + create_account = await lightweight_admin_client.post( + "/api/v1/admin/accounts", + json={"account_id": acct, "admin_user_id": "alice"}, + headers=root_headers(), + ) + assert create_account.status_code == 200, create_account.text + + create_user = await lightweight_admin_client.post( + f"/api/v1/admin/accounts/{acct}/users", + json={ + "user_id": "bob", + "role": "user", + "user_config": { + "memory_policy": { + "self": {"enabled": True, "memory_types": ["profile"]}, + "peer": {"enabled": False, "memory_types": []}, + } + }, + }, + headers=root_headers(), + ) + assert create_user.status_code == 200, create_user.text + + get_settings = await lightweight_admin_client.get( + f"/api/v1/admin/accounts/{acct}/users/bob/settings", + headers=root_headers(), + ) + assert get_settings.status_code == 200, get_settings.text + assert get_settings.json()["result"]["overrides"]["memory_policy"]["self"] == { + "enabled": True, + "memory_types": ["profile"], + } + + patch_settings = await lightweight_admin_client.patch( + f"/api/v1/admin/accounts/{acct}/users/bob/settings", + json={ + "memory_policy": { + "self": {"enabled": True, "memory_types": ["experiences"]}, + "peer": {"enabled": True, "memory_types": ["events"]}, + } + }, + headers=root_headers(), + ) + assert patch_settings.status_code == 200, patch_settings.text + configured = patch_settings.json()["result"]["overrides"]["memory_policy"] + assert configured["self"]["memory_types"] == ["experiences"] + assert configured["peer"]["memory_types"] == ["events"] + assert patch_settings.json()["result"]["settings"]["memory_policy"]["self"]["memory_types"] == [ + "cases", + "experiences", + "trajectories", + ] + + viking_fs = lightweight_admin_app.state.fake_service.viking_fs + persisted = await read_user_config( + viking_fs, + RequestContext(user=UserIdentifier(acct, "bob"), role=Role.USER), + ) + assert persisted.memory_policy == configured + + async def test_create_user_paths_ignore_deprecated_agent_evolution_config( lightweight_admin_client: httpx.AsyncClient, lightweight_admin_app: FastAPI, diff --git a/tests/session/test_memory_extraction_scope.py b/tests/session/test_memory_extraction_scope.py index ded2bcba99..7133c03c81 100644 --- a/tests/session/test_memory_extraction_scope.py +++ b/tests/session/test_memory_extraction_scope.py @@ -67,3 +67,50 @@ def test_actor_memory_extraction_scope_still_uses_policy_and_messages(): assert scope.allow_self_memory is True assert scope.allowed_peer_ids == {"bob"} assert scope.include_session_skills is True + + +def test_memory_extraction_scope_keeps_self_and_peer_type_filters_separate(): + scope = _resolve_memory_extraction_scope( + _ctx(), + MemoryPolicy.from_dict( + { + "self": {"memory_types": ["profile", "experiences"]}, + "peer": {"memory_types": ["events"]}, + } + ).resolve( + {"cases", "events", "experiences", "profile", "trajectories"}, + agent_evolution_enabled=True, + ), + [_message("alice")], + config_session_skill_extraction_enabled=True, + ) + + assert scope.self_memory_types == { + "cases", + "experiences", + "profile", + "trajectories", + } + assert scope.peer_memory_types == {"events"} + + +def test_disabled_peer_types_do_not_expand_the_active_extraction_scope(): + scope = _resolve_memory_extraction_scope( + _ctx(), + MemoryPolicy.from_dict( + { + "self": {"memory_types": ["experiences"]}, + "peer": { + "enabled": False, + "memory_types": ["profile", "events"], + }, + } + ).resolve( + {"cases", "events", "experiences", "profile", "trajectories"}, + agent_evolution_enabled=True, + ), + [_message("alice")], + config_session_skill_extraction_enabled=True, + ) + + assert scope.memory_types == {"cases", "experiences", "trajectories"} diff --git a/tests/session/test_session_commit.py b/tests/session/test_session_commit.py index db5949a4b7..a4a9be9af3 100644 --- a/tests/session/test_session_commit.py +++ b/tests/session/test_session_commit.py @@ -127,6 +127,28 @@ async def test_commit_uses_global_setting_and_enables_agent_memory( assert call_kwargs["agent_evolution_enabled"] is True assert call_kwargs["allowed_memory_types"] is None + async def test_commit_reads_latest_user_memory_policy_when_session_has_no_override( + self, session_with_messages: Session + ): + session_with_messages._memory_policy_provider = lambda: { + "self": {"enabled": True, "memory_types": ["profile"]}, + "peer": {"enabled": False, "memory_types": []}, + } + session_with_messages._session_compressor.extract_long_term_memories = AsyncMock( + return_value=[] + ) + + result = await session_with_messages.commit_async() + task_result = await _wait_for_task(result["task_id"]) + + assert task_result["status"] == "completed" + call_kwargs = ( + session_with_messages._session_compressor.extract_long_term_memories.call_args.kwargs + ) + assert call_kwargs["allowed_self_memory_types"] == {"profile"} + assert call_kwargs["allowed_peer_memory_types"] == set() + assert call_kwargs["allowed_peer_ids"] == set() + async def test_disabled_agent_evolution_keeps_working_memory( self, session_with_messages: Session, monkeypatch ): diff --git a/tests/unit/session/test_memory_policy.py b/tests/unit/session/test_memory_policy.py index f1116b4502..c8d9830a0b 100644 --- a/tests/unit/session/test_memory_policy.py +++ b/tests/unit/session/test_memory_policy.py @@ -43,9 +43,8 @@ def test_memory_policy_uses_top_level_memory_types(): assert policy.peer_enabled is True assert policy.memory_types == {"profile", "events"} assert policy.to_dict() == { - "self": {"enabled": False}, - "peer": {"enabled": True}, - "memory_types": ["events", "profile"], + "self": {"enabled": False, "memory_types": ["events", "profile"]}, + "peer": {"enabled": True, "memory_types": ["events", "profile"]}, } @@ -95,7 +94,29 @@ def test_memory_policy_rejects_invalid_memory_types(): with pytest.raises(InvalidArgumentError, match="missing"): policy.validate_memory_types({"profile"}) - assert MemoryPolicy.from_dict({"memory_types": ["experiences"]}).memory_types == {"experiences"} + policy = MemoryPolicy.from_dict({"memory_types": ["experiences"]}) + assert policy.self_memory_types == {"experiences"} + assert policy.resolve( + {"cases", "events", "experiences", "profile", "trajectories"}, + agent_evolution_enabled=True, + ).self_memory_types == {"cases", "experiences", "trajectories"} + + +def test_memory_policy_supports_independent_self_and_peer_types(): + policy = MemoryPolicy.from_dict( + { + "self": {"enabled": True, "memory_types": ["profile", "experiences"]}, + "peer": {"enabled": True, "memory_types": ["events"]}, + } + ) + + assert policy.self_memory_types == {"experiences", "profile"} + assert policy.peer_memory_types == {"events"} + + +def test_memory_policy_rejects_agent_types_for_peer_memory(): + with pytest.raises(InvalidArgumentError, match="does not support agent memory types"): + MemoryPolicy.from_dict({"peer": {"enabled": True, "memory_types": ["experiences"]}}) async def test_initialize_memory_files_respects_memory_type_filter(monkeypatch):