diff --git a/openviking/server/routers/sessions.py b/openviking/server/routers/sessions.py index ef6d5748c7..0497ca0511 100644 --- a/openviking/server/routers/sessions.py +++ b/openviking/server/routers/sessions.py @@ -858,7 +858,7 @@ async def record_used( resolve_path_variables(resolved_skill["uri"]), _ctx ) - session.used(contexts=resolved_contexts, skill=resolved_skill) + await session.used_async(contexts=resolved_contexts, skill=resolved_skill) return Response( status="ok", result={ diff --git a/openviking/session/session.py b/openviking/session/session.py index e1c7a3de21..ca7d89f2e2 100644 --- a/openviking/session/session.py +++ b/openviking/session/session.py @@ -500,8 +500,10 @@ class SessionMeta: # session. Maps to config.memory_extraction_config.events.tags in the API. # None means no session default; a commit may still override per-call. event_search_tags: Optional[List[str]] = None + # Usage waiting to be consumed by the next successful commit. + pending_usage_records: List[Dict[str, Any]] = field(default_factory=list) - def to_dict(self) -> Dict[str, Any]: + def to_dict(self, *, include_internal: bool = False) -> Dict[str, Any]: data = { "session_id": self.session_id, "created_at": self.created_at, @@ -531,6 +533,8 @@ def to_dict(self) -> Dict[str, Any]: data["total_message_count"] = self.total_message_count if self.event_search_tags is not None: data["event_search_tags"] = list(self.event_search_tags) + if include_internal and self.pending_usage_records: + data["pending_usage_records"] = [dict(item) for item in self.pending_usage_records] return data @classmethod @@ -581,6 +585,11 @@ def from_dict(cls, data: Dict[str, Any]) -> "SessionMeta": last_message_at=data.get("last_message_at", ""), last_auto_commit_at=data.get("last_auto_commit_at", ""), event_search_tags=data.get("event_search_tags"), + pending_usage_records=[ + dict(item) + for item in (data.get("pending_usage_records") or []) + if isinstance(item, dict) + ], ) @@ -591,11 +600,34 @@ class Usage: uri: str type: str # "context" | "skill" contribution: float = 0.0 - input: str = "" - output: str = "" + input: Any = "" + output: Any = "" success: bool = True timestamp: str = field(default_factory=get_current_timestamp) + def to_dict(self) -> Dict[str, Any]: + return { + "uri": self.uri, + "type": self.type, + "contribution": self.contribution, + "input": self.input, + "output": self.output, + "success": self.success, + "timestamp": self.timestamp, + } + + @classmethod + def from_dict(cls, data: Dict[str, Any]) -> "Usage": + return cls( + uri=str(data.get("uri", "")), + type=str(data.get("type", "")), + contribution=float(data.get("contribution", 0.0) or 0.0), + input=data.get("input", ""), + output=data.get("output", ""), + success=bool(data.get("success", True)), + timestamp=str(data.get("timestamp", "")) or get_current_timestamp(), + ) + class Session: """Session management class - Message = role + parts.""" @@ -710,6 +742,7 @@ async def load(self): # message_count mirrors the live message list, maintained by every write # path. Recompute on load so a stale persisted value can't drift. self._meta.message_count = len(self._messages) + self._restore_pending_usage() if not self._meta.created_by_account_id: self._meta.created_by_account_id = self.ctx.account_id @@ -812,7 +845,7 @@ async def _save_meta(self, lease_ref: Optional[Any] = None) -> None: self._meta.updated_at = get_current_timestamp() await self._viking_fs.write_file( uri=f"{self._session_uri}/.meta.json", - content=json.dumps(self._meta.to_dict(), ensure_ascii=False), + content=json.dumps(self._meta.to_dict(include_internal=True), ensure_ascii=False), ctx=self.ctx, lease_ref=lease_ref, ) @@ -875,11 +908,54 @@ def used( skill: Optional[Dict[str, Any]] = None, ) -> None: """Record actually used contexts and skills.""" + run_async(self.used_async(contexts=contexts, skill=skill)) + + async def used_async( + self, + contexts: Optional[List[str]] = None, + skill: Optional[Dict[str, Any]] = None, + ) -> None: + """Persist actually used contexts and skills for the next commit.""" + records = [Usage(uri=uri, type="context") for uri in contexts or []] + if skill: + records.append( + Usage( + uri=skill.get("uri", ""), + type="skill", + input=skill.get("input", ""), + output=skill.get("output", ""), + success=skill.get("success", True), + ) + ) + if not records: + return + + if self._viking_fs: + session_path = self._viking_fs._uri_to_path(self._session_uri, ctx=self.ctx) + lease = await self._viking_fs._async_agfs.pathlock_acquire_exact( + session_path, timeout_secs=_SESSION_PHASE1_LOCK_TIMEOUT_SECONDS + ) + try: + try: + meta_content = await self._viking_fs.read_file( + f"{self._session_uri}/.meta.json", ctx=self.ctx + ) + self._meta = SessionMeta.from_dict(json.loads(meta_content)) + except Exception as exc: + if not _is_storage_not_found(exc): + raise + self._restore_pending_usage() + self._usage_records.extend(records) + self._meta.pending_usage_records = [item.to_dict() for item in self._usage_records] + await self._save_meta() + finally: + await self._viking_fs._async_agfs.pathlock_release(lease) + else: + self._usage_records.extend(records) + + self._restore_usage_stats() if contexts: for uri in contexts: - usage = Usage(uri=uri, type="context") - self._usage_records.append(usage) - self._stats.contexts_used += 1 logger.debug(f"Tracked context usage: {uri}") try: from openviking.metrics.datasources.session import SessionLifecycleDataSource @@ -891,15 +967,6 @@ def used( pass if skill: - usage = Usage( - uri=skill.get("uri", ""), - type="skill", - input=skill.get("input", ""), - output=skill.get("output", ""), - success=skill.get("success", True), - ) - self._usage_records.append(usage) - self._stats.skills_used += 1 logger.debug(f"Tracked skill usage: {skill.get('uri')}") try: from openviking.metrics.datasources.session import SessionLifecycleDataSource @@ -908,6 +975,14 @@ def used( except Exception: pass + def _restore_pending_usage(self) -> None: + self._usage_records = [Usage.from_dict(item) for item in self._meta.pending_usage_records] + self._restore_usage_stats() + + def _restore_usage_stats(self) -> None: + self._stats.contexts_used = sum(item.type == "context" for item in self._usage_records) + self._stats.skills_used = sum(item.type == "skill" for item in self._usage_records) + def _tool_result_store(self) -> Optional[ToolResultStore]: if not self._viking_fs: return None @@ -1920,6 +1995,7 @@ async def commit_async( ctx=self.ctx, ) self._meta = SessionMeta.from_dict(json.loads(meta_content)) + self._restore_pending_usage() if ( memory_policy is None and self._meta.memory_policy is None @@ -2148,7 +2224,10 @@ async def commit_async( # commit boundary, so an idle scan and a concurrent worker # never see a stale state. self._meta.last_auto_commit_at = get_current_timestamp() + self._meta.pending_usage_records = [] await self._save_meta() + self._usage_records = [] + self._restore_usage_stats() await self._write_phase1_ready_marker(archive_uri) except Exception as e: logger.error(f"[commit] Failed during {phase1_stage}: {e}") diff --git a/tests/api_test/sessions/slow/test_session_concurrency.py b/tests/api_test/sessions/slow/test_session_concurrency.py index 62138ac46f..f3fe54ff80 100644 --- a/tests/api_test/sessions/slow/test_session_concurrency.py +++ b/tests/api_test/sessions/slow/test_session_concurrency.py @@ -144,18 +144,23 @@ def test_session_used_multiple_times_accumulates(self, api_client): api_client.add_message(session_id, "user", "Used accumulation test") - api_client.session_used( + used1 = api_client.session_used( session_id, contexts=["viking://resources/ctx1"], skill={"name": "skill-a"}, ) - api_client.session_used( + used2 = api_client.session_used( session_id, contexts=["viking://resources/ctx2", "viking://resources/ctx3"], skill={"name": "skill-b"}, ) + assert used1.json()["result"]["contexts_used"] == 1 + assert used1.json()["result"]["skills_used"] == 1 + assert used2.json()["result"]["contexts_used"] == 3 + assert used2.json()["result"]["skills_used"] == 2 + get_resp = api_client.get_session(session_id) assert get_resp.status_code == 200 result = get_resp.json().get("result", {}) diff --git a/tests/session/test_session_commit.py b/tests/session/test_session_commit.py index 640633d34b..cfa5b6fdd6 100644 --- a/tests/session/test_session_commit.py +++ b/tests/session/test_session_commit.py @@ -705,11 +705,11 @@ async def test_active_count_incremented_after_commit(self, client_with_resource_ session_id="active_count_regression_test", ) await session.ensure_exists() - session._session_compressor.extract_long_term_memories = AsyncMock(return_value=[]) + service.sessions._session_compressor.extract_long_term_memories = AsyncMock(return_value=[]) session.add_message("user", [TextPart("Query")]) - session.used(contexts=[uri]) + await session.used_async(contexts=[uri]) session.add_message("assistant", [TextPart("Answer")]) - result = await session.commit_async() + result = await service.sessions.commit_async(session.session_id, client_ctx) # Wait for background task to complete (active_count is updated there) task_result = await _wait_for_task(result["task_id"]) diff --git a/tests/unit/session/test_event_tag_concurrency.py b/tests/unit/session/test_event_tag_concurrency.py index 0895284e27..d5952db301 100644 --- a/tests/unit/session/test_event_tag_concurrency.py +++ b/tests/unit/session/test_event_tag_concurrency.py @@ -32,7 +32,7 @@ def __init__( persisted_meta: SessionMeta, ): self.meta_uri = f"{session_uri}/.meta.json" - self.files = {self.meta_uri: json.dumps(persisted_meta.to_dict())} + self.files = {self.meta_uri: json.dumps(persisted_meta.to_dict(include_internal=True))} self._async_agfs = _PathLock() self.writes = [] @@ -46,6 +46,10 @@ async def read_file(self, uri, ctx=None): raise FileNotFoundError(uri) return self.files[uri] + async def exists(self, uri, ctx=None): + del uri, ctx + return False + async def write_file(self, uri, content, ctx=None, lease_ref=None): del ctx self.files[uri] = content @@ -54,10 +58,12 @@ async def write_file(self, uri, content, ctx=None, lease_ref=None): @pytest.mark.asyncio async def test_commit_uses_event_tags_from_lock_protected_meta_snapshot(monkeypatch): + monkeypatch.setattr("openviking.session.session._enabled_memory_types", lambda: set()) session_uri = "viking://user/default/sessions/session-1" persisted_meta = SessionMeta( session_id="session-1", event_search_tags=["channel=app"], + pending_usage_records=[{"uri": "viking://resources/context-1", "type": "context"}], ) viking_fs = _MetaVikingFS(session_uri, persisted_meta) session = Session( @@ -92,6 +98,7 @@ async def capture_phase1_marker(archive_uri, *, queue_message, **kwargs): await session.commit_async() assert captured_queue_message["event_search_tags"] == ["channel=app"] + assert captured_queue_message["usage_uris"] == ["viking://resources/context-1"] assert viking_fs._async_agfs.acquired == 1 assert viking_fs._async_agfs.released == 1 @@ -136,3 +143,35 @@ async def test_update_config_updates_policy_and_tags_in_one_locked_write(): assert viking_fs.writes[0][2] is None assert viking_fs._async_agfs.acquired == 1 assert viking_fs._async_agfs.released == 1 + + +@pytest.mark.asyncio +async def test_used_accumulates_across_recreated_session_instances(): + session_uri = "viking://user/default/sessions/session-1" + viking_fs = _MetaVikingFS(session_uri, SessionMeta(session_id="session-1")) + + first_request = Session( + viking_fs=viking_fs, + session_id="session-1", + session_uri=session_uri, + ) + await first_request.used_async( + contexts=["viking://resources/context-1"], + skill={"uri": "viking://skills/skill-1"}, + ) + + second_request = Session( + viking_fs=viking_fs, + session_id="session-1", + session_uri=session_uri, + ) + await second_request.used_async( + contexts=["viking://resources/context-2", "viking://resources/context-3"], + skill={"uri": "viking://skills/skill-2"}, + ) + + assert second_request.stats.contexts_used == 3 + assert second_request.stats.skills_used == 2 + saved_meta = SessionMeta.from_dict(json.loads(viking_fs.files[viking_fs.meta_uri])) + assert len(saved_meta.pending_usage_records) == 5 + assert "pending_usage_records" not in saved_meta.to_dict()