"""Tests for authentication middleware.""" import pytest from unittest.mock import patch, AsyncMock, MagicMock from fastapi import HTTPException from fastapi.security import HTTPAuthorizationCredentials from app.middleware.auth import get_current_client, AuthMiddleware from app.models import ClientKey class TestGetCurrentClient: """Test get_current_client dependency.""" @pytest.mark.asyncio async def test_missing_credentials(self): """Test with missing credentials.""" with pytest.raises(HTTPException) as exc_info: await get_current_client(None) assert exc_info.value.status_code == 401 assert "Missing API key" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_invalid_key(self, mock_db): """Test with invalid API key.""" credentials = HTTPAuthorizationCredentials( scheme="Bearer", credentials="invalid-key" ) with patch("app.middleware.auth.AuthService.validate_key", return_value=None): with pytest.raises(HTTPException) as exc_info: await get_current_client(credentials) assert exc_info.value.status_code == 401 assert "Invalid API key" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_disabled_key(self): """Test with disabled API key.""" credentials = HTTPAuthorizationCredentials( scheme="Bearer", credentials="disabled-key" ) disabled_client = ClientKey( id=1, key="disabled-key", name="disabled", is_active=False, created_at="2024-01-01" ) with patch("app.middleware.auth.AuthService.validate_key", return_value=disabled_client): with pytest.raises(HTTPException) as exc_info: await get_current_client(credentials) assert exc_info.value.status_code == 403 assert "disabled" in str(exc_info.value.detail).lower() @pytest.mark.asyncio async def test_valid_key(self): """Test with valid API key.""" credentials = HTTPAuthorizationCredentials( scheme="Bearer", credentials="test-api-key" ) valid_client = ClientKey( id=1, key="test-api-key", name="test-client", is_active=True, created_at="2024-01-01" ) with patch("app.middleware.auth.AuthService.validate_key", return_value=valid_client): result = await get_current_client(credentials) assert result == valid_client assert result.is_active is True class TestAuthMiddleware: """Test AuthMiddleware.""" def test_default_exclude_paths(self): """Test default exclude paths.""" middleware = AuthMiddleware() assert "/health" in middleware.exclude_paths assert "/docs" in middleware.exclude_paths assert "/openapi.json" in middleware.exclude_paths def test_custom_exclude_paths(self): """Test custom exclude paths.""" middleware = AuthMiddleware(exclude_paths=["/custom"]) assert "/custom" in middleware.exclude_paths assert "/health" not in middleware.exclude_paths @pytest.mark.asyncio async def test_middleware_calls_next(self): """Test middleware calls next handler.""" middleware = AuthMiddleware() request = MagicMock() call_next = AsyncMock(return_value=MagicMock()) response = await middleware(request, call_next) call_next.assert_called_once_with(request)