fix(ui): use durable org for role session replies
This commit is contained in:
@@ -4967,6 +4967,28 @@ class WSHandler:
|
|||||||
def _resolve_task_org_id(self, task: Any | None) -> str:
|
def _resolve_task_org_id(self, task: Any | None) -> str:
|
||||||
return self._ensure_office_services().session.resolve_task_org_id(task)
|
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
|
@staticmethod
|
||||||
def _is_company_session_exec_mode(exec_mode: Any) -> bool:
|
def _is_company_session_exec_mode(exec_mode: Any) -> bool:
|
||||||
return str(exec_mode or "").strip().lower() in {"company", "org", "custom"}
|
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_preferred_agent = self._task_preferred_agent
|
||||||
session_org_id = ""
|
session_org_id = ""
|
||||||
task = None
|
task = None
|
||||||
|
config_task = None
|
||||||
store = engine.store
|
store = engine.store
|
||||||
if self._store_is_ready(store):
|
if self._store_is_ready(store):
|
||||||
from opc.core.models import TaskStatus
|
from opc.core.models import TaskStatus
|
||||||
task = await store.get_task(task_id)
|
task = await store.get_task(task_id)
|
||||||
if task:
|
if task:
|
||||||
session_exec_mode, session_company_profile = self._resolve_task_session_config(task)
|
config_task = await self._resolve_session_runtime_config_task(
|
||||||
session_org_id = self._resolve_task_org_id(task)
|
task_id,
|
||||||
session_preferred_agent = self._resolve_task_preferred_agent(task)
|
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(
|
if await self._try_lock_free_parked_checkpoint_answer(
|
||||||
task_id=task_id,
|
task_id=task_id,
|
||||||
@@ -8729,9 +8759,16 @@ class WSHandler:
|
|||||||
from opc.core.models import TaskStatus
|
from opc.core.models import TaskStatus
|
||||||
task = await store.get_task(task_id)
|
task = await store.get_task(task_id)
|
||||||
if task:
|
if task:
|
||||||
session_exec_mode, session_company_profile = self._resolve_task_session_config(task)
|
config_task = await self._resolve_session_runtime_config_task(
|
||||||
session_org_id = self._resolve_task_org_id(task)
|
task_id,
|
||||||
session_preferred_agent = self._resolve_task_preferred_agent(task)
|
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):
|
if task.status == TaskStatus.DONE and self._is_company_session_exec_mode(session_exec_mode):
|
||||||
task.status = TaskStatus.IDLE
|
task.status = TaskStatus.IDLE
|
||||||
task.metadata = dict(getattr(task, "metadata", {}) or {})
|
task.metadata = dict(getattr(task, "metadata", {}) or {})
|
||||||
@@ -8757,12 +8794,17 @@ class WSHandler:
|
|||||||
except Exception:
|
except Exception:
|
||||||
logger.opt(exception=True).debug("failed to mark company session runtime running")
|
logger.opt(exception=True).debug("failed to mark company session runtime running")
|
||||||
try:
|
try:
|
||||||
engine_mode, company_profile = self._resolve_engine_mode(
|
config_task_id = str(getattr(config_task, "id", "") or "").strip()
|
||||||
session_exec_mode,
|
selected_task_id = str(getattr(task, "id", "") or "").strip()
|
||||||
session_company_profile,
|
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 should_persist_selected_config:
|
||||||
if task is not None:
|
|
||||||
await self._persist_session_config(
|
await self._persist_session_config(
|
||||||
task,
|
task,
|
||||||
exec_mode=session_exec_mode,
|
exec_mode=session_exec_mode,
|
||||||
@@ -8771,6 +8813,22 @@ class WSHandler:
|
|||||||
org_id=session_org_id,
|
org_id=session_org_id,
|
||||||
engine=engine,
|
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 = dict(message_metadata or {})
|
||||||
engine_message_metadata.update(_ui_message_identity_metadata(
|
engine_message_metadata.update(_ui_message_identity_metadata(
|
||||||
message_id=user_message_id,
|
message_id=user_message_id,
|
||||||
|
|||||||
@@ -1998,6 +1998,115 @@ class TestWSHandlerSessionSend(unittest.IsolatedAsyncioTestCase):
|
|||||||
self.assertIn("company_session_reopened_at", refreshed.metadata)
|
self.assertIn("company_session_reopened_at", refreshed.metadata)
|
||||||
self.engine.process_message.assert_called_once()
|
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:
|
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."""
|
"""Session replies should reuse the task session even for simple stub stores."""
|
||||||
ws = MagicMock()
|
ws = MagicMock()
|
||||||
|
|||||||
Reference in New Issue
Block a user