diff --git a/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py b/api/tests/unit_tests/services/auth/test_api_key_auth_service.py similarity index 70% rename from api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py rename to api/tests/unit_tests/services/auth/test_api_key_auth_service.py index e22aa102328..4bf1913dd5d 100644 --- a/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py +++ b/api/tests/unit_tests/services/auth/test_api_key_auth_service.py @@ -1,3 +1,5 @@ +"""Unit tests for API-key authentication using an SQLite binding table.""" + from __future__ import annotations import json @@ -5,13 +7,15 @@ from unittest.mock import MagicMock, Mock, patch from uuid import uuid4 import pytest -from flask import Flask +from sqlalchemy import event from sqlalchemy.orm import Session from models.source import DataSourceApiKeyAuthBinding from services.auth.api_key_auth_service import ApiKeyAuthService +@pytest.mark.parametrize("sqlite_session", [(DataSourceApiKeyAuthBinding,)], indirect=True) +@pytest.mark.usefixtures("sqlite_session") class TestApiKeyAuthService: @pytest.fixture def tenant_id(self) -> str: @@ -45,36 +49,28 @@ class TestApiKeyAuthService: db_session.commit() return binding - def test_get_provider_auth_list_success( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider - ): - self._create_binding(db_session_with_containers, tenant_id=tenant_id, category=category, provider=provider) - db_session_with_containers.expire_all() + def test_get_provider_auth_list_success(self, sqlite_session: Session, tenant_id, category, provider): + self._create_binding(sqlite_session, tenant_id=tenant_id, category=category, provider=provider) + sqlite_session.expire_all() - result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=sqlite_session) assert len(result) >= 1 tenant_results = [r for r in result if r.tenant_id == tenant_id] assert len(tenant_results) == 1 assert tenant_results[0].provider == provider - def test_get_provider_auth_list_empty( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id - ): - result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) + def test_get_provider_auth_list_empty(self, sqlite_session: Session, tenant_id): + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=sqlite_session) tenant_results = [r for r in result if r.tenant_id == tenant_id] assert tenant_results == [] - def test_get_provider_auth_list_filters_disabled( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider - ): - self._create_binding( - db_session_with_containers, tenant_id=tenant_id, category=category, provider=provider, disabled=True - ) - db_session_with_containers.expire_all() + def test_get_provider_auth_list_filters_disabled(self, sqlite_session: Session, tenant_id, category, provider): + self._create_binding(sqlite_session, tenant_id=tenant_id, category=category, provider=provider, disabled=True) + sqlite_session.expire_all() - result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=sqlite_session) tenant_results = [r for r in result if r.tenant_id == tenant_id] assert tenant_results == [] @@ -85,8 +81,7 @@ class TestApiKeyAuthService: self, mock_encrypter: MagicMock, mock_factory: MagicMock, - flask_app_with_containers: Flask, - db_session_with_containers: Session, + sqlite_session: Session, tenant_id, mock_args, ): @@ -95,22 +90,21 @@ class TestApiKeyAuthService: mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.return_value = "encrypted_test_key_123" - ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=sqlite_session) mock_factory.assert_called_once() mock_auth_instance.validate_credentials.assert_called_once() mock_encrypter.encrypt_token.assert_called_once_with(tenant_id, "test_secret_key_123") - db_session_with_containers.expire_all() - bindings = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id).all() + sqlite_session.expire_all() + bindings = sqlite_session.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id).all() assert len(bindings) == 1 @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") def test_create_provider_auth_validation_failed( self, mock_factory: MagicMock, - flask_app_with_containers: Flask, - db_session_with_containers: Session, + sqlite_session: Session, tenant_id, mock_args, ): @@ -118,10 +112,10 @@ class TestApiKeyAuthService: mock_auth_instance.validate_credentials.return_value = False mock_factory.return_value = mock_auth_instance - ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=sqlite_session) - db_session_with_containers.expire_all() - bindings = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id).all() + sqlite_session.expire_all() + bindings = sqlite_session.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id).all() assert len(bindings) == 0 @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") @@ -130,8 +124,7 @@ class TestApiKeyAuthService: self, mock_encrypter: MagicMock, mock_factory: MagicMock, - flask_app_with_containers: Flask, - db_session_with_containers: Session, + sqlite_session: Session, tenant_id, mock_args, ): @@ -142,7 +135,7 @@ class TestApiKeyAuthService: original_key = mock_args["credentials"]["config"]["api_key"] - ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=sqlite_session) assert mock_args["credentials"]["config"]["api_key"] == "encrypted_test_key_123" assert mock_args["credentials"]["config"]["api_key"] != original_key @@ -150,77 +143,60 @@ class TestApiKeyAuthService: def test_get_auth_credentials_success( self, - flask_app_with_containers: Flask, - db_session_with_containers: Session, + sqlite_session: Session, tenant_id, category, provider, mock_credentials, ): self._create_binding( - db_session_with_containers, + sqlite_session, tenant_id=tenant_id, category=category, provider=provider, credentials=mock_credentials, ) - db_session_with_containers.expire_all() + sqlite_session.expire_all() - result = ApiKeyAuthService.get_auth_credentials( - tenant_id, category, provider, session=db_session_with_containers - ) + result = ApiKeyAuthService.get_auth_credentials(tenant_id, category, provider, session=sqlite_session) assert result == mock_credentials - def test_get_auth_credentials_not_found( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider - ): - result = ApiKeyAuthService.get_auth_credentials( - tenant_id, category, provider, session=db_session_with_containers - ) + def test_get_auth_credentials_not_found(self, sqlite_session: Session, tenant_id, category, provider): + result = ApiKeyAuthService.get_auth_credentials(tenant_id, category, provider, session=sqlite_session) assert result is None - def test_get_auth_credentials_json_parsing( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider - ): + def test_get_auth_credentials_json_parsing(self, sqlite_session: Session, tenant_id, category, provider): special_credentials = {"auth_type": "api_key", "config": {"api_key": "key_with_中文_and_special_chars_!@#$%"}} self._create_binding( - db_session_with_containers, + sqlite_session, tenant_id=tenant_id, category=category, provider=provider, credentials=special_credentials, ) - db_session_with_containers.expire_all() + sqlite_session.expire_all() - result = ApiKeyAuthService.get_auth_credentials( - tenant_id, category, provider, session=db_session_with_containers - ) + result = ApiKeyAuthService.get_auth_credentials(tenant_id, category, provider, session=sqlite_session) assert result == special_credentials assert result["config"]["api_key"] == "key_with_中文_and_special_chars_!@#$%" - def test_delete_provider_auth_success( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider - ): - binding = self._create_binding( - db_session_with_containers, tenant_id=tenant_id, category=category, provider=provider - ) + def test_delete_provider_auth_success(self, sqlite_session: Session, tenant_id, category, provider): + binding = self._create_binding(sqlite_session, tenant_id=tenant_id, category=category, provider=provider) binding_id = binding.id - db_session_with_containers.expire_all() + sqlite_session.expire_all() - ApiKeyAuthService.delete_provider_auth(tenant_id, binding_id, session=db_session_with_containers) + ApiKeyAuthService.delete_provider_auth(tenant_id, binding_id, session=sqlite_session) - db_session_with_containers.expire_all() - remaining = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(id=binding_id).first() + sqlite_session.expire_all() + remaining = sqlite_session.query(DataSourceApiKeyAuthBinding).filter_by(id=binding_id).first() assert remaining is None - def test_delete_provider_auth_not_found( - self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id - ): + def test_delete_provider_auth_not_found(self, sqlite_session: Session, tenant_id): # Should not raise when binding not found - ApiKeyAuthService.delete_provider_auth(tenant_id, str(uuid4()), session=db_session_with_containers) + ApiKeyAuthService.delete_provider_auth(tenant_id, str(uuid4()), session=sqlite_session) def test_validate_api_key_auth_args_success(self, mock_args): ApiKeyAuthService.validate_api_key_auth_args(mock_args) @@ -287,33 +263,42 @@ class TestApiKeyAuthService: @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") @patch("services.auth.api_key_auth_service.encrypter") def test_create_provider_auth_database_error_handling( - self, mock_encrypter, mock_factory, flask_app_with_containers: Flask, tenant_id, mock_args - ): + self, mock_encrypter, mock_factory, tenant_id, mock_args, sqlite_session: Session + ) -> None: mock_auth_instance = Mock() mock_auth_instance.validate_credentials.return_value = True mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.return_value = "encrypted_key" - mock_session = MagicMock() - mock_session.commit.side_effect = Exception("Database error") - with pytest.raises(Exception, match="Database error"): - ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=mock_session) + def raise_database_error(_session: Session) -> None: + raise Exception("Database error") + + event.listen(sqlite_session, "before_commit", raise_database_error) + try: + with pytest.raises(Exception, match="Database error"): + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=sqlite_session) + finally: + event.remove(sqlite_session, "before_commit", raise_database_error) @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") - def test_create_provider_auth_factory_exception(self, mock_factory: MagicMock, tenant_id, mock_args): + def test_create_provider_auth_factory_exception( + self, mock_factory: MagicMock, tenant_id, mock_args, sqlite_session: Session + ) -> None: mock_factory.side_effect = Exception("Factory error") with pytest.raises(Exception, match="Factory error"): - ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=MagicMock()) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=sqlite_session) @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") @patch("services.auth.api_key_auth_service.encrypter") - def test_create_provider_auth_encryption_exception(self, mock_encrypter, mock_factory, tenant_id, mock_args): + def test_create_provider_auth_encryption_exception( + self, mock_encrypter, mock_factory, tenant_id, mock_args, sqlite_session: Session + ) -> None: mock_auth_instance = Mock() mock_auth_instance.validate_credentials.return_value = True mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.side_effect = Exception("Encryption error") with pytest.raises(Exception, match="Encryption error"): - ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=MagicMock()) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=sqlite_session) def test_validate_api_key_auth_args_none_input(self): with pytest.raises(TypeError):