diff --git a/api/callback.py b/api/callback.py index dbffc54..c76f8ca 100644 --- a/api/callback.py +++ b/api/callback.py @@ -1,18 +1,17 @@ -from base64 import b64decode - -from dotenv import find_dotenv, load_dotenv -from flask import Flask, Response, jsonify, redirect, render_template, request - -load_dotenv(find_dotenv()) - +import hmac import json import os +from base64 import b64decode import firebase_admin +from dotenv import find_dotenv, load_dotenv from firebase_admin import credentials, firestore +from flask import Flask, Response, make_response, render_template, request from util import spotify +load_dotenv(find_dotenv()) + print("Starting Server") firebase_config = os.getenv("FIREBASE") @@ -27,30 +26,46 @@ app = Flask(__name__) +def _clear_oauth_state_cookie(response): + response.delete_cookie(spotify.OAUTH_STATE_COOKIE_NAME, path="/") + return response + + @app.route("/", defaults={"path": ""}) @app.route("/") def catch_all(path): code = request.args.get("code") if code is None: - # TODO: no code return Response("not ok") + state = request.args.get("state") + expected_state = request.cookies.get(spotify.OAUTH_STATE_COOKIE_NAME) + if ( + not state + or not expected_state + or not hmac.compare_digest(state, expected_state) + ): + return Response("Invalid OAuth state", status=400) + token_info = spotify.generate_token(code) if "access_token" not in token_info: error = token_info.get("error", "unknown") desc = token_info.get("error_description", "") - return Response(f"Token exchange failed: {error} - {desc}", status=400) + response = Response(f"Token exchange failed: {error} - {desc}", status=400) + return _clear_oauth_state_cookie(response) + token_info = spotify.normalize_token_info(token_info) access_token = token_info["access_token"] profile_resp = spotify.get_user_profile_raw(access_token) if profile_resp.status_code != 200 or not profile_resp.text.strip(): - return Response( + response = Response( f"Spotify profile fetch failed: HTTP {profile_resp.status_code} - {profile_resp.text[:300]}", status=502, ) + return _clear_oauth_state_cookie(response) spotify_user = profile_resp.json() user_id = spotify_user["id"] @@ -63,7 +78,8 @@ def catch_all(path): "BASE_URL": spotify.BASE_URL, } - return render_template("callback.html.j2", **rendered_data) + response = make_response(render_template("callback.html.j2", **rendered_data)) + return _clear_oauth_state_cookie(response) if __name__ == "__main__": diff --git a/api/login.py b/api/login.py index f99bc23..b49e385 100644 --- a/api/login.py +++ b/api/login.py @@ -1,4 +1,6 @@ -from flask import Flask, Response, jsonify, render_template, redirect +import secrets + +from flask import Flask, redirect from util import spotify @@ -8,10 +10,27 @@ @app.route("/", defaults={"path": ""}) @app.route("/") def catch_all(path): - - login_url = f"https://accounts.spotify.com/authorize?client_id={spotify.SPOTIFY_CLIENT_ID}&response_type=code&scope=user-read-currently-playing,user-read-recently-played&redirect_uri={spotify.REDIRECT_URI}" - - return redirect(login_url) + state = secrets.token_urlsafe(32) + login_url = ( + "https://accounts.spotify.com/authorize" + f"?client_id={spotify.SPOTIFY_CLIENT_ID}" + "&response_type=code" + "&scope=user-read-currently-playing,user-read-recently-played" + f"&redirect_uri={spotify.REDIRECT_URI}" + f"&state={state}" + ) + + response = redirect(login_url) + response.set_cookie( + spotify.OAUTH_STATE_COOKIE_NAME, + state, + max_age=spotify.OAUTH_STATE_MAX_AGE_SECONDS, + httponly=True, + secure=bool(spotify.REDIRECT_URI and spotify.REDIRECT_URI.startswith("https://")), + samesite="Lax", + path="/", + ) + return response if __name__ == "__main__": diff --git a/api/view.py b/api/view.py index 37bfd98..c6f6a13 100644 --- a/api/view.py +++ b/api/view.py @@ -254,46 +254,50 @@ def get_access_token(uid): print("not exist data in firebase: {}".format(uid)) return None - token_info = doc.to_dict() - + token_info = doc.to_dict() or {} CACHE_TOKEN_INFO[uid] = token_info current_ts = int(time()) - access_token = token_info.get("access_token", None) - print(access_token) - - # Check token expired + access_token = token_info.get("access_token") expired_ts = token_info.get("expired_ts") - if expired_ts is None or current_ts >= expired_ts: - # Refresh token - refresh_token = token_info["refresh_token"] - new_token = spotify.refresh_token(refresh_token) + # Reuse a valid access token without exposing it in application logs. + if access_token and expired_ts is not None and current_ts < expired_ts: + return access_token - # Handle refresh token revoke - if new_token.get("error") == "invalid_grant": - # Delete token in firebase - doc_ref = db.collection("users").document(uid) - doc_ref.delete() + refresh_token_value = token_info.get("refresh_token") + if not refresh_token_value: + return None - # Delete token in memory cache - delete_cache_token_info(uid) - return None + new_token = spotify.refresh_token(refresh_token_value) - expired_ts = int(time()) + new_token["expires_in"] - update_data = { - "access_token": new_token["access_token"], - "expired_ts": expired_ts, - } + # A revoked refresh token requires a fresh authorization flow. + if new_token.get("error") == "invalid_grant": doc_ref = db.collection("users").document(uid) - doc_ref.update(update_data) + doc_ref.delete() + delete_cache_token_info(uid) + return None + + if "access_token" not in new_token or "expires_in" not in new_token: + return None - access_token = new_token["access_token"] + refreshed_token_info = spotify.normalize_token_info( + new_token, + existing_refresh_token=refresh_token_value, + now=current_ts, + ) + update_data = { + key: refreshed_token_info[key] + for key in ("access_token", "refresh_token", "expires_in", "expired_ts") + if key in refreshed_token_info + } - # Save in memory cache - CACHE_TOKEN_INFO[uid] = update_data + doc_ref = db.collection("users").document(uid) + doc_ref.update(update_data) - return access_token + merged_token_info = {**token_info, **update_data} + CACHE_TOKEN_INFO[uid] = merged_token_info + return merged_token_info["access_token"] def get_song_info(uid, show_offline): diff --git a/tests/test_api_callback.py b/tests/test_api_callback.py index b30d8da..a3dc452 100644 --- a/tests/test_api_callback.py +++ b/tests/test_api_callback.py @@ -48,7 +48,7 @@ def test_callback_with_code(client): def test_callback_with_empty_code_param(client): - """Test callback with empty code parameter.""" + """Test callback with empty authorization code parameter.""" response = client.get("/?code=") # Empty string should be processed (though would fail in real Spotify auth) @@ -97,7 +97,7 @@ def test_multiple_query_parameters(client): "UPPERCASE_CODE" ]) def test_special_code_values(client, code): - """Test callback with special authorization code values.""" + """Test callback with special code values.""" response = client.get(f"/?code={code}") assert response.status_code == 200 @@ -200,11 +200,16 @@ def test_successful_integration_callback(self, mock_generate_token, mock_get_use mock_db.collection.return_value = mock_collection mock_collection.document.return_value = mock_document - # Make request with authorization code - response = real_app_client.get("/?code=test_auth_code") + # Make request with a matching OAuth state cookie/query pair. + real_app_client.set_cookie("spotify_oauth_state", "test_state") + with patch('util.spotify.time', return_value=1000): + response = real_app_client.get( + "/?code=test_auth_code&state=test_state" + ) # Verify response assert response.status_code == 200 + assert "spotify_oauth_state=;" in response.headers.get("Set-Cookie", "") # Verify Spotify API calls mock_generate_token.assert_called_once_with("test_auth_code") @@ -216,8 +221,20 @@ def test_successful_integration_callback(self, mock_generate_token, mock_get_use mock_document.set.assert_called_once_with({ "access_token": "test_access_token", "refresh_token": "test_refresh_token", - "expires_in": 3600 + "expires_in": 3600, + "expired_ts": 4600 }) + + def test_integration_callback_rejects_missing_or_mismatched_state(self, real_app_client): + """OAuth callback must be bound to the browser that started the flow.""" + response = real_app_client.get("/?code=test_code&state=test_state") + assert response.status_code == 400 + assert response.data == b"Invalid OAuth state" + + real_app_client.set_cookie("spotify_oauth_state", "expected_state") + response = real_app_client.get("/?code=test_code&state=wrong_state") + assert response.status_code == 400 + assert response.data == b"Invalid OAuth state" def test_integration_callback_without_code(self, real_app_client): """Test integration callback without authorization code.""" @@ -231,6 +248,7 @@ def test_integration_spotify_error_handling(self, mock_generate_token, real_app_ """Test integration error handling for Spotify API failures.""" # Mock token generation to raise an exception mock_generate_token.side_effect = Exception("Token generation failed") + real_app_client.set_cookie("spotify_oauth_state", "test_state") with pytest.raises(Exception, match="Token generation failed"): - real_app_client.get("/?code=test_code") + real_app_client.get("/?code=test_code&state=test_state") diff --git a/tests/test_auth_security.py b/tests/test_auth_security.py new file mode 100644 index 0000000..4146e7a --- /dev/null +++ b/tests/test_auth_security.py @@ -0,0 +1,170 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +# Add the repository root to the import path. +sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..')) + + +def test_login_generates_state_and_sets_hardened_cookie(): + """The authorization request must carry a state bound to an HttpOnly cookie.""" + from api.login import app + + with patch('api.login.secrets.token_urlsafe', return_value='test_state'), \ + patch('util.spotify.SPOTIFY_CLIENT_ID', 'test_client_id'), \ + patch('util.spotify.REDIRECT_URI', 'https://example.com/api/callback'): + with app.test_client() as client: + response = client.get('/') + + assert response.status_code == 302 + assert 'state=test_state' in response.headers['Location'] + + cookie = response.headers.get('Set-Cookie', '') + assert 'spotify_oauth_state=test_state' in cookie + assert 'Max-Age=600' in cookie + assert 'HttpOnly' in cookie + assert 'Secure' in cookie + assert 'SameSite=Lax' in cookie + + +def test_normalize_token_info_persists_expiry_and_existing_refresh_token(): + """Refresh responses that omit refresh_token must keep the previous token.""" + from util import spotify + + normalized = spotify.normalize_token_info( + { + 'access_token': 'new_access_token', + 'expires_in': '3600', + }, + existing_refresh_token='existing_refresh_token', + now=1000, + ) + + assert normalized == { + 'access_token': 'new_access_token', + 'expires_in': 3600, + 'refresh_token': 'existing_refresh_token', + 'expired_ts': 4600, + } + + +def _load_view_module(): + # Keep this test independent from the CI environment and from test order. + with patch.dict(os.environ, {'TESTING': 'true'}): + from api import view + return view + + +def _mock_token_document(token_info): + document = MagicMock() + snapshot = MagicMock() + snapshot.exists = True + snapshot.to_dict.return_value = token_info + document.get.return_value = snapshot + + collection = MagicMock() + collection.document.return_value = document + + database = MagicMock() + database.collection.return_value = collection + return database, document + + +def test_get_access_token_does_not_log_bearer_token(capsys): + """Bearer tokens are credentials and must never be printed to application logs.""" + view = _load_view_module() + view.CACHE_TOKEN_INFO.clear() + view.CACHE_TOKEN_INFO['test_uid'] = { + 'access_token': 'super_secret_access_token', + 'refresh_token': 'refresh_token', + 'expired_ts': 4600, + } + + with patch.object(view, 'time', return_value=1000): + assert view.get_access_token('test_uid') == 'super_secret_access_token' + + captured = capsys.readouterr() + assert 'super_secret_access_token' not in captured.out + assert 'super_secret_access_token' not in captured.err + + +def test_refresh_keeps_existing_refresh_token_when_spotify_omits_rotation(): + """A refresh response may omit refresh_token; the stored token must survive.""" + view = _load_view_module() + view.CACHE_TOKEN_INFO.clear() + + database, document = _mock_token_document({ + 'access_token': 'expired_access_token', + 'refresh_token': 'existing_refresh_token', + 'expires_in': 3600, + 'expired_ts': 900, + }) + + with patch.object(view, 'db', database), \ + patch.object(view, 'time', return_value=1000), \ + patch.object( + view.spotify, + 'refresh_token', + return_value={ + 'access_token': 'new_access_token', + 'expires_in': 3600, + }, + ): + access_token = view.get_access_token('test_uid') + + assert access_token == 'new_access_token' + document.update.assert_called_once_with({ + 'access_token': 'new_access_token', + 'refresh_token': 'existing_refresh_token', + 'expires_in': 3600, + 'expired_ts': 4600, + }) + assert view.CACHE_TOKEN_INFO['test_uid']['refresh_token'] == 'existing_refresh_token' + + +def test_refresh_persists_rotated_refresh_token(): + """If Spotify rotates the refresh token, the new value must replace the old one.""" + view = _load_view_module() + view.CACHE_TOKEN_INFO.clear() + + database, document = _mock_token_document({ + 'access_token': 'expired_access_token', + 'refresh_token': 'old_refresh_token', + 'expires_in': 3600, + 'expired_ts': 900, + }) + + with patch.object(view, 'db', database), \ + patch.object(view, 'time', return_value=1000), \ + patch.object( + view.spotify, + 'refresh_token', + return_value={ + 'access_token': 'new_access_token', + 'refresh_token': 'rotated_refresh_token', + 'expires_in': 1800, + }, + ): + access_token = view.get_access_token('test_uid') + + assert access_token == 'new_access_token' + document.update.assert_called_once_with({ + 'access_token': 'new_access_token', + 'refresh_token': 'rotated_refresh_token', + 'expires_in': 1800, + 'expired_ts': 2800, + }) + assert view.CACHE_TOKEN_INFO['test_uid']['refresh_token'] == 'rotated_refresh_token' + + +def test_normalize_token_info_rejects_invalid_expiry(): + """Invalid token metadata should fail clearly instead of poisoning the cache.""" + from util import spotify + + with pytest.raises(ValueError, match='expires_in'): + spotify.normalize_token_info( + {'access_token': 'token', 'expires_in': 'not-a-number'}, + now=1000, + ) diff --git a/util/spotify.py b/util/spotify.py index 8ce2900..c06aa18 100644 --- a/util/spotify.py +++ b/util/spotify.py @@ -1,12 +1,11 @@ from base64 import b64encode +from time import time from dotenv import find_dotenv, load_dotenv load_dotenv(find_dotenv()) -import json import os -import random import requests @@ -16,6 +15,9 @@ REDIRECT_URI = "{}/callback".format(BASE_URL) +OAUTH_STATE_COOKIE_NAME = "spotify_oauth_state" +OAUTH_STATE_MAX_AGE_SECONDS = 600 + # scope user-read-currently-playing,user-read-recently-played SPOTIFY_URL_REFRESH_TOKEN = "https://accounts.spotify.com/api/token" SPOTIFY_URL_NOW_PLAYING = "https://api.spotify.com/v1/me/player/currently-playing?additional_types=track,episode" @@ -31,6 +33,33 @@ class InvalidTokenError(Exception): pass +def normalize_token_info(token_info, existing_refresh_token=None, now=None): + """Normalize Spotify token metadata before persisting or caching it. + + Spotify may omit ``refresh_token`` from a refresh response. In that case, + the previous refresh token must be retained. ``expired_ts`` is stored as an + absolute timestamp so callers do not need to refresh a freshly-issued token + on its first use. + """ + normalized = dict(token_info) + + if existing_refresh_token and not normalized.get("refresh_token"): + normalized["refresh_token"] = existing_refresh_token + + expires_in = normalized.get("expires_in") + if expires_in is not None: + try: + expires_in = max(0, int(expires_in)) + except (TypeError, ValueError) as exc: + raise ValueError("Spotify token expires_in must be an integer") from exc + + normalized["expires_in"] = expires_in + current_ts = int(time() if now is None else now) + normalized["expired_ts"] = current_ts + expires_in + + return normalized + + def get_authorization(): return b64encode(f"{SPOTIFY_CLIENT_ID}:{SPOTIFY_SECRET_ID}".encode()).decode(