diff --git a/opc/plugins/office_ui/ws_handler.py b/opc/plugins/office_ui/ws_handler.py index b3faa54..72e3dc1 100644 --- a/opc/plugins/office_ui/ws_handler.py +++ b/opc/plugins/office_ui/ws_handler.py @@ -4967,6 +4967,28 @@ class WSHandler: def _resolve_task_org_id(self, task: Any | None) -> str: return self._ensure_office_services().session.resolve_task_org_id(task) + async def _resolve_session_runtime_config_task( + self, + task_id: str, + task: Any | None, + *, + engine: Any, + ) -> Any | None: + """Resolve company-session config from the durable runtime identity.""" + if task is None: + return None + exec_mode, _ = self._resolve_task_session_config(task) + if not self._is_company_session_exec_mode(exec_mode): + return task + try: + target = await self._resolve_company_runtime_target(task_id, engine=engine) + except Exception: + logger.opt(exception=True).debug( + "failed to resolve durable session config task" + ) + return task + return (target or {}).get("config_task") or task + @staticmethod def _is_company_session_exec_mode(exec_mode: Any) -> bool: return str(exec_mode or "").strip().lower() in {"company", "org", "custom"} @@ -8694,14 +8716,22 @@ class WSHandler: session_preferred_agent = self._task_preferred_agent session_org_id = "" task = None + config_task = None store = engine.store if self._store_is_ready(store): from opc.core.models import TaskStatus task = await store.get_task(task_id) if task: - session_exec_mode, session_company_profile = self._resolve_task_session_config(task) - session_org_id = self._resolve_task_org_id(task) - session_preferred_agent = self._resolve_task_preferred_agent(task) + config_task = await self._resolve_session_runtime_config_task( + task_id, + task, + engine=engine, + ) + session_exec_mode, session_company_profile = self._resolve_task_session_config( + config_task + ) + session_org_id = self._resolve_task_org_id(config_task) + session_preferred_agent = self._resolve_task_preferred_agent(config_task) if await self._try_lock_free_parked_checkpoint_answer( task_id=task_id, @@ -8729,9 +8759,16 @@ class WSHandler: from opc.core.models import TaskStatus task = await store.get_task(task_id) if task: - session_exec_mode, session_company_profile = self._resolve_task_session_config(task) - session_org_id = self._resolve_task_org_id(task) - session_preferred_agent = self._resolve_task_preferred_agent(task) + config_task = await self._resolve_session_runtime_config_task( + task_id, + task, + engine=engine, + ) + session_exec_mode, session_company_profile = self._resolve_task_session_config( + config_task + ) + session_org_id = self._resolve_task_org_id(config_task) + session_preferred_agent = self._resolve_task_preferred_agent(config_task) if task.status == TaskStatus.DONE and self._is_company_session_exec_mode(session_exec_mode): task.status = TaskStatus.IDLE task.metadata = dict(getattr(task, "metadata", {}) or {}) @@ -8757,12 +8794,17 @@ class WSHandler: except Exception: logger.opt(exception=True).debug("failed to mark company session runtime running") try: - engine_mode, company_profile = self._resolve_engine_mode( - session_exec_mode, - session_company_profile, + config_task_id = str(getattr(config_task, "id", "") or "").strip() + selected_task_id = str(getattr(task, "id", "") or "").strip() + should_persist_selected_config = bool( + task is not None + and ( + not config_task_id + or config_task_id == selected_task_id + or (session_exec_mode == "org" and not session_org_id) + ) ) - engine_preferred_agent = session_preferred_agent if session_exec_mode == "task" else None - if task is not None: + if should_persist_selected_config: await self._persist_session_config( task, exec_mode=session_exec_mode, @@ -8771,6 +8813,22 @@ class WSHandler: org_id=session_org_id, engine=engine, ) + persisted_identity = self._resolve_task_identity( + task, + default_exec_mode=session_exec_mode, + default_company_profile=session_company_profile, + default_preferred_agent=session_preferred_agent, + default_org_id=session_org_id, + ) + session_exec_mode = persisted_identity.exec_mode + session_company_profile = persisted_identity.company_profile + session_preferred_agent = persisted_identity.preferred_agent + session_org_id = persisted_identity.org_id + engine_mode, company_profile = self._resolve_engine_mode( + session_exec_mode, + session_company_profile, + ) + engine_preferred_agent = session_preferred_agent if session_exec_mode == "task" else None engine_message_metadata = dict(message_metadata or {}) engine_message_metadata.update(_ui_message_identity_metadata( message_id=user_message_id, diff --git a/tests/test_session_integration.py b/tests/test_session_integration.py index 4ace294..26b04e1 100644 --- a/tests/test_session_integration.py +++ b/tests/test_session_integration.py @@ -1998,6 +1998,115 @@ class TestWSHandlerSessionSend(unittest.IsolatedAsyncioTestCase): self.assertIn("company_session_reopened_at", refreshed.metadata) self.engine.process_message.assert_called_once() + async def test_process_session_message_uses_durable_org_for_role_task(self) -> None: + runtime_session_id = "runtime-org-session" + anchor = await self.store.get_task(self.task_id) + assert anchor is not None + anchor.session_id = runtime_session_id + anchor.metadata = { + "exec_mode": "org", + "mode": "company", + "company_profile": "custom", + "org_id": "vc-investment-firm", + } + await self.store.save_task(anchor) + + role_task = Task( + id="role-task-without-org-id", + title="Sector Analyst", + project_id="test-project", + session_id=f"{runtime_session_id}:role:sector-analyst", + parent_session_id=runtime_session_id, + metadata={ + "exec_mode": "org", + "mode": "company", + "company_profile": "custom", + "shared_role_session": True, + "shared_role_id": "sector_analyst", + "company_runtime_root_session_id": runtime_session_id, + }, + ) + await self.store.save_task(role_task) + self.handler.services_context.get_active_saved_org_name = AsyncMock( + return_value="wrong-active-org" + ) + + await self.handler._process_session_message( + role_task.id, + "approve", + session_id=role_task.session_id, + ) + + self.engine.process_message.assert_called_once() + call_kwargs = self.engine.process_message.call_args.kwargs + self.assertEqual(call_kwargs["mode"], "org") + self.assertEqual(call_kwargs["company_profile"], "custom") + self.assertEqual(call_kwargs["org_id"], "vc-investment-firm") + + async def test_lock_free_process_session_message_uses_durable_org_for_role_task(self) -> None: + runtime_session_id = "runtime-org-lock-free-session" + anchor = await self.store.get_task(self.task_id) + assert anchor is not None + anchor.session_id = runtime_session_id + anchor.metadata = { + "exec_mode": "org", + "mode": "company", + "company_profile": "custom", + "org_id": "vc-investment-firm", + } + await self.store.save_task(anchor) + + role_task = Task( + id="role-task-lock-free-without-org-id", + title="Sector Analyst", + project_id="test-project", + session_id=f"{runtime_session_id}:role:sector-analyst", + parent_session_id=runtime_session_id, + metadata={ + "exec_mode": "org", + "mode": "company", + "company_profile": "custom", + "shared_role_session": True, + "shared_role_id": "sector_analyst", + "company_runtime_root_session_id": runtime_session_id, + }, + ) + await self.store.save_task(role_task) + checkpoint = ExecutionCheckpoint( + checkpoint_id="org-lock-free-gate", + project_id="test-project", + session_id=runtime_session_id, + checkpoint_type="company_work_item_gate", + status="pending", + task_id=role_task.id, + ) + await self.store.save_execution_checkpoint(checkpoint) + self.engine._load_execution_checkpoint_by_id = AsyncMock(return_value=checkpoint) + self.handler.services_context.get_active_saved_org_name = AsyncMock( + return_value="wrong-active-org" + ) + + lock = self.handler._get_task_lock(role_task.id) + await lock.acquire() + try: + await self.handler._process_session_message( + role_task.id, + "approve", + session_id=role_task.session_id, + message_metadata={ + "response_to_checkpoint_id": checkpoint.checkpoint_id, + "response_to_checkpoint_type": checkpoint.checkpoint_type, + }, + ) + finally: + lock.release() + + self.engine.process_message.assert_called_once() + call_kwargs = self.engine.process_message.call_args.kwargs + self.assertEqual(call_kwargs["mode"], "org") + self.assertEqual(call_kwargs["company_profile"], "custom") + self.assertEqual(call_kwargs["org_id"], "vc-investment-firm") + async def test_session_send_reuses_task_session_without_is_ready_flag(self) -> None: """Session replies should reuse the task session even for simple stub stores.""" ws = MagicMock()