diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index b8266b7d6eb..73f69488499 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -23,6 +23,7 @@ from .claude_settings import ( write_claude_settings, ) from .pkce_login import ( + Http, PkceFailure, fresh_api_key, pkce_token_record, @@ -672,12 +673,26 @@ def _finish_login(base_url: str, api_key: str, config_claude: bool) -> None: show_commands() +def _replace_stored_token(record: CliTokenData, http: Http) -> None: + previous: Final = load_token() + save_token(record) + if previous is None: + return + revocation: Final = revoke_stored_credential(previous, http) + if revocation is not None: + click.echo( + f"Could not revoke the previous login's refresh token on the proxy ({revocation.reason}); " + "it expires on its own." + ) + + def _pkce_login(base_url: str, config_claude: bool) -> None: - credential: Final = run_pkce_login(base_url, requests.Session(), echo=click.echo) + http: Final = requests.Session() + credential: Final = run_pkce_login(base_url, http, echo=click.echo) if isinstance(credential, PkceFailure): click.echo(f"Authentication failed: {credential.reason}") return - save_token(pkce_token_record(base_url, credential)) + _replace_stored_token(pkce_token_record(base_url, credential), http) _finish_login(base_url, credential.access_token, config_claude) @@ -739,7 +754,7 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool) -> None: # Save token data. base_url is stored so we can verify origin # before reusing the key on a subsequent CLI invocation. - save_token( + _replace_stored_token( { "base_url": base_url.rstrip("/"), "key": api_key, @@ -749,7 +764,8 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool) -> None: "auth_header_name": "Authorization", "jwt_token": "", "timestamp": time.time(), - } + }, + requests.Session(), ) _finish_login(base_url, api_key, config_claude) diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 0fc42e7158e..454cf6f8df1 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -371,10 +371,50 @@ class TestTokenUtilities: class TestLoginCommand: """Test login CLI command""" + @pytest.fixture(autouse=True) + def isolated_home(self, monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + return tmp_path + def setup_method(self): """Setup for each test""" self.runner = CliRunner() + def test_login_replaces_a_pkce_record_and_revokes_its_refresh_token(self): + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "status": "ready", + "key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt", + "user_id": "test-user-123", + "team_id": "team-1", + "teams": ["team-1"], + } + _FakeSession.instances.clear() + + with ( + patch("webbrowser.open"), + patch("requests.post", return_value=_mock_cli_sso_start_response()), + patch("requests.get", return_value=mock_response), + patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession), + patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record()), + patch("litellm.proxy.client.cli.commands.auth.save_token") as mock_save, + patch("litellm.proxy.client.cli.interface.show_commands"), + ): + result = self.runner.invoke(login, obj={"base_url": "https://test.example.com"}) + + assert result.exit_code == 0, result.output + assert "Login successful!" in result.output + assert "Could not revoke" not in result.output + assert mock_save.call_args.args[0]["key"] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt" + assert _FakeSession.instances[0].posts == [ + ( + f"{PKCE_BASE_URL}/revoke", + {"token": "llm_srefresh_old", "token_type_hint": "refresh_token", "client_id": "llm_dcrc_abc"}, + ) + ] + def test_login_success(self): """Test successful login flow with single team (JWT generated immediately)""" mock_context = Mock() @@ -1237,10 +1277,83 @@ def _pkce_credential(): class TestPkceLoginCommand: """``lite login --pkce`` swaps the proxy-mediated SSO poll for the browser PKCE flow.""" + @pytest.fixture(autouse=True) + def isolated_home(self, monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + return tmp_path + def setup_method(self): self.runner = CliRunner() _FakeSession.instances.clear() + def test_pkce_login_saves_the_new_record_then_revokes_the_refresh_token_it_replaced(self): + posts_when_saved = [] + + with ( + patch("litellm.proxy.client.cli.commands.auth.run_pkce_login", return_value=_pkce_credential()), + patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record(team_id="team-a")), + patch( + "litellm.proxy.client.cli.commands.auth.save_token", + side_effect=lambda record: posts_when_saved.append(list(_FakeSession.instances[0].posts)), + ) as save, + patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession), + patch("litellm.proxy.client.cli.interface.show_commands"), + ): + result = self.runner.invoke(login, ["--pkce"], obj={"base_url": PKCE_BASE_URL}) + + assert result.exit_code == 0, result.output + assert "Login successful!" in result.output + assert "Could not revoke" not in result.output + assert save.call_args.args[0]["refresh_token"] == "llm_srefresh_fresh" + assert save.call_args.args[0]["team_id"] == "team-b" + assert posts_when_saved == [[]] + assert _FakeSession.instances[0].posts == [ + ( + f"{PKCE_BASE_URL}/revoke", + {"token": "llm_srefresh_old", "token_type_hint": "refresh_token", "client_id": "llm_dcrc_abc"}, + ) + ] + + def test_pkce_login_keeps_the_new_record_when_the_old_refresh_token_cannot_be_revoked(self): + class _FailingSession(_FakeSession): + def __init__(self): + super().__init__() + self.response = _FakeHttpResponse(503, {"error": "temporarily_unavailable"}) + + with ( + patch("litellm.proxy.client.cli.commands.auth.run_pkce_login", return_value=_pkce_credential()), + patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record()), + patch("litellm.proxy.client.cli.commands.auth.save_token") as save, + patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FailingSession), + patch("litellm.proxy.client.cli.interface.show_commands"), + ): + result = self.runner.invoke(login, ["--pkce"], obj={"base_url": PKCE_BASE_URL}) + + assert result.exit_code == 0, result.output + assert ( + "Could not revoke the previous login's refresh token on the proxy (revocation failed with 503" + in result.output + ) + assert "Login successful!" in result.output + assert save.call_args.args[0]["refresh_token"] == "llm_srefresh_fresh" + + @pytest.mark.parametrize("previous", [None, {"key": "sk-classic", "base_url": PKCE_BASE_URL}]) + def test_pkce_login_without_a_previous_refresh_token_makes_no_revocation_request(self, previous): + with ( + patch("litellm.proxy.client.cli.commands.auth.run_pkce_login", return_value=_pkce_credential()), + patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=previous), + patch("litellm.proxy.client.cli.commands.auth.save_token") as save, + patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession), + patch("litellm.proxy.client.cli.interface.show_commands"), + ): + result = self.runner.invoke(login, ["--pkce"], obj={"base_url": PKCE_BASE_URL}) + + assert result.exit_code == 0, result.output + assert "Login successful!" in result.output + save.assert_called_once() + assert _FakeSession.instances[0].posts == [] + def test_pkce_login_saves_the_refreshable_record_and_skips_the_sso_poll(self): with ( patch("litellm.proxy.client.cli.commands.auth.run_pkce_login", return_value=_pkce_credential()) as run,