Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 39 additions & 10 deletions amplifier_app_cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -4518,6 +4518,38 @@ async def after_cleanup():
console.print()


def _configured_model_for_reporting(session: Any) -> tuple[str, str]:
"""Display-only configured conversation model, not proof of actual routing.

Honor the conversation pin and streaming-loop priority ordering, retaining
insertion order for ties. Never borrow a model from an unselected mount.
"""
try:
providers = session.coordinator.get("providers") or {}
if not providers:
return "unknown", "unknown"
pinned = _pinned_provider_name(session)
if pinned is not None:
name = pinned
provider = providers.get(name)
if provider is None:
return "unknown", "unknown"
else:
name, provider = min(
providers.items(),
key=lambda item: CommandProcessor._provider_priority_for_display(item[1]),
)
model = CommandProcessor._provider_model_for_display(provider)
if not model:
model = getattr(provider, "model", None) or getattr(provider, "default_model", None)
if not isinstance(model, str) or not model.strip():
return "unknown", "unknown"
return f"{name}/{model}", "configured_default"
except Exception:
# Reporting must not turn a successful response into a failed session.
return "unknown", "unknown"


async def execute_single(
prompt: str,
config: dict,
Expand Down Expand Up @@ -4624,14 +4656,7 @@ async def execute_single(
execution_started = False
session_save_attempted = False

def _resolve_model_name() -> str:
providers = session.coordinator.get("providers") or {}
for prov_name, prov in providers.items():
if hasattr(prov, "model"):
return f"{prov_name}/{prov.model}"
if hasattr(prov, "default_model"):
return f"{prov_name}/{prov.default_model}"
return "unknown"
model_name, model_source = "unknown", "unknown"

async def _persist_session() -> int:
"""Write transcript + metadata. Returns the message count saved (0 = nothing)."""
Expand All @@ -4650,7 +4675,8 @@ async def _persist_session() -> int:
"session_id": actual_session_id,
"created": existing_metadata.get("created", datetime.now(UTC).isoformat()),
"bundle": bundle_name,
"model": _resolve_model_name(),
"model": model_name,
"model_source": model_source,
"turn_count": len([m for m in messages if m.get("role") == "user"]),
# Store working_dir for session sync between CLI and web
"working_dir": str(Path.cwd().resolve()),
Expand All @@ -4663,6 +4689,9 @@ async def _persist_session() -> int:
return len(messages)

try:
# Freeze the configured label under the cleanup boundary, before any
# execution hooks. It is not observed provider-side call telemetry.
model_name, model_source = _configured_model_for_reporting(session)
await handoff.start()
# Register trace collector hooks if in json-trace mode
if trace_collector:
Expand Down Expand Up @@ -4832,7 +4861,6 @@ def _goal_sigint_handler(signum, frame):

# Get metadata for output
actual_session_id = session.session_id
model_name = _resolve_model_name()

# Emit prompt:complete (canonical kernel event) BEFORE formatting output
# This ensures hook output goes to stderr in JSON mode
Expand Down Expand Up @@ -4862,6 +4890,7 @@ def _goal_sigint_handler(signum, frame):
"session_id": actual_session_id,
"bundle": bundle_name,
"model": model_name,
"model_source": model_source,
"timestamp": datetime.now(UTC).isoformat(),
}
# Add trace data if collecting
Expand Down
124 changes: 122 additions & 2 deletions tests/test_headless_session_persistence.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
import io
import json
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch

import pytest
Expand All @@ -28,7 +29,7 @@ def _private_console() -> Console:
return Console(file=io.StringIO())


def _initialized_session(*, response: str = "saved response") -> MagicMock:
def _initialized_session(*, response: str = "saved response", providers=None) -> MagicMock:
"""Return the smallest session double that exercises the final save path."""
context = MagicMock()
context.get_messages = AsyncMock(
Expand All @@ -42,7 +43,7 @@ def get_capability(name: str):
if name == "context":
return context
if name == "providers":
return {}
return providers or {}
return None

session = MagicMock()
Expand Down Expand Up @@ -727,3 +728,122 @@ def print_again() -> None:
second = CliRunner().invoke(cli, ["print-again"])
assert second.exit_code == 0, second.output
assert "SECOND INVOCATION" in second.stdout


@pytest.mark.asyncio
@pytest.mark.parametrize("output_format", ["json", "json-trace"])
@pytest.mark.parametrize("model", ["gpt-6.1-sol", "gpt-6-astra"])
async def test_headless_reports_nonfirst_selected_model_in_output_and_saved_metadata(
split_homes, tmp_path, capfd, output_format, model
):
from amplifier_app_cli.main import execute_single

providers = {
"vllm": SimpleNamespace(default_model="openai/gpt-oss-120b", priority=10),
"selected-openai": SimpleNamespace(
default_model=model, priority=0,
get_info=lambda: SimpleNamespace(defaults={"model": model}),
),
}
initialized = _initialized_session(providers=providers)
with (
patch(f"{_MAIN}.create_initialized_session", new=AsyncMock(return_value=initialized)),
patch(f"{_MAIN}.console", new=_private_console()),
):
await execute_single(
prompt="persist this", config={}, search_paths=[tmp_path], verbose=False,
session_id=_SESSION_ID, bundle_name="test-bundle", output_format=output_format,
)
output = json.loads(capfd.readouterr().out)
_, metadata = SessionStore().load(_SESSION_ID)
assert output["model"] == metadata["model"] == f"selected-openai/{model}"
assert output["model_source"] == metadata["model_source"] == "configured_default"
assert list(providers) == ["vllm", "selected-openai"]
assert providers["vllm"].priority == 10
assert providers["selected-openai"].priority == 0


@pytest.mark.parametrize(
"providers,expected",
[
({}, "unknown"),
({"first": SimpleNamespace(default_model="gpt-6-astra"),
"second": SimpleNamespace(default_model="gpt-6.1-sol")}, "first/gpt-6-astra"),
({"first": SimpleNamespace(default_model="gpt-6-astra", config={"priority": 10}),
"second": SimpleNamespace(default_model="gpt-6.1-sol", config={"priority": 0})},
"second/gpt-6.1-sol"),
({"first": SimpleNamespace(default_model="gpt-6-astra", priority=5, config={"priority": 0}),
"second": SimpleNamespace(default_model="gpt-6.1-sol", priority=1)}, "second/gpt-6.1-sol"),
({"selected": SimpleNamespace(priority=0),
"other": SimpleNamespace(default_model="gpt-6.1-sol", priority=10)}, "unknown"),
({"selected": SimpleNamespace(priority=0, default_model=None)}, "unknown"),
],
)
def test_configured_model_reporting_priority_and_unknown(providers, expected):
from amplifier_app_cli.main import _configured_model_for_reporting

coordinator = SimpleNamespace(
get=lambda name: providers, get_capability=lambda name: None,
)
label, source = _configured_model_for_reporting(SimpleNamespace(coordinator=coordinator))
assert label == expected
assert source == ("unknown" if expected == "unknown" else "configured_default")


@pytest.mark.parametrize("pinned,expected", [
("second", "second/gpt-6.1-sol"), ("unmounted", "unknown"),
])
def test_configured_model_reporting_honors_pin_without_silent_fallback(pinned, expected):
from amplifier_app_cli.main import _configured_model_for_reporting

providers = {
"first": SimpleNamespace(default_model="gpt-6-astra", priority=0),
"second": SimpleNamespace(default_model="gpt-6.1-sol", priority=100),
}
coordinator = SimpleNamespace(
get=lambda name: providers,
get_capability=lambda name: SimpleNamespace(current=lambda: pinned),
)
label, source = _configured_model_for_reporting(SimpleNamespace(coordinator=coordinator))
assert label == expected
assert source == ("unknown" if expected == "unknown" else "configured_default")


@pytest.mark.asyncio
@pytest.mark.parametrize("lookup_failure", ["pin", "model"])
async def test_headless_reporting_lookup_failure_preserves_response_and_cleanup(
split_homes, tmp_path, capfd, lookup_failure
):
from amplifier_app_cli.main import execute_single

class BrokenModel:
priority = 0

@property
def model(self):
raise OSError("Model display unavailable")

def broken_pin():
raise KeyError("Pin display unavailable")

providers = {"selected": BrokenModel() if lookup_failure == "model" else
SimpleNamespace(default_model="gpt-6.1-sol", priority=0)}
initialized = _initialized_session(providers=providers)
initialized.session.coordinator.get_capability = lambda name: (
SimpleNamespace(current=broken_pin) if lookup_failure == "pin" else None
)
with (
patch(f"{_MAIN}.create_initialized_session", new=AsyncMock(return_value=initialized)),
patch(f"{_MAIN}.console", new=_private_console()),
):
await execute_single(
prompt="persist this", config={}, search_paths=[tmp_path], verbose=False,
session_id=_SESSION_ID, bundle_name="test-bundle", output_format="json",
)
output = json.loads(capfd.readouterr().out)
_, metadata = SessionStore().load(_SESSION_ID)
assert output["status"] == "success"
assert output["response"] == "saved response"
assert output["model"] == metadata["model"] == "unknown"
assert output["model_source"] == metadata["model_source"] == "unknown"
initialized.cleanup.assert_awaited_once()
Loading