fix(ui): preserve durable org identity for runtime approvals

This commit is contained in:
cgycorey
2026-08-01 23:00:32 +01:00
parent 34770373c2
commit 734a2d5969
11 changed files with 719 additions and 67 deletions
+109 -20
View File
@@ -4978,16 +4978,38 @@ class WSHandler:
if task is None:
return None
exec_mode, _ = self._resolve_task_session_config(task)
if not self._is_company_session_exec_mode(exec_mode):
runtime_bound = is_company_runtime_task(task)
if not runtime_bound and 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:
except ServiceError:
if runtime_bound:
raise
logger.opt(exception=True).debug(
"failed to resolve durable session config task"
)
return task
return (target or {}).get("config_task") or task
except Exception as exc:
if runtime_bound:
raise ServiceError(
"company_runtime_identity_mismatch",
"Company runtime identity could not be resolved",
{"task_id": str(task_id or "").strip()},
) from exc
logger.opt(exception=True).debug(
"failed to resolve durable session config task"
)
return task
if target is not None:
return target.get("config_task") or task
if runtime_bound:
raise ServiceError(
"company_runtime_identity_mismatch",
"Company runtime identity could not be resolved",
{"task_id": str(task_id or "").strip()},
)
return task
@staticmethod
def _is_company_session_exec_mode(exec_mode: Any) -> bool:
@@ -5059,14 +5081,24 @@ class WSHandler:
# Look up session_id from task
session_id: str | None = None
task = None
config_task = None
preferred_agent = self._task_preferred_agent
session_org_id = self._normalize_session_org_id(org_id)
if task_id and getattr(engine, "store", None):
task = await engine.store.get_task(task_id)
if task:
session_id = task.session_id
identity = self._resolve_task_identity(
company_runtime_target: dict[str, Any] | None = None
try:
if task is not None:
config_task = await self._resolve_session_runtime_config_task(
task_id,
task,
engine=engine,
)
identity = self._resolve_task_identity(
config_task,
default_exec_mode=mode,
default_company_profile=profile,
default_preferred_agent=preferred_agent,
@@ -5076,9 +5108,12 @@ class WSHandler:
profile = identity.company_profile
session_org_id = identity.org_id
preferred_agent = identity.preferred_agent
company_runtime_target: dict[str, Any] | None = None
try:
if identity.is_custom_org and not identity.org_id:
raise ServiceError(
"org_id_required",
"org_id_required",
{"project_id": pid, "task_id": task_id},
)
content = f"{title}\n{description}".strip()
engine_mode, company_profile = self._resolve_engine_mode(mode, profile)
engine_preferred_agent = preferred_agent if engine_mode == "project" else None
@@ -8124,17 +8159,6 @@ class WSHandler:
pid = self._normalize_project_id(run_project_id or getattr(run_engine, "project_id", None))
async with lock:
try:
try:
await self._set_company_runtime_control(
target,
state="resuming",
checkpoint_id=str(
getattr(checkpoint, "checkpoint_id", "") or ""
).strip(),
)
except Exception:
logger.opt(exception=True).debug("failed to broadcast company suspend reply routing state")
config_task = target.get("config_task")
session_exec_mode = self._normalize_session_exec_mode(self._exec_mode)
session_company_profile = self._normalize_session_company_profile(self._company_profile)
@@ -8146,6 +8170,27 @@ class WSHandler:
session_exec_mode,
session_company_profile,
)
if engine_mode == "org" and not session_org_id:
raise ServiceError(
"org_id_required",
"org_id_required",
{
"project_id": pid,
"task_id": str(target.get("config_source_task_id", "") or "").strip(),
},
)
try:
await self._set_company_runtime_control(
target,
state="resuming",
checkpoint_id=str(
getattr(checkpoint, "checkpoint_id", "") or ""
).strip(),
)
except Exception:
logger.opt(exception=True).debug("failed to broadcast company suspend reply routing state")
engine_message_metadata = dict(message_metadata or {})
engine_message_metadata.update({
"response_to_checkpoint_id": str(getattr(checkpoint, "checkpoint_id", "") or ""),
@@ -8612,13 +8657,37 @@ class WSHandler:
item
for item in pending or []
if str(getattr(item, "checkpoint_id", "") or "").strip() == checkpoint_id
and str(getattr(item, "checkpoint_type", "") or "").strip()
in self._LOCK_FREE_CHECKPOINT_ANSWER_TYPES
and str(getattr(item, "checkpoint_type", "") or "").strip() == checkpoint_type
),
None,
)
if checkpoint is None:
return False
checkpoint_payload = dict(getattr(checkpoint, "payload", {}) or {})
checkpoint_task_id = str(
getattr(checkpoint, "task_id", "")
or checkpoint_payload.get("waiting_task_id")
or checkpoint_payload.get("task_id")
or ""
).strip()
checkpoint_session_id = str(
getattr(checkpoint, "session_id", "")
or checkpoint_payload.get("session_id")
or ""
).strip()
payload_task_ids = {
str(item).strip()
for item in list(checkpoint_payload.get("task_ids", []) or [])
if str(item).strip()
}
if (
not checkpoint_task_id
or checkpoint_task_id != str(task_id or "").strip()
or not checkpoint_session_id
or checkpoint_session_id != str(session_id or "").strip()
or (payload_task_ids and str(task_id or "").strip() not in payload_task_ids)
):
return False
logger.info(
f"Lock-free checkpoint answer: task lock for {task_id} is held by a "
f"live turn; delivering reply to pending checkpoint {checkpoint_id} "
@@ -8732,6 +8801,16 @@ class WSHandler:
)
session_org_id = self._resolve_task_org_id(config_task)
session_preferred_agent = self._resolve_task_preferred_agent(config_task)
if (
session_exec_mode in {"org", "custom"}
and not session_org_id
and is_company_runtime_task(task)
):
raise ServiceError(
"org_id_required",
"org_id_required",
{"project_id": pid, "task_id": task_id},
)
if await self._try_lock_free_parked_checkpoint_answer(
task_id=task_id,
@@ -8769,6 +8848,16 @@ class WSHandler:
)
session_org_id = self._resolve_task_org_id(config_task)
session_preferred_agent = self._resolve_task_preferred_agent(config_task)
if (
session_exec_mode in {"org", "custom"}
and not session_org_id
and is_company_runtime_task(task)
):
raise ServiceError(
"org_id_required",
"org_id_required",
{"project_id": pid, "task_id": task_id},
)
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 {})