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:
|
||||
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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user