diff --git a/src/octopal/infrastructure/providers/codex_provider.py b/src/octopal/infrastructure/providers/codex_provider.py index 81d2f39b..c1013a29 100644 --- a/src/octopal/infrastructure/providers/codex_provider.py +++ b/src/octopal/infrastructure/providers/codex_provider.py @@ -4,13 +4,19 @@ import asyncio import contextlib +import hashlib import json import os import shutil import subprocess +import tempfile from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta from pathlib import Path -from typing import Any +from typing import Any, cast + +import structlog from octopal.infrastructure.config.models import LLMConfig from octopal.infrastructure.config.settings import Settings @@ -19,12 +25,109 @@ CODEX_REQUEST_TIMEOUT_SECONDS = 30.0 CODEX_TURN_TIMEOUT_SECONDS = 180.0 +CODEX_SESSION_TTL_DAYS = 30 +CODEX_SESSION_STATE_VERSION = 1 + +logger = structlog.get_logger(__name__) class CodexAppServerError(RuntimeError): pass +@dataclass(frozen=True) +class _CodexSession: + thread_id: str + config_fingerprint: str + message_fingerprints: tuple[str, ...] + updated_at: datetime + + +class _CodexSessionStore: + """Small provider-local registry for resumable Codex thread IDs.""" + + def __init__(self, state_dir: Path) -> None: + self._path = state_dir / "codex_sessions.json" + + def get(self, session_ref: str) -> tuple[_CodexSession | None, str | None]: + payload = self._read() + raw = (payload.get("sessions") or {}).get(session_ref) + if not isinstance(raw, dict): + return None, None + try: + updated_at = datetime.fromisoformat(str(raw["updated_at"])) + if updated_at.tzinfo is None: + updated_at = updated_at.replace(tzinfo=UTC) + session = _CodexSession( + thread_id=str(raw["thread_id"]), + config_fingerprint=str(raw["config_fingerprint"]), + message_fingerprints=tuple( + str(value) for value in (raw.get("message_fingerprints") or []) + ), + updated_at=updated_at, + ) + except (KeyError, TypeError, ValueError): + self.delete(session_ref) + return None, "invalid_state" + if datetime.now(UTC) - session.updated_at > timedelta(days=CODEX_SESSION_TTL_DAYS): + self.delete(session_ref) + return None, "expired" + return session, None + + def put(self, session_ref: str, session: _CodexSession) -> None: + payload = self._read() + sessions = payload.setdefault("sessions", {}) + if not isinstance(sessions, dict): + sessions = {} + payload["sessions"] = sessions + sessions[session_ref] = { + "thread_id": session.thread_id, + "config_fingerprint": session.config_fingerprint, + "message_fingerprints": list(session.message_fingerprints), + "updated_at": session.updated_at.astimezone(UTC).isoformat(), + } + self._write(payload) + + def delete(self, session_ref: str) -> None: + payload = self._read() + sessions = payload.get("sessions") + if not isinstance(sessions, dict) or session_ref not in sessions: + return + sessions.pop(session_ref, None) + self._write(payload) + + def _read(self) -> dict[str, Any]: + try: + payload = json.loads(self._path.read_text(encoding="utf-8")) + except FileNotFoundError: + return {"version": CODEX_SESSION_STATE_VERSION, "sessions": {}} + except (OSError, json.JSONDecodeError): + logger.warning("Codex session registry is unreadable; starting clean") + return {"version": CODEX_SESSION_STATE_VERSION, "sessions": {}} + if not isinstance(payload, dict) or payload.get("version") != CODEX_SESSION_STATE_VERSION: + return {"version": CODEX_SESSION_STATE_VERSION, "sessions": {}} + return payload + + def _write(self, payload: dict[str, Any]) -> None: + self._path.parent.mkdir(parents=True, exist_ok=True) + fd, temporary_name = tempfile.mkstemp( + prefix=f".{self._path.name}.", suffix=".tmp", dir=str(self._path.parent) + ) + temporary_path = Path(temporary_name) + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + json.dump(payload, handle, ensure_ascii=False, indent=2, sort_keys=True) + handle.write("\n") + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary_path, self._path) + if os.name != "nt": + os.chmod(self._path, 0o600) + except Exception: + temporary_path.unlink(missing_ok=True) + raise + + class _CodexAppServerClient: def __init__(self, command: str, args: list[str], env: dict[str, str]) -> None: self._command = command @@ -226,7 +329,9 @@ def __init__( self._profile = resolve_litellm_profile( settings, model_override=model, config_override=config ) - self._model = self._profile.raw_model or self._profile.model + self._model = cast(str, self._profile.raw_model or self._profile.model) + self._sessions = _CodexSessionStore(Path(getattr(settings, "state_dir", Path("data")))) + self._session_locks: dict[str, asyncio.Lock] = {} @property def provider_id(self) -> str: @@ -237,8 +342,13 @@ def model_id(self) -> str: return self._model async def complete(self, messages: list[Message | dict], **kwargs: object) -> str: - result = await self._run_turn(messages, tools=None, on_partial=None) - return result["content"] + result = await self._run_turn( + messages, + tools=None, + on_partial=None, + session_key=_session_key_from_kwargs(kwargs), + ) + return cast(str, result["content"]) async def complete_stream( self, @@ -247,8 +357,13 @@ async def complete_stream( on_partial: Callable[[str], Awaitable[None]], **kwargs: object, ) -> str: - result = await self._run_turn(messages, tools=None, on_partial=on_partial) - return result["content"] + result = await self._run_turn( + messages, + tools=None, + on_partial=on_partial, + session_key=_session_key_from_kwargs(kwargs), + ) + return cast(str, result["content"]) async def complete_with_tools( self, @@ -258,7 +373,12 @@ async def complete_with_tools( tool_choice: str = "auto", **kwargs: object, ) -> dict: - result = await self._run_turn(messages, tools=tools, on_partial=None) + result = await self._run_turn( + messages, + tools=tools, + on_partial=None, + session_key=_session_key_from_kwargs(kwargs), + ) return { "content": result["content"], "tool_calls": result["tool_calls"], @@ -271,6 +391,26 @@ async def _run_turn( *, tools: list[dict] | None, on_partial: Callable[[str], Awaitable[None]] | None, + session_key: str | None, + ) -> dict[str, Any]: + if session_key: + session_ref = _fingerprint(session_key) + lock = self._session_locks.setdefault(session_ref, asyncio.Lock()) + async with lock: + return await self._run_session_turn( + messages, + tools=tools, + on_partial=on_partial, + session_ref=session_ref, + ) + return await self._run_ephemeral_turn(messages, tools=tools, on_partial=on_partial) + + async def _run_ephemeral_turn( + self, + messages: list[Message | dict], + *, + tools: list[dict] | None, + on_partial: Callable[[str], Awaitable[None]] | None, ) -> dict[str, Any]: client = _CodexAppServerClient(_codex_command(), _codex_args(), _codex_env()) await client.start() @@ -325,6 +465,136 @@ async def _run_turn( finally: await client.close() + async def _run_session_turn( + self, + messages: list[Message | dict], + *, + tools: list[dict] | None, + on_partial: Callable[[str], Awaitable[None]] | None, + session_ref: str, + ) -> dict[str, Any]: + client = _CodexAppServerClient(_codex_command(), _codex_args(), _codex_env()) + await client.start() + resumed = False + try: + instructions, full_input_items = _messages_to_codex_input(messages) + message_fingerprints = _message_fingerprints(messages) + dynamic_tools = _tools_to_dynamic_tools(tools or []) + cwd = str(Path.cwd()) + config_fingerprint = _session_config_fingerprint( + model=self._model, + cwd=cwd, + effort=_normalize_effort(getattr(self._settings, "codex_reasoning_effort", None)), + dynamic_tools=dynamic_tools, + ) + session, reset_reason = self._sessions.get(session_ref) + if session is not None and session.config_fingerprint != config_fingerprint: + self._sessions.delete(session_ref) + session = None + reset_reason = "configuration_changed" + if reset_reason: + _log_session_state("reset", session_ref, reason=reset_reason) + + thread_id: str | None = None + input_items = full_input_items + if session is not None: + try: + resumed_thread = await client.request( + "thread/resume", + { + "threadId": session.thread_id, + "model": self._model, + "cwd": cwd, + "approvalPolicy": "never", + "sandbox": "read-only", + "personality": "none", + }, + timeout=CODEX_TURN_TIMEOUT_SECONDS, + ) + thread_id = ((resumed_thread or {}).get("thread") or {}).get("id") + if not thread_id: + raise CodexAppServerError("Codex did not return a resumed thread id") + input_items = _incremental_codex_input( + messages, + previous_fingerprints=session.message_fingerprints, + ) + resumed = True + _log_session_state("resumed", session_ref, thread_id=thread_id) + except Exception as exc: + self._sessions.delete(session_ref) + _log_session_state( + "reset", + session_ref, + thread_id=session.thread_id, + reason="resume_failed", + error=type(exc).__name__, + ) + + if thread_id is None: + thread = await client.request( + "thread/start", + _compact( + { + "model": self._model, + "cwd": cwd, + "approvalPolicy": "never", + "sandbox": "read-only", + "developerInstructions": instructions or None, + "personality": "none", + "serviceName": "octopal", + "ephemeral": False, + "environments": [], + "dynamicTools": dynamic_tools or None, + } + ), + timeout=CODEX_TURN_TIMEOUT_SECONDS, + ) + thread_id = ((thread or {}).get("thread") or {}).get("id") + if not thread_id: + raise CodexAppServerError("Codex did not return a thread id") + input_items = full_input_items + _log_session_state("created", session_ref, thread_id=thread_id) + + turn = await client.request( + "turn/start", + _compact( + { + "threadId": thread_id, + "input": input_items, + "cwd": cwd, + "model": self._model, + "approvalPolicy": "never", + "sandboxPolicy": {"type": "readOnly", "networkAccess": False}, + "effort": _normalize_effort( + getattr(self._settings, "codex_reasoning_effort", None) + ), + "environments": [], + } + ), + timeout=CODEX_TURN_TIMEOUT_SECONDS, + ) + turn_id = ((turn or {}).get("turn") or {}).get("id") + result = await _collect_turn( + client, thread_id=thread_id, turn_id=turn_id, on_partial=on_partial + ) + self._sessions.put( + session_ref, + _CodexSession( + thread_id=thread_id, + config_fingerprint=config_fingerprint, + message_fingerprints=message_fingerprints, + updated_at=datetime.now(UTC), + ), + ) + return result + except Exception: + if resumed: + self._sessions.delete(session_ref) + _log_session_state("reset", session_ref, reason="resumed_turn_failed") + raise + finally: + await client.close() + async def _collect_turn( client: _CodexAppServerClient, @@ -438,12 +708,109 @@ def _content_to_text(content: Any) -> str: return json.dumps(content, ensure_ascii=False) +def _session_key_from_kwargs(kwargs: dict[str, object]) -> str | None: + value = kwargs.get("codex_session_key") + if value is None: + return None + normalized = str(value).strip() + return normalized or None + + +def _message_fingerprints(messages: list[Message | dict]) -> tuple[str, ...]: + return tuple(_fingerprint(_message_payload(message)) for message in messages) + + +def _message_payload(message: Message | dict) -> str: + data = message.to_dict() if isinstance(message, Message) else dict(message) + return json.dumps(data, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str) + + +def _incremental_codex_input( + messages: list[Message | dict], + *, + previous_fingerprints: tuple[str, ...], +) -> list[dict[str, str]]: + current_fingerprints = _message_fingerprints(messages) + suffix: list[Message | dict] + prefix_len = len(previous_fingerprints) + if ( + prefix_len <= len(current_fingerprints) + and current_fingerprints[:prefix_len] == previous_fingerprints + ): + suffix = messages[prefix_len:] + else: + suffix = [] + for index in range(len(messages) - 1, -1, -1): + message = messages[index] + data = message.to_dict() if isinstance(message, Message) else dict(message) + if str(data.get("role") or "").lower() == "user": + suffix = messages[index:] + break + + chunks: list[str] = [] + for message in suffix: + data = message.to_dict() if isinstance(message, Message) else dict(message) + role = str(data.get("role") or "message").upper() + content = _content_to_text(data.get("content")) + if content: + chunks.append(f"{role}:\n{content}") + text = "\n\n".join(chunks).strip() or "Continue." + return [{"type": "text", "text": text}] + + +def _session_config_fingerprint( + *, + model: str, + cwd: str, + effort: str | None, + dynamic_tools: list[dict[str, Any]], +) -> str: + tool_catalog = sorted(dynamic_tools, key=lambda item: str(item.get("name") or "")) + payload = { + "bridgeVersion": CODEX_SESSION_STATE_VERSION, + "model": model, + "cwd": cwd, + "effort": effort, + "approvalPolicy": "never", + "sandboxPolicy": {"type": "readOnly", "networkAccess": False}, + "dynamicTools": tool_catalog, + "command": _codex_command(), + "args": _codex_args(), + } + return _fingerprint( + json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + ) + + +def _fingerprint(value: str) -> str: + return hashlib.sha256(value.encode("utf-8")).hexdigest() + + +def _log_session_state( + state: str, + session_ref: str, + *, + thread_id: str | None = None, + reason: str | None = None, + error: str | None = None, +) -> None: + logger.info( + "Codex provider session state changed", + state=state, + session_ref=session_ref[:12], + thread_ref=_fingerprint(thread_id)[:12] if thread_id else None, + reason=reason, + error=error, + ) + + def _tools_to_dynamic_tools(tools: list[dict[str, Any]]) -> list[dict[str, Any]]: dynamic_tools: list[dict[str, Any]] = [] for tool in tools: if tool.get("type") != "function": continue - function = tool.get("function") if isinstance(tool.get("function"), dict) else {} + raw_function = tool.get("function") + function: dict[str, Any] = raw_function if isinstance(raw_function, dict) else {} name = str(function.get("name") or "").strip() if not name: continue diff --git a/src/octopal/runtime/octo/route_completion.py b/src/octopal/runtime/octo/route_completion.py index ce8138aa..7029e1fd 100644 --- a/src/octopal/runtime/octo/route_completion.py +++ b/src/octopal/runtime/octo/route_completion.py @@ -2,7 +2,7 @@ import json from collections.abc import Awaitable, Callable -from typing import Any +from typing import Any, cast import structlog @@ -18,14 +18,19 @@ async def _complete_text( *, context: str, on_partial: Callable[[str], Awaitable[None]] | None = None, + provider_kwargs: dict[str, object] | None = None, ) -> str: sanitized = _sanitize_messages_for_complete(messages) + call_kwargs = dict(provider_kwargs or {}) try: if callable(on_partial): stream_callable = getattr(provider, "complete_stream", None) if callable(stream_callable): - return await stream_callable(sanitized, on_partial=on_partial) - text = await provider.complete(sanitized) + return cast( + str, + await stream_callable(sanitized, on_partial=on_partial, **call_kwargs), + ) + text = cast(str, await provider.complete(sanitized, **call_kwargs)) if callable(on_partial) and text: try: await on_partial(text) diff --git a/src/octopal/runtime/octo/router.py b/src/octopal/runtime/octo/router.py index 0a0dfe47..78fdaac3 100644 --- a/src/octopal/runtime/octo/router.py +++ b/src/octopal/runtime/octo/router.py @@ -3,7 +3,7 @@ import asyncio import json from collections.abc import Awaitable, Callable -from typing import Any +from typing import Any, cast import structlog @@ -78,6 +78,86 @@ _build_generic_worker_completion_message = _worker_results._build_generic_worker_completion_message _build_worker_result_payload = _worker_results._build_worker_result_payload _extract_worker_artifact_paths = _worker_results._extract_worker_artifact_paths + + +def _codex_session_key( + *, + octo: Any, + chat_id: int, + conversation_scope: str | None, + channel_context: dict[str, object] | None, +) -> str: + context = channel_context or {} + source_channel = str(context.get("source_channel") or "").strip() + if not source_channel: + settings = getattr(octo, "settings", None) + source_channel = str(getattr(settings, "user_channel", "chat") or "chat").strip() + scope = str(conversation_scope or "conversation").strip() or "conversation" + return f"{source_channel}:{scope}:{chat_id}" + + +def _provider_session_kwargs(ctx: dict[str, object]) -> dict[str, object]: + session_key = str(ctx.get("codex_session_key") or "").strip() + return {"codex_session_key": session_key} if session_key else {} + + +class _ProviderWithCallDefaults: + def __init__( + self, + provider: InferenceProvider, + call_defaults: dict[str, object], + ) -> None: + self._provider = provider + self._call_defaults = call_defaults + + def _kwargs(self, kwargs: dict[str, object]) -> dict[str, object]: + return {**self._call_defaults, **kwargs} + + async def complete( + self, + messages: list[Message | dict[str, Any]], + **kwargs: object, + ) -> str: + return cast( + str, + await self._provider.complete(messages, **self._kwargs(kwargs)), + ) + + async def complete_stream( + self, + messages: list[Message | dict[str, Any]], + *, + on_partial: Callable[[str], Awaitable[None]], + **kwargs: object, + ) -> str: + return cast( + str, + await self._provider.complete_stream( + messages, + on_partial=on_partial, + **self._kwargs(kwargs), + ), + ) + + async def complete_with_tools( + self, + messages: list[Message | dict[str, Any]], + *, + tools: list[dict], + tool_choice: str = "auto", + **kwargs: object, + ) -> dict: + return cast( + dict[str, Any], + await self._provider.complete_with_tools( + messages, + tools=tools, + tool_choice=tool_choice, + **self._kwargs(kwargs), + ), + ) + + _is_durable_workspace_artifact_path = _worker_results._is_durable_workspace_artifact_path _normalize_worker_artifact_path = _worker_results._normalize_worker_artifact_path _normalize_worker_result_entry = _worker_results._normalize_worker_result_entry @@ -277,6 +357,12 @@ async def route_or_reply( "internal_followup": internal_followup, "background_delivery": background_delivery, "conversation_scope": conversation_scope, + "codex_session_key": _codex_session_key( + octo=octo, + chat_id=chat_id, + conversation_scope=conversation_scope, + channel_context=channel_context, + ), } ) resolution_report = ctx.get("tool_resolution_report") @@ -378,7 +464,11 @@ async def route_or_reply( messages.append(Message(role="system", content=operational_memory_context)) _log_system_prompt(messages, "route") - plan = await _build_plan(provider, messages, bool(octo_tools)) + planner_provider = _ProviderWithCallDefaults( + provider, + _provider_session_kwargs(ctx), + ) + plan = await _build_plan(planner_provider, messages, bool(octo_tools)) if plan: routing_trace_metadata["planner_used"] = True logger.info( @@ -880,6 +970,7 @@ async def _complete_route_with_tools( tool_capable = getattr(provider, "complete_with_tools", None) trace_ctx = get_current_trace_context() trace_sink = getattr(octo, "trace_sink", None) + provider_kwargs = _provider_session_kwargs(ctx) if callable(tool_capable) and tool_specs: if trace_ctx is not None and trace_sink is not None: @@ -908,7 +999,10 @@ async def _complete_route_with_tools( for _ in range(max_attempts): try: result = await provider.complete_with_tools( - messages, tools=tools, tool_choice="auto" + messages, + tools=tools, + tool_choice="auto", + **provider_kwargs, ) except Exception as e: if ( @@ -960,6 +1054,7 @@ async def _complete_route_with_tools( provider, messages, context="saved_image_tool_retry_failed", + provider_kwargs=provider_kwargs, ) return await _finalize_response( provider=provider, @@ -1016,6 +1111,7 @@ async def _complete_route_with_tools( provider, messages, context="transient_tool_error_fallback", + provider_kwargs=provider_kwargs, ) return await _finalize_response( provider=provider, @@ -1189,6 +1285,7 @@ async def _complete_route_with_tools( provider, messages, context="octo_tool_loop_breaker", + provider_kwargs=provider_kwargs, ) return await _finalize_response( provider=provider, @@ -1337,6 +1434,7 @@ async def _complete_route_with_tools( provider, messages, context="empty_tool_response_fallback", + provider_kwargs=provider_kwargs, ) return await _finalize_response( provider=provider, @@ -1379,6 +1477,7 @@ async def _complete_route_with_tools( provider, messages, context="tool_limit_fallback", + provider_kwargs=provider_kwargs, ) return await _finalize_response( provider=provider, @@ -1401,6 +1500,7 @@ async def _complete_route_with_tools( provider, messages, context="tool_error_fallback", + provider_kwargs=provider_kwargs, ) return await _finalize_response( provider=provider, @@ -1416,6 +1516,7 @@ async def _complete_route_with_tools( messages, context="plain_completion", on_partial=on_plain_partial, + provider_kwargs=provider_kwargs, ) logger.debug("Octo output", output=response_raw) return await _finalize_response( diff --git a/tests/test_codex_provider_sessions.py b/tests/test_codex_provider_sessions.py new file mode 100644 index 00000000..1cf10cc2 --- /dev/null +++ b/tests/test_codex_provider_sessions.py @@ -0,0 +1,510 @@ +from __future__ import annotations + +import asyncio +import json +from pathlib import Path +from typing import Any + +import pytest + +from octopal.infrastructure.config.settings import Settings +from octopal.infrastructure.providers import codex_provider +from octopal.infrastructure.providers.base import Message +from octopal.infrastructure.providers.codex_provider import CodexAppServerError, CodexProvider +from octopal.runtime.octo.router import _complete_route_with_tools, route_or_reply +from octopal.tools.registry import ToolSpec + + +class _FakeCodexClient: + instances: list[_FakeCodexClient] = [] + next_thread = 1 + fail_resume_once = False + fail_turn = False + active_turns = 0 + max_active_turns = 0 + + def __init__(self, command: str, args: list[str], env: dict[str, str]) -> None: + del command, args, env + self.calls: list[tuple[str, dict[str, Any]]] = [] + self.thread_id = "" + self.turn_event_index = 0 + self.turn_active = False + type(self).instances.append(self) + + async def start(self) -> None: + return None + + async def request( + self, + method: str, + params: dict[str, Any] | None = None, + *, + timeout: float, + ) -> Any: + del timeout + payload = dict(params or {}) + self.calls.append((method, payload)) + if method == "thread/start": + self.thread_id = f"thread-{type(self).next_thread}" + type(self).next_thread += 1 + return {"thread": {"id": self.thread_id}} + if method == "thread/resume": + if type(self).fail_resume_once: + type(self).fail_resume_once = False + raise CodexAppServerError("missing thread") + self.thread_id = str(payload["threadId"]) + return {"thread": {"id": self.thread_id}} + if method == "turn/start": + self.turn_event_index = 0 + self.turn_active = True + type(self).active_turns += 1 + type(self).max_active_turns = max(type(self).max_active_turns, type(self).active_turns) + return {"turn": {"id": f"turn-{len(type(self).instances)}"}} + raise AssertionError(f"unexpected request: {method}") + + async def next_event(self, timeout: float) -> tuple[str, dict[str, Any]]: + del timeout + await asyncio.sleep(0.01) + if type(self).fail_turn: + self._finish_turn() + raise CodexAppServerError("turn failed") + self.turn_event_index += 1 + if self.turn_event_index == 1: + return ( + "notification", + { + "method": "item/agentMessage/delta", + "params": {"threadId": self.thread_id, "delta": "ok"}, + }, + ) + self._finish_turn() + return ( + "notification", + {"method": "turn/completed", "params": {"threadId": self.thread_id}}, + ) + + async def close(self) -> None: + self._finish_turn() + + def _finish_turn(self) -> None: + if self.turn_active: + self.turn_active = False + type(self).active_turns -= 1 + + +@pytest.fixture(autouse=True) +def _fake_codex_client(monkeypatch: pytest.MonkeyPatch) -> None: + _FakeCodexClient.instances = [] + _FakeCodexClient.next_thread = 1 + _FakeCodexClient.fail_resume_once = False + _FakeCodexClient.fail_turn = False + _FakeCodexClient.active_turns = 0 + _FakeCodexClient.max_active_turns = 0 + monkeypatch.setattr(codex_provider, "_CodexAppServerClient", _FakeCodexClient) + + +def _settings(state_dir: Path, *, model: str = "gpt-5.4") -> Settings: + return Settings( + OCTOPAL_STATE_DIR=state_dir, + OCTOPAL_LITELLM_PROVIDER_ID="codex", + OCTOPAL_LITELLM_MODEL=model, + ) + + +def _tool(name: str = "lookup") -> dict[str, Any]: + return { + "type": "function", + "function": { + "name": name, + "description": name, + "parameters": {"type": "object", "properties": {}}, + }, + } + + +def _call(client: _FakeCodexClient, method: str) -> dict[str, Any]: + return next(payload for candidate, payload in client.calls if candidate == method) + + +def test_persistent_session_resumes_after_provider_restart_with_only_new_input( + tmp_path: Path, +) -> None: + async def scenario() -> None: + first = CodexProvider(_settings(tmp_path)) + await first.complete_with_tools( + [ + {"role": "system", "content": "developer and memory context"}, + {"role": "user", "content": "first request"}, + ], + tools=[_tool()], + codex_session_key="telegram:primary:101", + ) + + created = _FakeCodexClient.instances[0] + assert _call(created, "thread/start")["ephemeral"] is False + assert _call(created, "thread/start")["developerInstructions"] == ( + "developer and memory context" + ) + + restarted = CodexProvider(_settings(tmp_path)) + await restarted.complete_with_tools( + [ + {"role": "system", "content": "repacked memory must not be resent"}, + {"role": "user", "content": "first request"}, + {"role": "assistant", "content": "first answer"}, + {"role": "user", "content": "second request"}, + ], + tools=[_tool()], + codex_session_key="telegram:primary:101", + ) + + resumed = _FakeCodexClient.instances[1] + assert _call(resumed, "thread/resume")["threadId"] == "thread-1" + assert "developerInstructions" not in _call(resumed, "thread/resume") + assert _call(resumed, "turn/start")["input"] == [ + {"type": "text", "text": "USER:\nsecond request"} + ] + + registry = json.loads((tmp_path / "codex_sessions.json").read_text()) + assert len(registry["sessions"]) == 1 + assert "telegram:primary:101" not in json.dumps(registry) + + asyncio.run(scenario()) + + +def test_tool_catalog_change_starts_a_fresh_persistent_thread(tmp_path: Path) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + messages = [{"role": "user", "content": "request"}] + await provider.complete_with_tools( + messages, + tools=[_tool("lookup")], + codex_session_key="telegram:primary:102", + ) + await provider.complete_with_tools( + messages, + tools=[_tool("write")], + codex_session_key="telegram:primary:102", + ) + + second = _FakeCodexClient.instances[1] + assert [method for method, _ in second.calls] == ["thread/start", "turn/start"] + assert _call(second, "thread/start")["ephemeral"] is False + assert _call(second, "thread/start")["dynamicTools"][0]["name"] == "write" + + asyncio.run(scenario()) + + +def test_warm_tool_loop_sends_only_the_appended_tool_result(tmp_path: Path) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + initial_messages = [ + {"role": "system", "content": "large initial context"}, + {"role": "user", "content": "look this up"}, + ] + await provider.complete_with_tools( + initial_messages, + tools=[_tool()], + codex_session_key="telegram:default:108", + ) + await provider.complete_with_tools( + [ + *initial_messages, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "name": "lookup", + "tool_call_id": "call-1", + "content": "fresh tool result", + }, + ], + tools=[_tool()], + codex_session_key="telegram:default:108", + ) + + resumed = _FakeCodexClient.instances[1] + assert _call(resumed, "turn/start")["input"] == [ + {"type": "text", "text": "TOOL:\nfresh tool result"} + ] + + asyncio.run(scenario()) + + +def test_mismatched_prompt_prefix_keeps_terminal_tool_result(tmp_path: Path) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + await provider.complete_with_tools( + [ + {"role": "system", "content": "initial context"}, + {"role": "user", "content": "look this up"}, + ], + tools=[_tool()], + codex_session_key="telegram:default:109", + ) + await provider.complete_with_tools( + [ + {"role": "system", "content": "changed packed context"}, + {"role": "user", "content": "look this up"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "name": "lookup", + "tool_call_id": "call-1", + "content": "terminal tool result", + }, + ], + tools=[_tool()], + codex_session_key="telegram:default:109", + ) + + resumed = _FakeCodexClient.instances[1] + assert _call(resumed, "thread/resume")["threadId"] == "thread-1" + assert _call(resumed, "turn/start")["input"] == [ + { + "type": "text", + "text": "USER:\nlook this up\n\nTOOL:\nterminal tool result", + } + ] + + asyncio.run(scenario()) + + +def test_resume_failure_falls_back_to_full_context_and_replaces_mapping(tmp_path: Path) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + await provider.complete_with_tools( + [{"role": "user", "content": "first"}], + tools=[_tool()], + codex_session_key="desktop:primary:103", + ) + _FakeCodexClient.fail_resume_once = True + await provider.complete_with_tools( + [ + {"role": "system", "content": "full current context"}, + {"role": "user", "content": "second"}, + ], + tools=[_tool()], + codex_session_key="desktop:primary:103", + ) + + fallback = _FakeCodexClient.instances[1] + assert [method for method, _ in fallback.calls] == [ + "thread/resume", + "thread/start", + "turn/start", + ] + assert _call(fallback, "thread/start")["developerInstructions"] == ("full current context") + registry = json.loads((tmp_path / "codex_sessions.json").read_text()) + assert next(iter(registry["sessions"].values()))["thread_id"] == "thread-2" + + asyncio.run(scenario()) + + +def test_concurrent_turns_for_one_session_are_serialized(tmp_path: Path) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + await asyncio.gather( + provider.complete_with_tools( + [{"role": "user", "content": "one"}], + tools=[_tool()], + codex_session_key="telegram:primary:104", + ), + provider.complete_with_tools( + [{"role": "user", "content": "two"}], + tools=[_tool()], + codex_session_key="telegram:primary:104", + ), + ) + assert _FakeCodexClient.max_active_turns == 1 + assert "thread/resume" in [method for method, _ in _FakeCodexClient.instances[1].calls] + + asyncio.run(scenario()) + + +def test_distinct_conversation_keys_never_share_a_thread(tmp_path: Path) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + await provider.complete_with_tools( + [{"role": "user", "content": "telegram request"}], + tools=[_tool()], + codex_session_key="telegram:default:107", + ) + await provider.complete_with_tools( + [{"role": "user", "content": "desktop request"}], + tools=[_tool()], + codex_session_key="desktop:default:107", + ) + + assert [method for method, _ in _FakeCodexClient.instances[0].calls] == [ + "thread/start", + "turn/start", + ] + assert [method for method, _ in _FakeCodexClient.instances[1].calls] == [ + "thread/start", + "turn/start", + ] + registry = json.loads((tmp_path / "codex_sessions.json").read_text()) + assert len(registry["sessions"]) == 2 + + asyncio.run(scenario()) + + +def test_ephemeral_calls_and_failed_first_turn_do_not_create_session_state( + tmp_path: Path, +) -> None: + async def scenario() -> None: + provider = CodexProvider(_settings(tmp_path)) + await provider.complete([{"role": "user", "content": "isolated helper"}]) + assert _call(_FakeCodexClient.instances[0], "thread/start")["ephemeral"] is True + assert not (tmp_path / "codex_sessions.json").exists() + + _FakeCodexClient.fail_turn = True + with pytest.raises(CodexAppServerError, match="turn failed"): + await provider.complete_with_tools( + [{"role": "user", "content": "failed"}], + tools=[_tool()], + codex_session_key="telegram:primary:105", + ) + assert not (tmp_path / "codex_sessions.json").exists() + + asyncio.run(scenario()) + + +def test_main_octo_route_passes_the_scoped_session_key_to_provider() -> None: + class _Provider: + def __init__(self) -> None: + self.tool_kwargs: dict[str, object] = {} + + async def complete_with_tools( + self, messages, *, tools, tool_choice="auto", **kwargs: object + ) -> dict[str, Any]: + del messages, tools, tool_choice + self.tool_kwargs = dict(kwargs) + return {"content": "The result is ready.", "tool_calls": []} + + async def complete(self, messages, **kwargs: object) -> str: + del messages, kwargs + return '{"verdict":"final","confidence":1.0,"reason":"complete"}' + + class _Octo: + trace_sink = None + + async def scenario() -> None: + provider = _Provider() + result = await _complete_route_with_tools( + octo=_Octo(), + provider=provider, + messages=[{"role": "user", "content": "finish the request"}], + tool_specs=[ + ToolSpec( + name="lookup", + description="lookup", + parameters={"type": "object", "properties": {}}, + permission="read", + handler=lambda: None, + ) + ], + ctx={ + "chat_id": 106, + "codex_session_key": "telegram:default:106", + }, + internal_followup=False, + user_text="finish the request", + images=None, + allow_tool_catalog_expansion=False, + ) + assert result == "The result is ready." + assert provider.tool_kwargs == {"codex_session_key": "telegram:default:106"} + + asyncio.run(scenario()) + + +def test_main_conversation_planner_reply_passes_scoped_session_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _Provider: + def __init__(self) -> None: + self.complete_kwargs: list[dict[str, object]] = [] + + async def complete(self, messages, **kwargs: object) -> str: + del messages + self.complete_kwargs.append(dict(kwargs)) + return '{"mode":"reply","steps":[],"response":"Hello from Alice."}' + + class _Memory: + async def add_message(self, role, content, metadata=None) -> None: + del role, content, metadata + + class _Octo: + store = object() + canon = object() + is_ws_active = False + internal_progress_send = None + trace_sink = None + + async def set_typing(self, chat_id: int, active: bool) -> None: + del chat_id, active + + async def set_thinking(self, active: bool) -> None: + del active + + def peek_context_wakeup(self, chat_id: int) -> str: + del chat_id + return "" + + async def fake_build_octo_prompt(**kwargs: object) -> list[Message]: + return [Message(role="user", content=str(kwargs["user_text"]))] + + async def no_action_retry(**kwargs: object) -> bool: + del kwargs + return False + + async def finalize_response(**kwargs: object) -> str: + return str(kwargs["response_text"]) + + import octopal.runtime.octo.router as router + + monkeypatch.setattr(router, "build_octo_prompt", fake_build_octo_prompt) + monkeypatch.setattr(router, "_needs_action_or_blocked_retry", no_action_retry) + monkeypatch.setattr(router, "_finalize_response", finalize_response) + monkeypatch.setattr( + router, + "_get_octo_tools", + lambda octo, chat_id: ([], {"octo": octo, "chat_id": chat_id}), + ) + + async def scenario() -> None: + provider = _Provider() + response = await route_or_reply( + _Octo(), + provider, + _Memory(), + "hello", + 211619002, + "", + conversation_scope="primary", + channel_context={"source_channel": "telegram"}, + ) + + assert response == "Hello from Alice." + assert provider.complete_kwargs == [{"codex_session_key": "telegram:primary:211619002"}] + + asyncio.run(scenario())