diff --git a/services/api/api/agent.py b/services/api/api/agent.py index 48f5ae407..f23ce13eb 100644 --- a/services/api/api/agent.py +++ b/services/api/api/agent.py @@ -100,9 +100,23 @@ _ENGINE_HARNESSES = {"amp", "claude-code", "codex", "pi-mono"} _REUSABLE_DB_STATES = {"running", "idle", "delivering", "error", "suspended"} +_CODEX_MODEL_PROFILES = frozenset({"fast", "think"}) IDLE_TTL_S = int(os.getenv("IDLE_TTL_S", "86400")) # 24 hours SUSPENDED_RETENTION_S = int(os.getenv("SUSPENDED_RETENTION_S", str(7 * 24 * 60 * 60))) + + +def _runtime_model_for_engine(engine: str | None, model: str | None) -> str | None: + normalized = (model or "").strip().lower() or None + if not normalized: + return None + if engine == "codex": + return normalized + if normalized in _CODEX_MODEL_PROFILES: + return None + return normalized + + MAX_ACTIVE_SANDBOX_SESSIONS = int(os.getenv("MAX_ACTIVE_SANDBOX_SESSIONS", "45")) STREAM_EOF_REATTACH_MAX = int(os.getenv("STREAM_EOF_REATTACH_MAX", "6")) STREAM_EOF_REATTACH_BACKOFF_S = float(os.getenv("STREAM_EOF_REATTACH_BACKOFF_S", "1.0")) @@ -200,6 +214,7 @@ async def _db_get_session(thread_key: str) -> SandboxSession | None: pool = _get_pool() row = await pool.fetchrow( "SELECT thread_key, sandbox_id, harness, engine, state, started_at, " + "model, " "agent_thread_id, last_delivered_id, inflight_turn_id, inflight_turn_input, " "inflight_attempts, last_result, trace_id " "FROM sandbox_sessions WHERE thread_key = $1", @@ -212,6 +227,7 @@ async def _db_get_session(thread_key: str) -> SandboxSession | None: thread_key=row["thread_key"], harness=row["harness"], engine=row["engine"], + model=row["model"] or "", started_at=row["started_at"].timestamp() if row["started_at"] else 0.0, backend_name="kubernetes", db_state=row["state"], @@ -254,18 +270,19 @@ async def _db_insert_session( session.trace_id = trace_id row = await pool.fetchrow( "INSERT INTO sandbox_sessions (" - "thread_key, sandbox_id, harness, engine, state, started_at, " + "thread_key, sandbox_id, harness, engine, model, state, started_at, " "agent_thread_id, last_delivered_id, inflight_turn_id, inflight_turn_input, " "inflight_started_at, inflight_attempts, last_result, last_result_at, trace_id" - ") VALUES ($1, $2, $3, $4, $5, NOW(), $6, $7, $8::text, $9::jsonb, " - "CASE WHEN $8::text IS NULL THEN NULL ELSE NOW() END, $10, $11, " - "CASE WHEN $11::text = '' THEN NULL ELSE NOW() END, $12::uuid) " + ") VALUES ($1, $2, $3, $4, $5, $6, NOW(), $7, $8, $9::text, $10::jsonb, " + "CASE WHEN $9::text IS NULL THEN NULL ELSE NOW() END, $11, $12, " + "CASE WHEN $12::text = '' THEN NULL ELSE NOW() END, $13::uuid) " "ON CONFLICT (thread_key) DO NOTHING " "RETURNING thread_key", session.thread_key, session.sandbox_id, harness, engine, + session.model or None, initial_state, agent_thread_id or None, last_delivered_id or None, @@ -836,7 +853,28 @@ async def get_or_spawn( old_last_result: str = "" old_trace_id: str = "" pool = _get_pool() + effective_harness = harness or default_harness() + resolved_engine, resolved_persona, repo = _resolve_harness_profile( + effective_harness, persona=persona, engine_override=engine + ) + requested_model = _runtime_model_for_engine(resolved_engine, model) + session = await _db_get_session(thread_key) + if session and requested_model and (session.model or "") != requested_model: + old_sandbox_id = session.sandbox_id + backend = get_backend() + with contextlib.suppress(Exception): + await backend.stop_by_id(old_sandbox_id) + await _db_delete_session(thread_key) + _drop_runtime(old_sandbox_id) + session = None + log.info( + "sandbox_replaced_for_model_profile", + thread_key=thread_key, + sandbox=old_sandbox_id[:12], + model=requested_model, + ) + if session: if session.db_state in _REUSABLE_DB_STATES: backend = get_backend() @@ -913,17 +951,10 @@ async def get_or_spawn( thread_trace_id = await get_or_create_thread_trace_id(pool, thread_key) - effective_harness = harness or default_harness() - - # Resolve harness profile (engine, persona, repo) once for both warm and cold paths - resolved_engine, resolved_persona, repo = _resolve_harness_profile( - effective_harness, persona=persona, engine_override=engine - ) - # Try warm pool first should_try_warm = ( not engine - and not model + and not requested_model and not old_agent_thread_id and not old_inflight_turn_id and not (effective_harness == "amp" and resolved_engine == "codex") @@ -958,9 +989,6 @@ async def get_or_spawn( return claimed # Cold spawn - resolved_engine, resolved_persona, repo = _resolve_harness_profile( - effective_harness, persona=persona, engine_override=engine - ) backend = get_backend() await _evict_idle_sessions_for_capacity(backend) trace_id = old_trace_id or thread_trace_id or str(uuid.uuid4()) @@ -970,7 +998,7 @@ async def get_or_spawn( resolved_engine, persona=resolved_persona, repo=repo, - model=model, + model=requested_model, resume_thread_id=old_agent_thread_id or None, trace_id=trace_id, ) diff --git a/services/api/api/runtime_control.py b/services/api/api/runtime_control.py index 260e8ab0c..9a3c1c85e 100644 --- a/services/api/api/runtime_control.py +++ b/services/api/api/runtime_control.py @@ -303,6 +303,7 @@ def _agent_session_title( # ── Per-message header (rendered italic at the top of every assistant message) ── _DEFAULT_CLAUDE_MODEL = "claude-opus-4-8" +_DEFAULT_CODEX_MODEL_PROFILE = "fast" _CLAUDE_MODEL_ALIASES: dict[str, str] = { "opus": _DEFAULT_CLAUDE_MODEL, @@ -327,6 +328,19 @@ def _resolve_codex_model_label(model: str | None) -> str: return f"codex-{raw}" +def _default_model_for_assignment( + *, + harness: str | None, + engine: str | None, + model: str | None, +) -> str | None: + if model: + return model + if (engine or harness or "").strip().lower() == "codex": + return _DEFAULT_CODEX_MODEL_PROFILE + return None + + def _engine_model_label( *, engine: str | None, @@ -576,7 +590,7 @@ async def _write_agents_override(runtime_id: str, agents_md_override: str) -> No async def get_active_assignment(pool, thread_key: str) -> dict[str, Any] | None: row = await pool.fetchrow( - "SELECT thread_key, assignment_generation, runtime_id, harness, engine, persona_id, " + "SELECT thread_key, assignment_generation, runtime_id, harness, engine, model, persona_id, " "prompt_ref, effective_agents_md_sha256, agents_md_override, state " "FROM agent_runtime_assignments " "WHERE thread_key = $1 AND state = 'active' " @@ -636,7 +650,11 @@ async def spawn_assignment( effective_harness = active_assignment.get("harness") or default_harness() effective_engine = active_assignment.get("engine") effective_persona_id = active_assignment.get("persona_id") - effective_model = None + effective_model = _default_model_for_assignment( + harness=effective_harness, + engine=effective_engine, + model=model, + ) effective_agents_md_override = active_assignment.get("agents_md_override") else: # Explicit harness wins; otherwise inherit from the persona's declared @@ -649,7 +667,11 @@ async def spawn_assignment( effective_harness = default_harness() effective_engine = engine effective_persona_id = persona_id - effective_model = model + effective_model = _default_model_for_assignment( + harness=effective_harness, + engine=effective_engine, + model=model, + ) effective_agents_md_override = agents_md_override payload = { @@ -738,7 +760,7 @@ async def spawn_assignment( return decode_jsonb(existing_idem["response_json"], {}) active = await conn.fetchrow( - "SELECT assignment_generation, runtime_id, harness, engine, persona_id, " + "SELECT assignment_generation, runtime_id, harness, engine, model, persona_id, " "prompt_ref, effective_agents_md_sha256, agents_md_override " "FROM agent_runtime_assignments " "WHERE thread_key = $1 AND state = 'active' " @@ -758,11 +780,15 @@ async def spawn_assignment( ) generation = int(active["assignment_generation"]) runtime_id = active["runtime_id"] - if runtime_id != session.sandbox_id: + if ( + runtime_id != session.sandbox_id + or (active["model"] or None) != (session.model or None) + ): await conn.execute( - "UPDATE agent_runtime_assignments SET runtime_id = $1, updated_at = NOW() " - "WHERE thread_key = $2 AND assignment_generation = $3", + "UPDATE agent_runtime_assignments SET runtime_id = $1, model = $2, updated_at = NOW() " + "WHERE thread_key = $3 AND assignment_generation = $4", session.sandbox_id, + session.model or None, thread_key, generation, ) @@ -784,14 +810,15 @@ async def spawn_assignment( ) await conn.execute( "INSERT INTO agent_runtime_assignments (" - "thread_key, assignment_generation, runtime_id, harness, engine, " + "thread_key, assignment_generation, runtime_id, harness, engine, model, " "persona_id, prompt_ref, effective_agents_md_sha256, agents_md_override, state" - ") VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, 'active')", + ") VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, 'active')", thread_key, generation, session.sandbox_id, session.harness, session.engine, + session.model or None, effective_persona_id, prompt_ref, prompt_sha, @@ -811,6 +838,7 @@ async def spawn_assignment( "assignment_state": assignment_state, "assignment_generation": generation, "persona_id": resolved_persona, + "model": session.model or None, "prompt_ref": resolved_prompt_ref, "effective_agents_md_sha256": resolved_prompt_sha, } @@ -2797,7 +2825,7 @@ async def _process_execution_impl(pool, row: dict[str, Any]) -> None: ) assignment = await pool.fetchrow( - "SELECT harness, engine, runtime_id, agents_md_override, persona_id, prompt_ref, effective_agents_md_sha256 " + "SELECT harness, engine, model, runtime_id, agents_md_override, persona_id, prompt_ref, effective_agents_md_sha256 " "FROM agent_runtime_assignments " "WHERE thread_key = $1 AND assignment_generation = $2", thread_key, @@ -2890,6 +2918,7 @@ async def _process_execution_impl(pool, row: dict[str, Any]) -> None: assignment["harness"], engine=assignment["engine"], persona=assignment["persona_id"], + model=assignment["model"], ) if session.sandbox_id != assignment["runtime_id"]: await pool.execute( diff --git a/services/api/api/sandbox/base.py b/services/api/api/sandbox/base.py index 5ce95785c..1d385f72f 100644 --- a/services/api/api/sandbox/base.py +++ b/services/api/api/sandbox/base.py @@ -36,6 +36,7 @@ class SandboxSession: thread_key: str harness: str engine: str + model: str = "" started_at: float = 0.0 backend_name: str = "" # "kubernetes", etc. db_state: str = "" diff --git a/services/api/api/sandbox/kubernetes.py b/services/api/api/sandbox/kubernetes.py index 6d4b51042..821a31175 100644 --- a/services/api/api/sandbox/kubernetes.py +++ b/services/api/api/sandbox/kubernetes.py @@ -1498,9 +1498,10 @@ async def create( if engine == "claude-code" and model: env.append(f"CLAUDE_MODEL={model}") if engine == "codex" and model: - if model != "fast": + if model not in {"fast", "think"}: raise ValueError(f"unknown Codex model profile: {model}") - env.append("CODEX_MODEL_PROFILE=fast") + if model == "fast": + env.append("CODEX_MODEL_PROFILE=fast") if engine == "claude-code" and resume_thread_id: env.append(f"CLAUDE_CONTINUE_SESSION_ID={resume_thread_id}") if persona: @@ -1722,6 +1723,7 @@ async def create( thread_key=thread_key, harness=harness, engine=engine, + model=model or "", started_at=time.time(), backend_name=self.name, trace_id=trace_id or "", diff --git a/services/api/api/workflow_engine.py b/services/api/api/workflow_engine.py index 3eff4d048..c3c04bf1a 100644 --- a/services/api/api/workflow_engine.py +++ b/services/api/api/workflow_engine.py @@ -993,6 +993,7 @@ async def _compute_agent_session_header( persona = selector.get("persona_id") harness = selector.get("harness") + model = selector.get("model") engine: str | None = None if not persona or not harness: active = await get_active_assignment(pool, thread_key) @@ -1000,12 +1001,14 @@ async def _compute_agent_session_header( persona = persona or _nonempty(active.get("persona_id")) harness = harness or _nonempty(active.get("harness")) engine = _nonempty(active.get("engine")) + model = model or _nonempty(active.get("model")) if persona and not engine: engine = _persona_default_engine(persona) return _agent_session_header( persona_id=persona, engine=engine, harness=harness, + model=model, ) @@ -1216,7 +1219,7 @@ async def _dispatch() -> dict[str, Any]: else: effective_delivery = dict(run_in.get("delivery") or {}) effective_history = history_messages or run_in.get("history_messages") or [] - selector = {"persona_id": persona, "harness": harness} + selector = {"persona_id": persona, "harness": harness, "model": model} slackbot_session_id: str | None = None try: diff --git a/services/api/api/workflows/slack_thread_turn.py b/services/api/api/workflows/slack_thread_turn.py index 373bf329c..d96094954 100644 --- a/services/api/api/workflows/slack_thread_turn.py +++ b/services/api/api/workflows/slack_thread_turn.py @@ -17,7 +17,10 @@ "claude": "claude-code", "pi": "pi-mono", } -_MODEL_PROFILE_FLAGS = frozenset({"fast"}) +_MODEL_PROFILE_FLAGS = { + "fast": "fast", + "think": "think", +} _PROMPT_FLAG_SKIP = frozenset({"engine", "model", "opus", "sonnet", "haiku"}) _PROMPT_FLAG_VALUE_SKIP = frozenset({"engine", "model"}) _PROMPT_FLAG_RE = re.compile( @@ -68,7 +71,8 @@ class PromptSelection: Both fields are optional and orthogonal: ``--invest`` sets only ``persona``, ``--claude`` sets only ``harness``, and ``--invest --claude`` - sets both. ``--fast`` selects the fast Codex model profile for the sandbox + sets both. Slack turns default to the fast Codex model profile; ``--think`` + opts into the deployment's higher-reasoning Codex profile for the sandbox spawned for this turn. The downstream resolver applies ``harness`` as the engine override and ``persona`` as the system-prompt overlay. """ @@ -140,7 +144,7 @@ def _classify_flag( if resolved in personas or flag in personas: return None, resolved, None if resolved in _MODEL_PROFILE_FLAGS: - return None, None, resolved + return None, None, _MODEL_PROFILE_FLAGS[resolved] return None, None, None diff --git a/services/api/db/migrations/041_track_runtime_model_profiles.sql b/services/api/db/migrations/041_track_runtime_model_profiles.sql new file mode 100644 index 000000000..aec58fd32 --- /dev/null +++ b/services/api/db/migrations/041_track_runtime_model_profiles.sql @@ -0,0 +1,15 @@ +-- migrate:up + +ALTER TABLE sandbox_sessions + ADD COLUMN IF NOT EXISTS model TEXT; + +ALTER TABLE agent_runtime_assignments + ADD COLUMN IF NOT EXISTS model TEXT; + +-- migrate:down + +ALTER TABLE agent_runtime_assignments + DROP COLUMN IF EXISTS model; + +ALTER TABLE sandbox_sessions + DROP COLUMN IF EXISTS model; diff --git a/services/api/tests/test_agent_control_plane.py b/services/api/tests/test_agent_control_plane.py index 9f9da03cb..667aa44a7 100644 --- a/services/api/tests/test_agent_control_plane.py +++ b/services/api/tests/test_agent_control_plane.py @@ -143,6 +143,7 @@ async def test_spawn_assignment_defaults_to_codex_when_no_selector( thread_key=thread_key, harness="codex", engine="codex", + model="fast", ) get_or_spawn = AsyncMock(return_value=session) @@ -158,16 +159,19 @@ async def test_spawn_assignment_defaults_to_codex_when_no_selector( agents_md_override=None, ) - get_or_spawn.assert_awaited_once_with(thread_key, "codex", engine=None) + get_or_spawn.assert_awaited_once_with( + thread_key, "codex", engine=None, model="fast" + ) assert result["persona_id"] is None assignment = await db_pool.fetchrow( - "SELECT harness, engine, persona_id FROM agent_runtime_assignments WHERE thread_key = $1", + "SELECT harness, engine, persona_id, model FROM agent_runtime_assignments WHERE thread_key = $1", thread_key, ) assert assignment is not None assert assignment["harness"] == "codex" assert assignment["engine"] == "codex" assert assignment["persona_id"] is None + assert assignment["model"] == "fast" @pytest.mark.asyncio @@ -218,6 +222,7 @@ async def test_spawn_assignment_forwards_model_profile(db_pool, monkeypatch): thread_key=thread_key, harness="codex", engine="codex", + model="fast", ) get_or_spawn = AsyncMock(return_value=session) @@ -236,6 +241,68 @@ async def test_spawn_assignment_forwards_model_profile(db_pool, monkeypatch): get_or_spawn.assert_awaited_once_with( thread_key, "codex", engine=None, model="fast" ) + assignment = await db_pool.fetchrow( + "SELECT model FROM agent_runtime_assignments WHERE thread_key = $1", + thread_key, + ) + assert assignment is not None + assert assignment["model"] == "fast" + + +@pytest.mark.asyncio +async def test_spawn_assignment_defaults_active_codex_assignment_to_fast_profile( + db_pool, monkeypatch +): + from api.runtime_control import prompt_identity, spawn_assignment + + monkeypatch.delenv("CENTAUR_DEFAULT_HARNESS", raising=False) + thread_key = f"slack:C-test:{uuid.uuid4().hex}:active-codex-fast-profile" + prompt_ref, prompt_sha = prompt_identity( + harness="codex", + persona_id=None, + agents_md_override=None, + ) + await db_pool.execute( + "INSERT INTO agent_runtime_assignments (" + "thread_key, assignment_generation, runtime_id, harness, engine, model, " + "persona_id, prompt_ref, effective_agents_md_sha256, state" + ") VALUES ($1, 1, $2, 'codex', 'codex', 'think', NULL, $3, $4, 'active')", + thread_key, + f"rt-{uuid.uuid4().hex[:8]}", + prompt_ref, + prompt_sha, + ) + session = SandboxSession( + sandbox_id=f"rt-{uuid.uuid4().hex[:8]}", + thread_key=thread_key, + harness="codex", + engine="codex", + model="fast", + ) + get_or_spawn = AsyncMock(return_value=session) + + with patch("api.runtime_control.get_or_spawn", new=get_or_spawn): + await spawn_assignment( + db_pool, + thread_key=thread_key, + spawn_id="spawn-existing-fast", + harness=None, + engine=None, + persona_id=None, + model=None, + agents_md_override=None, + ) + + get_or_spawn.assert_awaited_once_with( + thread_key, "codex", engine="codex", model="fast" + ) + assignment = await db_pool.fetchrow( + "SELECT assignment_generation, model FROM agent_runtime_assignments WHERE thread_key = $1", + thread_key, + ) + assert assignment is not None + assert assignment["assignment_generation"] == 1 + assert assignment["model"] == "fast" @pytest.mark.asyncio diff --git a/services/api/tests/test_sandbox_kubernetes_backend.py b/services/api/tests/test_sandbox_kubernetes_backend.py index 8df9199d7..99df8195a 100644 --- a/services/api/tests/test_sandbox_kubernetes_backend.py +++ b/services/api/tests/test_sandbox_kubernetes_backend.py @@ -2117,6 +2117,30 @@ async def test_create_sets_codex_fast_profile_env( assert env["CODEX_MODEL_PROFILE"] == "fast" +@pytest.mark.asyncio +async def test_create_accepts_codex_think_profile_without_fast_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + backend = KubernetesExecutorBackend() + fake_core = FakeCoreApi() + backend._core = fake_core + backend._networking = FakeNetworkingApi() + _stub_create_dependencies( + monkeypatch, + backend, + extra_env=[{"name": "CODEX_AUTH_MODE", "value": "access_token"}], + harness_cmd="codex-app-wrapper", + ) + + session = await backend.create("slack:C123:123.456", "codex", "codex", model="think") + + pod_body = fake_core.created_pods[-1][1] + container = pod_body["spec"]["containers"][0] + env = {item["name"]: item["value"] for item in container["env"]} + assert "CODEX_MODEL_PROFILE" not in env + assert session.model == "think" + + @pytest.mark.asyncio async def test_create_uses_api_key_in_default_mode( monkeypatch: pytest.MonkeyPatch, diff --git a/services/api/tests/test_workflows.py b/services/api/tests/test_workflows.py index 0c7749d86..dfe066b90 100644 --- a/services/api/tests/test_workflows.py +++ b/services/api/tests/test_workflows.py @@ -407,12 +407,14 @@ def test_recovery_command_paraphrases_are_recognized(): ("--claude review this", "claude-code", None, None, "review this"), ("--pi analyze this", "pi-mono", None, None, "analyze this"), ("--fast summarize this", None, None, "fast", "summarize this"), + ("--think summarize this", None, None, "think", "summarize this"), # Persona + harness compose orthogonally. ("--invest --claude review this", "claude-code", "invest", None, "review this"), ("--claude --invest review this", "claude-code", "invest", None, "review this"), ("--invest --amp review this", "amp", "invest", None, "review this"), ("--invest --codex review this", "codex", "invest", None, "review this"), ("--invest --fast review this", None, "invest", "fast", "review this"), + ("--invest --think review this", None, "invest", "think", "review this"), ("please use --opus and review this", None, None, None, "please use and review this"), ("please use --model opus and review this", None, None, None, "please use and review this"), ("please use `--model opus` and review this", None, None, None, "please use and review this"), @@ -857,6 +859,49 @@ async def test_slack_thread_turn_fast_profile_releases_assignment(db_pool): ] +@pytest.mark.asyncio +async def test_slack_thread_turn_think_profile_releases_assignment(db_pool): + from api.workflow_engine import WorkflowContext + from api.workflows.slack_thread_turn import Input, handler + + run_id = f"wfr_{uuid.uuid4().hex[:16]}" + thread_key = f"slack:C-test:{uuid.uuid4().hex}" + ctx = WorkflowContext( + pool=db_pool, + run_id=run_id, + checkpoints={}, + lease_s=30.0, + worker_id="w1", + ) + do_agent_turn_mock = AsyncMock(return_value={"ok": True, "execution_id": "exe-1"}) + release_assignment_mock = AsyncMock(return_value={"ok": True, "released": True}) + + with ( + patch("api.workflow_engine.do_agent_turn", new=do_agent_turn_mock), + patch("api.runtime_control.release_assignment", new=release_assignment_mock), + ): + await handler( + Input( + thread_key=thread_key, + parts=[{"type": "text", "text": "--think summarize this carefully"}], + message_id="slack:current", + ), + ctx, + ) + + release_assignment_mock.assert_awaited_once_with( + db_pool, + thread_key=thread_key, + release_id="prompt-switch:slack:current", + cancel_inflight=True, + stop_runtime_background=True, + ) + assert do_agent_turn_mock.await_args.kwargs["model"] == "think" + assert do_agent_turn_mock.await_args.kwargs["parts"] == [ + {"type": "text", "text": "summarize this carefully"} + ] + + @pytest.mark.asyncio async def test_prompt_switch_retry_still_hydrates_prior_ask_from_history(db_pool): from api.workflow_engine import WorkflowContext