Skip to content
Open
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
3 changes: 2 additions & 1 deletion amplifier_app_cli/commands/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,8 @@ def _prepare_resume_context(
store = SessionStore()
transcript, metadata = store.load(session_id)

# Extract bundle from saved session metadata
# SessionStore normalizes malformed parseable metadata at the shared read
# boundary; extract_session_mode also defensively ignores unusable bundles.
saved_bundle, _ = extract_session_mode(metadata)

bundle_name = None
Expand Down
67 changes: 54 additions & 13 deletions amplifier_app_cli/effective_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,50 @@ def format_banner_line(self) -> str:
return f"Bundle: {bundle_name} | Provider: {self.provider_name} | {self.model}"


@dataclass(frozen=True)
class EffectiveProviderModel:
"""Canonical provider/model provenance for a resolved session config."""

provider: str
model: str

def as_metadata(self) -> dict[str, str]:
"""Return the fields persisted by every durable session writer."""
return {"provider": self.provider, "model": self.model}


def get_effective_provider_model(config: dict[str, Any]) -> EffectiveProviderModel:
"""Resolve the provider/model pair that will handle the session.

Provider selection follows orchestrator priority semantics. Within the
selected provider's config, ``model`` is the supported explicit runtime
setting and therefore takes precedence over the legacy/default
``default_model`` setting.
"""
providers = config.get("providers", [])
selected_provider = (
_select_provider_by_priority(providers) if isinstance(providers, list) else None
)
if selected_provider is None:
return EffectiveProviderModel(provider="none", model="none")

provider = selected_provider.get("module")
if not isinstance(provider, str) or not provider:
provider = "unknown"

provider_config = selected_provider.get("config", {})
if not isinstance(provider_config, dict):
provider_config = {}

model = provider_config.get("model")
if not isinstance(model, str) or not model:
model = provider_config.get("default_model")
if not isinstance(model, str) or not model:
model = "default"

return EffectiveProviderModel(provider=provider, model=model)


def get_effective_config_summary(
config: dict[str, Any],
config_source: str = "default",
Expand All @@ -52,22 +96,14 @@ def get_effective_config_summary(
Returns:
EffectiveConfigSummary with display-friendly information
"""
# Extract provider info - select by priority (lowest number wins)
# This matches the orchestrator's _select_provider() logic
providers = config.get("providers", [])
selected_provider = _select_provider_by_priority(providers)

if selected_provider:
provider_module = selected_provider.get("module", "unknown")
provider_config = selected_provider.get("config", {})
model = provider_config.get("default_model", "default")

provenance = get_effective_provider_model(config)
provider_module = provenance.provider
model = provenance.model
if provider_module != "none":
# Try to get friendly provider name
provider_name = _get_provider_display_name(provider_module)
else:
provider_module = "none"
provider_name = "None"
model = "none"

# Extract orchestrator
session_config = config.get("session", {})
Expand Down Expand Up @@ -158,4 +194,9 @@ def _get_provider_display_name(provider_module: str) -> str:
return name_map.get(name, name.replace("-", " ").title())


__all__ = ["EffectiveConfigSummary", "get_effective_config_summary"]
__all__ = [
"EffectiveConfigSummary",
"EffectiveProviderModel",
"get_effective_config_summary",
"get_effective_provider_model",
]
22 changes: 3 additions & 19 deletions amplifier_app_cli/incremental_save.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
if TYPE_CHECKING:
from amplifier_core import AmplifierSession

from .effective_config import get_effective_provider_model
from .session_store import SessionStore

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -95,8 +96,7 @@ async def on_tool_post(self, event: str, data: dict[str, Any]):
# Update debounce counter
self._last_message_count = current_count

# Extract model name from config
model_name = self._extract_model_name()
provenance = get_effective_provider_model(self.config)

# Load existing metadata to preserve fields like name, description
# that may have been set by other hooks (e.g., session-naming)
Expand All @@ -110,7 +110,7 @@ async def on_tool_post(self, event: str, data: dict[str, Any]):
"created", datetime.now(UTC).isoformat()
),
"bundle": self.bundle_name,
"model": model_name,
**provenance.as_metadata(),
"turn_count": len([m for m in messages if m.get("role") == "user"]),
"incremental": True, # Distinguish from final saves
# Store working_dir for session sync between CLI and web
Expand All @@ -131,22 +131,6 @@ async def on_tool_post(self, event: str, data: dict[str, Any]):

return HookResult(action="continue")

def _extract_model_name(self) -> str:
"""Extract model name from session config.

Returns:
Model name string or "unknown" if not found
"""
providers = self.config.get("providers", [])
if isinstance(providers, list) and providers:
first_provider = providers[0]
if isinstance(first_provider, dict) and "config" in first_provider:
provider_config = first_provider["config"]
return provider_config.get("model") or provider_config.get(
"default_model", "unknown"
)
return "unknown"


def register_incremental_save(
session: "AmplifierSession",
Expand Down
22 changes: 8 additions & 14 deletions amplifier_app_cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,10 @@
from .commands.version import version as version_cmd
from .console import Markdown, console
from .dedicated_tty_input import close_dedicated_tty_input, get_dedicated_tty_input
from .effective_config import get_effective_config_summary
from .effective_config import (
get_effective_config_summary,
get_effective_provider_model,
)
from .key_manager import KeyManager
from .session_runner import SessionConfig, create_initialized_session
from .session_store import SessionStore
Expand Down Expand Up @@ -2881,17 +2884,6 @@ async def interactive_chat(
)
)

# Helper to extract model name from config
def _extract_model_name() -> str:
if isinstance(config.get("providers"), list) and config["providers"]:
first_provider = config["providers"][0]
if isinstance(first_provider, dict) and "config" in first_provider:
provider_config = first_provider["config"]
return provider_config.get("model") or provider_config.get(
"default_model", "unknown"
)
return "unknown"

# Helper to save session after each turn
async def _save_session():
context = session.coordinator.get("context")
Expand All @@ -2900,14 +2892,15 @@ async def _save_session():
# Load existing metadata to preserve fields like name, description
# that may have been set by other hooks (e.g., session-naming)
existing_metadata = store.get_metadata(actual_session_id) or {}
provenance = get_effective_provider_model(config)
metadata = {
**existing_metadata, # Preserve name, description, etc.
"session_id": actual_session_id,
"created": existing_metadata.get(
"created", datetime.now(UTC).isoformat()
),
"bundle": bundle_name,
"model": _extract_model_name(),
**provenance.as_metadata(),
"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 Down Expand Up @@ -3647,14 +3640,15 @@ def _goal_sigint_handler(signum, frame):
# Load existing metadata to preserve fields like name, description
# that may have been set by other hooks (e.g., session-naming)
existing_metadata = store.get_metadata(actual_session_id) or {}
provenance = get_effective_provider_model(config)
metadata = {
**existing_metadata, # Preserve name, description, etc.
"session_id": actual_session_id,
"created": existing_metadata.get(
"created", datetime.now(UTC).isoformat()
),
"bundle": bundle_name,
"model": model_name,
**provenance.as_metadata(),
"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 Down
Loading
Loading