| """ |
| OAuth flow security tests (SECU-06). |
| |
| Tests cover: |
| - GitHub OAuth flow |
| - Google OAuth flow |
| - Microsoft OAuth flow |
| - State parameter CSRF prevention |
| - Token encryption at rest |
| - Token refresh |
| """ |
| import pytest |
| from unittest.mock import Mock, patch |
| from urllib.parse import urlparse, parse_qs |
| from fastapi.testclient import TestClient |
| from sqlalchemy.orm import Session |
| from tests.factories.user_factory import UserFactory |
| from core.models import OAuthToken |
| import json |
|
|
|
|
| class TestGitHubOAuthFlow: |
| """Test GitHub OAuth integration.""" |
|
|
| @patch('core.oauth_handler.OAuthHandler.get_authorization_url') |
| def test_github_oauth_authorize_redirect(self, mock_get_url, client: TestClient): |
| """Test GitHub OAuth initiate redirect.""" |
| mock_get_url.return_value = "https://github.com/login/oauth/authorize?client_id=test&state=test123" |
|
|
| response = client.get("/api/v1/integrations/github/authorize") |
|
|
| |
| assert response.status_code in [200, 404, 501] |
|
|
| if response.status_code == 200: |
| data = response.json() |
| |
| if "auth_url" in data: |
| auth_url = data["auth_url"] |
| |
| assert "github.com" in auth_url |
| assert "client_id" in auth_url |
| |
| assert "state" in auth_url or "state" in data |
|
|
| @patch('httpx.AsyncClient.post') |
| @patch('httpx.AsyncClient.get') |
| def test_github_oauth_callback_with_valid_code(self, mock_get, mock_post, client: TestClient, db_session: Session): |
| """Test GitHub OAuth callback with valid authorization code.""" |
| |
| mock_post_response = Mock() |
| mock_post_response.status_code = 200 |
| mock_post_response.json.return_value = { |
| "access_token": "github_access_token_123", |
| "token_type": "bearer", |
| "scope": "user:email" |
| } |
| mock_post.return_value = mock_post_response |
|
|
| |
| mock_get_response = Mock() |
| mock_get_response.status_code = 200 |
| mock_get_response.json.return_value = { |
| "id": 123456, |
| "login": "testuser", |
| "email": "test@example.com", |
| "name": "Test User" |
| } |
| mock_get.return_value = mock_get_response |
|
|
| response = client.get( |
| "/api/v1/integrations/github/callback?code=test_code_123&state=test_state_123" |
| ) |
|
|
| |
| assert response.status_code in [200, 302, 404, 501] |
|
|
| |
| if response.status_code in [200, 302]: |
| oauth_token = db_session.query(OAuthToken).filter( |
| OAuthToken.provider == "github" |
| ).first() |
| |
| if oauth_token: |
| assert oauth_token is not None |
|
|
| @patch('httpx.AsyncClient.post') |
| def test_github_oauth_state_parameter_validation(self, mock_post, client: TestClient): |
| """Test GitHub OAuth validates state parameter (CSRF prevention).""" |
| |
| mock_post_response = Mock() |
| mock_post_response.status_code = 200 |
| mock_post_response.json.return_value = {"access_token": "test_token"} |
| mock_post.return_value = mock_post_response |
|
|
| |
| response = client.get( |
| "/api/v1/integrations/github/callback?code=test_code&state=malicious_state" |
| ) |
|
|
| |
| if response.status_code in [400, 403, 401]: |
| |
| try: |
| detail = response.json().get("detail", "").lower() |
| assert "state" in detail or "csrf" in detail |
| except: |
| pass |
| else: |
| |
| assert response.status_code in [404, 501] |
|
|
| def test_github_oauth_error_handling(self, client: TestClient): |
| """Test GitHub OAuth error handling.""" |
| response = client.get( |
| "/api/v1/integrations/github/callback?error=access_denied&state=test_state" |
| ) |
|
|
| |
| assert response.status_code in [200, 302, 400, 401, 404, 501] |
|
|
|
|
| class TestGoogleOAuthFlow: |
| """Test Google OAuth integration.""" |
|
|
| @patch('core.oauth_handler.OAuthHandler.get_authorization_url') |
| def test_google_oauth_authorize_redirect(self, mock_get_url, client: TestClient): |
| """Test Google OAuth initiate redirect.""" |
| response = client.get("/api/v1/integrations/google/authorize") |
|
|
| |
| assert response.status_code in [200, 404, 501] |
|
|
| if response.status_code == 200: |
| data = response.json() |
| if "auth_url" in data: |
| auth_url = data["auth_url"] |
| assert "accounts.google.com" in auth_url or "googleapis.com" in auth_url |
| assert "state" in auth_url or "state" in data |
|
|
| @patch('httpx.AsyncClient.post') |
| @patch('httpx.AsyncClient.get') |
| def test_google_oauth_callback_success(self, mock_get, mock_post, client: TestClient, db_session: Session): |
| """Test Google OAuth callback with valid authorization.""" |
| |
| mock_post_response = Mock() |
| mock_post_response.status_code = 200 |
| mock_post_response.json.return_value = { |
| "access_token": "google_access_token", |
| "expires_in": 3600, |
| "refresh_token": "google_refresh_token", |
| "token_type": "Bearer" |
| } |
| mock_post.return_value = mock_post_response |
|
|
| |
| mock_get_response = Mock() |
| mock_get_response.status_code = 200 |
| mock_get_response.json.return_value = { |
| "id": "123456789", |
| "email": "test@gmail.com", |
| "verified_email": True, |
| "name": "Test User", |
| "given_name": "Test", |
| "family_name": "User" |
| } |
| mock_get.return_value = mock_get_response |
|
|
| response = client.get( |
| "/api/v1/integrations/google/callback?code=test_code&state=test_state" |
| ) |
|
|
| assert response.status_code in [200, 302, 404, 501] |
|
|
| @patch('core.oauth_handler.OAuthHandler.get_authorization_url') |
| def test_google_oauth_state_csrf_protection(self, mock_get_url, client: TestClient): |
| """Test Google OAuth state parameter prevents CSRF.""" |
| response = client.get("/api/v1/integrations/google/authorize") |
|
|
| |
| response = client.get( |
| "/api/v1/integrations/google/callback?code=test_code&state=attacker_state" |
| ) |
|
|
| |
| assert response.status_code in [200, 302, 400, 403, 404, 501] |
|
|
|
|
| class TestMicrosoftOAuthFlow: |
| """Test Microsoft OAuth integration.""" |
|
|
| @patch('httpx.AsyncClient.post') |
| @patch('httpx.AsyncClient.get') |
| def test_microsoft_oauth_callback_success(self, mock_get, mock_post, client: TestClient, db_session: Session): |
| """Test Microsoft OAuth callback with valid authorization.""" |
| |
| mock_post_response = Mock() |
| mock_post_response.status_code = 200 |
| mock_post_response.json.return_value = { |
| "access_token": "microsoft_access_token", |
| "expires_in": 3600, |
| "refresh_token": "microsoft_refresh_token", |
| "token_type": "Bearer" |
| } |
| mock_post.return_value = mock_post_response |
|
|
| |
| mock_get_response = Mock() |
| mock_get_response.status_code = 200 |
| mock_get_response.json.return_value = { |
| "id": "microsoft_id_123", |
| "mail": "test@outlook.com", |
| "displayName": "Test User" |
| } |
| mock_get.return_value = mock_get_response |
|
|
| response = client.get( |
| "/api/v1/integrations/microsoft/callback?code=test_code&state=test_state" |
| ) |
|
|
| assert response.status_code in [200, 302, 404, 501] |
|
|
|
|
| class TestTokenEncryption: |
| """Test OAuth token encryption at rest.""" |
|
|
| def test_tokens_encrypted_in_database(self, client: TestClient, db_session: Session): |
| """Test OAuth tokens are encrypted, not stored as plaintext.""" |
| |
| user = UserFactory(_session=db_session) |
|
|
| |
| oauth_token = OAuthToken( |
| user_id=user.id, |
| provider="github", |
| _encrypted_access_token="encrypted_token_123", |
| access_token="github_access_token_123" |
| ) |
| db_session.add(oauth_token) |
| db_session.commit() |
|
|
| |
| stored_token = db_session.query(OAuthToken).filter( |
| OAuthToken.provider == "github" |
| ).first() |
|
|
| assert stored_token is not None |
| |
| if hasattr(stored_token, '_encrypted_access_token'): |
| assert stored_token._encrypted_access_token != "github_access_token_123" |
| |
| assert stored_token.access_token == "github_access_token_123" |
|
|
| def test_token_property_decrypts_value(self, db_session: Session): |
| """Test token property returns decrypted value.""" |
| user = UserFactory(_session=db_session) |
|
|
| |
| oauth_token = OAuthToken( |
| user_id=user.id, |
| provider="test_provider", |
| _encrypted_access_token="encrypted_value_xyz", |
| access_token="decrypted_value_xyz" |
| ) |
| db_session.add(oauth_token) |
| db_session.commit() |
|
|
| |
| assert oauth_token.access_token == "decrypted_value_xyz" |
|
|
|
|
| class TestTokenRefresh: |
| """Test OAuth token refresh flow.""" |
|
|
| @patch('httpx.AsyncClient.post') |
| def test_token_refresh_with_valid_refresh_token(self, mock_post, client: TestClient, db_session: Session): |
| """Test refreshing OAuth token with valid refresh token.""" |
| |
| user = UserFactory(email="oauth@test.com", _session=db_session) |
| oauth_token = OAuthToken( |
| user_id=user.id, |
| provider="github", |
| _encrypted_access_token="old_encrypted_token", |
| refresh_token="valid_refresh_token", |
| expires_at="2026-02-01T10:00:00" |
| ) |
| db_session.add_all([user, oauth_token]) |
| db_session.commit() |
|
|
| |
| mock_post_response = Mock() |
| mock_post_response.status_code = 200 |
| mock_post_response.json.return_value = { |
| "access_token": "new_access_token", |
| "refresh_token": "new_refresh_token", |
| "expires_in": 7200 |
| } |
| mock_post.return_value = mock_post_response |
|
|
| |
| from tests.security.conftest import create_test_token |
| response = client.post( |
| "/api/v1/integrations/oauth/refresh", |
| json={"provider": "github"}, |
| headers={"Authorization": f"Bearer {create_test_token(user.id)}"} |
| ) |
|
|
| |
| assert response.status_code in [200, 404, 501] |
|
|
| if response.status_code == 200: |
| data = response.json() |
| assert "access_token" in data or "success" in data |
|
|
| @patch('httpx.AsyncClient.post') |
| def test_token_refresh_with_invalid_refresh_token(self, mock_post, client: TestClient, db_session: Session): |
| """Test token refresh fails with invalid refresh token.""" |
| user = UserFactory(email="oauth@test.com", _session=db_session) |
| oauth_token = OAuthToken( |
| user_id=user.id, |
| provider="github", |
| _encrypted_access_token="old_token", |
| refresh_token="invalid_refresh_token", |
| expires_at="2026-02-01T10:00:00" |
| ) |
| db_session.add_all([user, oauth_token]) |
| db_session.commit() |
|
|
| |
| mock_post_response = Mock() |
| mock_post_response.status_code = 200 |
| mock_post_response.json.return_value = { |
| "error": "invalid_grant", |
| "error_description": "The refresh token is invalid." |
| } |
| mock_post.return_value = mock_post_response |
|
|
| from tests.security.conftest import create_test_token |
| response = client.post( |
| "/api/v1/integrations/oauth/refresh", |
| json={"provider": "github"}, |
| headers={"Authorization": f"Bearer {create_test_token(user.id)}"} |
| ) |
|
|
| |
| assert response.status_code in [200, 400, 401, 404, 501] |
|
|
|
|
| class TestTokenRevocation: |
| """Test OAuth token revocation.""" |
|
|
| def test_revoke_oauth_token(self, client: TestClient, db_session: Session): |
| """Test revoking OAuth token.""" |
| user = UserFactory(email="oauth@test.com", _session=db_session) |
| oauth_token = OAuthToken( |
| user_id=user.id, |
| provider="github", |
| _encrypted_access_token="active_token", |
| refresh_token="refresh_token" |
| ) |
| db_session.add_all([user, oauth_token]) |
| db_session.commit() |
|
|
| from tests.security.conftest import create_test_token |
| response = client.post( |
| f"/api/v1/integrations/oauth/revoke/{oauth_token.provider}", |
| headers={"Authorization": f"Bearer {create_test_token(user.id)}"} |
| ) |
|
|
| |
| assert response.status_code in [200, 204, 404, 501] |
|
|
| if response.status_code in [200, 204]: |
| |
| db_session.refresh(oauth_token) |
| assert oauth_token.revoked_at is not None or \ |
| db_session.query(OAuthToken).filter(OAuthToken.id == oauth_token.id).first() is None |
|
|
|
|
| class TestOAuthStateParameterSecurity: |
| """Test OAuth state parameter security across providers.""" |
|
|
| def test_state_parameter_is_unpredictable(self, client: TestClient): |
| """Test state parameter is cryptographically random.""" |
| import secrets |
|
|
| |
| states = [secrets.token_urlsafe(32) for _ in range(10)] |
|
|
| |
| assert len(set(states)) == 10 |
|
|
| |
| for state in states: |
| |
| |
| assert len(state) >= 32 |
|
|
| def test_state_parameter_prevents_csuf(self): |
| """Test state parameter prevents CSRF attacks in theory.""" |
| |
| |
| |
| |
| |
|
|
| |
| import secrets |
| server_state = secrets.token_urlsafe(32) |
|
|
| |
| attacker_state = "malicious_state_123" |
|
|
| |
| assert server_state != attacker_state |
|
|
| |
| callback_state = server_state |
| assert callback_state == server_state |
|
|