From 6690d2c860850d8b335e7c65392487dc89621f86 Mon Sep 17 00:00:00 2001 From: cgycorey <4724788+cgycorey@users.noreply.github.com> Date: Sat, 1 Aug 2026 12:57:54 +0100 Subject: [PATCH] fix(ui): restore custom org fallback for role tasks --- opc/plugins/office_ui/services/session.py | 19 +++ tests/test_office_session_org_fallback.py | 167 ++++++++++++++++++++++ 2 files changed, 186 insertions(+) create mode 100644 tests/test_office_session_org_fallback.py diff --git a/opc/plugins/office_ui/services/session.py b/opc/plugins/office_ui/services/session.py index f980f19..c4229ad 100644 --- a/opc/plugins/office_ui/services/session.py +++ b/opc/plugins/office_ui/services/session.py @@ -618,6 +618,25 @@ class SessionService: default_preferred_agent=self.context.mode_state.task_preferred_agent, explicit_exec_mode=True, ) + if identity.is_custom_org and not identity.org_id: + # Role-task rows may lack org_id; mirror create()'s active-org fallback + fallback_org_id = "" + if self.context.get_active_saved_org_name is not None: + try: + fallback_org_id = await self.context.get_active_saved_org_name() + except Exception: + logger.opt(exception=True).debug( + "persist_session_config: failed to resolve active saved org" + ) + if fallback_org_id: + identity = canonicalize_execution_identity( + exec_mode=exec_mode, + company_profile=company_profile, + preferred_agent=preferred_agent, + org_id=fallback_org_id, + default_preferred_agent=self.context.mode_state.task_preferred_agent, + explicit_exec_mode=True, + ) if identity.is_custom_org and not identity.org_id: raise ServiceError("org_id_required", "org_id_required", { "task_id": str(getattr(task, "id", "") or ""), diff --git a/tests/test_office_session_org_fallback.py b/tests/test_office_session_org_fallback.py new file mode 100644 index 0000000..78dfa01 --- /dev/null +++ b/tests/test_office_session_org_fallback.py @@ -0,0 +1,167 @@ +from __future__ import annotations + +import unittest +from types import SimpleNamespace +from typing import Any + +from opc.plugins.office_ui.services.context import OfficeServiceContext +from opc.plugins.office_ui.services.models import ServiceError +from opc.plugins.office_ui.services.session import SessionService + + +class _Store: + def __init__(self) -> None: + self.saved: list[Any] = [] + + async def save_task(self, task: Any) -> None: + self.saved.append(task) + + +def _task(**overrides: Any) -> SimpleNamespace: + task = SimpleNamespace( + id="task-role-1", + session_id="sess-1", + project_id="demo", + title="Role task", + parent_session_id=None, + metadata={}, + org_id=None, + ) + for key, value in overrides.items(): + setattr(task, key, value) + return task + + +def _context(*, hook: Any | None = None) -> OfficeServiceContext: + engine = SimpleNamespace(project_id="demo", store=_Store(), memory=None) + context = OfficeServiceContext(engine=engine, agent_store=None, chat_store=None, event_adapter=None) + if hook is not None: + context.get_active_saved_org_name = hook + return context + + +class TestPersistSessionConfigOrgFallback(unittest.IsolatedAsyncioTestCase): + async def _persist( + self, + context: OfficeServiceContext, + task: Any, + *, + exec_mode: str = "org", + company_profile: str = "custom", + org_id: str = "", + ) -> None: + await SessionService(context).persist_session_config( + task, + exec_mode=exec_mode, + company_profile=company_profile, + preferred_agent="native", + org_id=org_id, + ) + + async def test_falls_back_to_active_saved_org_when_task_lacks_org_id(self) -> None: + async def active_org() -> str: + return "vc-investment-firm" + + context = _context(hook=active_org) + task = _task() + + await self._persist(context, task) + + assert task.metadata["org_id"] == "vc-investment-firm" + assert task.metadata["organization_id"] == "vc-investment-firm" + assert task.org_id == "vc-investment-firm" + assert task.metadata["exec_mode"] == "org" + assert task.metadata["company_profile"] == "custom" + assert context.engine.store.saved == [task] + + async def test_explicit_org_id_still_used_when_present(self) -> None: + async def active_org() -> str: + return "other-org" + + context = _context(hook=active_org) + task = _task() + + await self._persist(context, task, org_id="vc-investment-firm") + + assert task.metadata["org_id"] == "vc-investment-firm" + assert task.org_id == "vc-investment-firm" + + async def test_raises_when_no_active_org_available(self) -> None: + async def empty_org() -> str: + return "" + + context = _context(hook=empty_org) + task = _task() + + with self.assertRaises(ServiceError) as ctx: + await self._persist(context, task) + assert ctx.exception.code == "org_id_required" + + async def test_raises_when_hook_unset(self) -> None: + context = _context() + task = _task() + + with self.assertRaises(ServiceError) as ctx: + await self._persist(context, task) + assert ctx.exception.code == "org_id_required" + + async def test_raises_when_hook_fails(self) -> None: + async def broken() -> str: + raise RuntimeError("org index unreadable") + + context = _context(hook=broken) + task = _task() + + with self.assertRaises(ServiceError) as ctx: + await self._persist(context, task) + assert ctx.exception.code == "org_id_required" + + async def test_company_mode_clears_org_fields_without_fallback(self) -> None: + fallback_called = False + + async def active_org() -> str: + nonlocal fallback_called + fallback_called = True + return "vc-investment-firm" + + context = _context(hook=active_org) + task = _task(metadata={"org_id": "stale-org"}) + + await self._persist( + context, + task, + exec_mode="company", + company_profile="corporate", + ) + + assert not fallback_called + assert "org_id" not in task.metadata + assert "organization_id" not in task.metadata + assert task.org_id is None + + async def test_task_mode_ignores_org_entirely(self) -> None: + fallback_called = False + + async def active_org() -> str: + nonlocal fallback_called + fallback_called = True + return "vc-investment-firm" + + context = _context(hook=active_org) + task = _task() + + await self._persist( + context, + task, + exec_mode="task", + company_profile="corporate", + ) + + assert not fallback_called + assert task.metadata["execution_mode"] == "task_mode" + assert "org_id" not in task.metadata + assert task.org_id is None + + +if __name__ == "__main__": + unittest.main()