diff --git a/src/basic_memory/cli/commands/cloud/core_commands.py b/src/basic_memory/cli/commands/cloud/core_commands.py index 2dc2e25e..f689097a 100644 --- a/src/basic_memory/cli/commands/cloud/core_commands.py +++ b/src/basic_memory/cli/commands/cloud/core_commands.py @@ -31,10 +31,74 @@ from basic_memory.cli.commands.cloud.rclone_installer import ( RcloneInstallError, install_rclone, ) +from basic_memory.mcp.project_context import get_available_workspaces console = Console() +async def _select_default_workspace_on_login() -> None: + """Prompt workspace selection after login when multiple workspaces exist. + + Single workspace: auto-set as default silently. + Multiple workspaces: show a numbered list and prompt for selection. + Failure is non-fatal — user can always run 'bm cloud workspace set-default'. + """ + try: + workspaces = await get_available_workspaces() + except Exception: + console.print( + "[dim]Workspace discovery unavailable; run 'bm cloud workspace set-default' if needed.[/dim]" + ) + return + + if not workspaces: + return + + config_manager = ConfigManager() + config = config_manager.config + + if len(workspaces) == 1: + config.default_workspace = workspaces[0].tenant_id + config_manager.save_config(config) + console.print(f"[dim]Default workspace: {workspaces[0].name}[/dim]") + return + + # Multiple workspaces — prompt user to pick one. + console.print("\n[bold]Multiple workspaces available:[/bold]") + for i, ws in enumerate(workspaces, 1): + console.print(f" {i}. {ws.name} ({ws.workspace_type}) — {ws.tenant_id}") + + raw = typer.prompt( + "Select default workspace (number, or press Enter to skip)", + default="", + ) + raw = raw.strip() + + if not raw: + console.print( + "[dim]No default workspace set; run 'bm cloud workspace set-default' to choose.[/dim]" + ) + return + + try: + idx = int(raw) - 1 + except ValueError: + console.print( + f"[yellow]'{raw}' is not a valid number; run 'bm cloud workspace set-default' to choose.[/yellow]" + ) + return + + if 0 <= idx < len(workspaces): + selected = workspaces[idx] + config.default_workspace = selected.tenant_id + config_manager.save_config(config) + console.print(f"[green]Default workspace set to '{selected.name}'[/green]") + else: + console.print( + f"[yellow]Selection out of range; run 'bm cloud workspace set-default' to choose.[/yellow]" + ) + + @cloud_app.command() def login(): """Authenticate with WorkOS using OAuth Device Authorization flow.""" @@ -58,6 +122,11 @@ def login(): console.print("[green]Cloud authentication successful[/green]") console.print(f"[dim]Cloud host ready: {host_url}[/dim]") + # Prompt workspace selection when multiple are available so users + # don't get silently locked to a stale default_workspace from a + # previous session. + await _select_default_workspace_on_login() + except SubscriptionRequiredError as e: track(EVENT_CLOUD_LOGIN_SUB_REQUIRED) console.print("\n[red]Subscription Required[/red]\n") @@ -76,10 +145,21 @@ def login(): @cloud_app.command() def logout(): - """Remove stored OAuth tokens.""" - config = ConfigManager().config + """Remove stored OAuth tokens and reset workspace selection.""" + config_manager = ConfigManager() + config = config_manager.config auth = CLIAuth(client_id=config.cloud_client_id, authkit_domain=config.cloud_domain) auth.logout() + + # Trigger: session is ending, so any previously selected workspace is no + # longer meaningful for the next authenticated user. + # Why: prevents stale default_workspace from silently routing to the wrong + # tenant (e.g., an org workspace) on re-login. + # Outcome: next login will prompt workspace selection afresh. + if config.default_workspace is not None: + config.default_workspace = None + config_manager.save_config(config) + console.print("[dim]API key (if configured) remains available for cloud project routing.[/dim]") diff --git a/tests/cli/test_cloud_authentication.py b/tests/cli/test_cloud_authentication.py index 10c1526c..d1c59d8a 100644 --- a/tests/cli/test_cloud_authentication.py +++ b/tests/cli/test_cloud_authentication.py @@ -2,19 +2,25 @@ from __future__ import annotations +import tempfile from contextlib import asynccontextmanager +from pathlib import Path from typing import Any, cast import httpx import pytest from typer.testing import CliRunner +import basic_memory.config +import basic_memory.cli.commands.cloud.core_commands as core_cmd from basic_memory.cli.app import app from basic_memory.cli.commands.cloud.api_client import ( CloudAPIError, SubscriptionRequiredError, make_api_request, ) +from basic_memory.config import BasicMemoryConfig, ConfigManager +from basic_memory.schemas.cloud import WorkspaceInfo class _StubAuth: @@ -199,3 +205,207 @@ class TestLoginCommand: result = runner.invoke(app, ["cloud", "login"]) assert result.exit_code == 1 assert "Login failed" in result.stdout + + +# --------------------------------------------------------------------------- +# Shared config fixture for tests that inspect persisted config state +# --------------------------------------------------------------------------- + + +class _ConfigFixtureMixin: + """Sets up an isolated temp config directory for each test.""" + + @pytest.fixture(autouse=True) + def _setup_config(self, monkeypatch): + self.temp_dir = tempfile.mkdtemp() + temp_path = Path(self.temp_dir) + config_dir = temp_path / ".basic-memory" + config_dir.mkdir(parents=True, exist_ok=True) + monkeypatch.setenv("HOME", str(temp_path)) + monkeypatch.setenv("BASIC_MEMORY_CONFIG_DIR", str(config_dir)) + basic_memory.config._CONFIG_CACHE = None + basic_memory.config._CONFIG_MTIME = None + basic_memory.config._CONFIG_SIZE = None + self.config_manager = ConfigManager() + self.temp_path = temp_path + + def _reset_cache(self): + basic_memory.config._CONFIG_CACHE = None + basic_memory.config._CONFIG_MTIME = None + basic_memory.config._CONFIG_SIZE = None + + def _save_config(self, **kwargs): + cfg = BasicMemoryConfig( + projects={"main": {"path": str(self.temp_path / "main")}}, + **kwargs, + ) + self.config_manager.save_config(cfg) + self._reset_cache() + + +class _FakeAuthFactory: + """Produces a CLIAuth-compatible stub that never touches the filesystem.""" + + def __init__(self, login_ok: bool = True): + self._login_ok = login_ok + + def __call__(self, **_kwargs): + login_ok = self._login_ok + + class _Auth: + async def login(self) -> bool: + return login_ok + + def logout(self) -> None: + pass + + return _Auth() + + +class TestLogoutCommand(_ConfigFixtureMixin): + """Tests for 'bm cloud logout' — token clearing and workspace reset.""" + + def test_logout_clears_default_workspace(self, monkeypatch): + self._save_config(default_workspace="11111111-1111-1111-1111-111111111111") + + monkeypatch.setattr(core_cmd, "CLIAuth", _FakeAuthFactory()) + + runner = CliRunner() + result = runner.invoke(app, ["cloud", "logout"]) + assert result.exit_code == 0 + + self._reset_cache() + config = ConfigManager().config + assert config.default_workspace is None + + def test_logout_when_no_default_workspace(self, monkeypatch): + self._save_config() + + monkeypatch.setattr(core_cmd, "CLIAuth", _FakeAuthFactory()) + + runner = CliRunner() + result = runner.invoke(app, ["cloud", "logout"]) + assert result.exit_code == 0 + + self._reset_cache() + config = ConfigManager().config + assert config.default_workspace is None + + +class TestLoginWorkspaceSelection(_ConfigFixtureMixin): + """Tests for workspace selection step inside 'bm cloud login'.""" + + def _patch_login_deps(self, monkeypatch, workspaces, login_ok=True): + """Patch all login dependencies for a successful login scenario.""" + self._save_config() + + monkeypatch.setattr(core_cmd, "CLIAuth", _FakeAuthFactory(login_ok=login_ok)) + monkeypatch.setattr( + "basic_memory.cli.commands.cloud.core_commands.get_cloud_config", + lambda: ("client_id", "domain", "https://cloud.example.com"), + ) + + async def fake_make_api_request(*_args, **_kwargs): + return httpx.Response(200, json={"status": "healthy"}) + + monkeypatch.setattr(core_cmd, "make_api_request", fake_make_api_request) + + async def fake_get_workspaces(context=None): + return workspaces + + monkeypatch.setattr(core_cmd, "get_available_workspaces", fake_get_workspaces) + + def test_login_single_workspace_auto_sets_default(self, monkeypatch): + ws = WorkspaceInfo( + tenant_id="aaaa-1111", + workspace_type="personal", + name="Personal", + role="owner", + ) + self._patch_login_deps(monkeypatch, [ws]) + + runner = CliRunner() + result = runner.invoke(app, ["cloud", "login"]) + assert result.exit_code == 0 + assert "Personal" in result.stdout + + self._reset_cache() + config = ConfigManager().config + assert config.default_workspace == "aaaa-1111" + + def test_login_multiple_workspaces_user_selects(self, monkeypatch): + ws1 = WorkspaceInfo( + tenant_id="aaaa-1111", + workspace_type="personal", + name="Personal", + role="owner", + ) + ws2 = WorkspaceInfo( + tenant_id="bbbb-2222", + workspace_type="organization", + name="Team", + role="editor", + ) + self._patch_login_deps(monkeypatch, [ws1, ws2]) + + runner = CliRunner() + # User types "2" to select the second workspace + result = runner.invoke(app, ["cloud", "login"], input="2\n") + assert result.exit_code == 0 + assert "Team" in result.stdout + + self._reset_cache() + config = ConfigManager().config + assert config.default_workspace == "bbbb-2222" + + def test_login_multiple_workspaces_user_skips(self, monkeypatch): + ws1 = WorkspaceInfo( + tenant_id="aaaa-1111", + workspace_type="personal", + name="Personal", + role="owner", + ) + ws2 = WorkspaceInfo( + tenant_id="bbbb-2222", + workspace_type="organization", + name="Team", + role="editor", + ) + self._patch_login_deps(monkeypatch, [ws1, ws2]) + + runner = CliRunner() + # User presses Enter to skip selection + result = runner.invoke(app, ["cloud", "login"], input="\n") + assert result.exit_code == 0 + assert "bm cloud workspace set-default" in result.stdout + + self._reset_cache() + config = ConfigManager().config + # No workspace should be auto-set + assert config.default_workspace is None + + def test_login_workspace_discovery_failure_is_nonfatal(self, monkeypatch): + self._save_config() + + monkeypatch.setattr(core_cmd, "CLIAuth", _FakeAuthFactory()) + monkeypatch.setattr( + "basic_memory.cli.commands.cloud.core_commands.get_cloud_config", + lambda: ("client_id", "domain", "https://cloud.example.com"), + ) + + async def fake_make_api_request(*_args, **_kwargs): + return httpx.Response(200, json={"status": "healthy"}) + + monkeypatch.setattr(core_cmd, "make_api_request", fake_make_api_request) + + async def fail_get_workspaces(context=None): + raise RuntimeError("no connection") + + monkeypatch.setattr(core_cmd, "get_available_workspaces", fail_get_workspaces) + + runner = CliRunner() + result = runner.invoke(app, ["cloud", "login"]) + # Login should still succeed despite workspace discovery failure + assert result.exit_code == 0 + assert "Cloud authentication successful" in result.stdout + assert "Workspace discovery unavailable" in result.stdout