diff --git a/api/tests/unit_tests/controllers/service_api/test_wraps.py b/api/tests/unit_tests/controllers/service_api/test_wraps.py index 5857e5c639a..f5b8c15d6d1 100644 --- a/api/tests/unit_tests/controllers/service_api/test_wraps.py +++ b/api/tests/unit_tests/controllers/service_api/test_wraps.py @@ -3,10 +3,13 @@ Unit tests for Service API wraps (authentication decorators) """ import uuid +from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask +from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound, Unauthorized from controllers.service_api.wraps import ( @@ -21,12 +24,11 @@ from controllers.service_api.wraps import ( validate_dataset_token, ) from enums.cloud_plan import CloudPlan -from models.account import TenantStatus -from models.model import ApiToken -from tests.unit_tests.conftest import ( - setup_mock_dataset_owner_execute_result, - setup_mock_tenant_owner_execute_result, -) +from models import Account, Tenant, TenantAccountJoin +from models.account import TenantAccountRole +from models.dataset import Dataset, RateLimitLog +from models.enums import ApiTokenType +from models.model import ApiToken, App, AppMode, IconType def _configure_current_app_mock(mock_current_app): @@ -34,6 +36,51 @@ def _configure_current_app_mock(mock_current_app): mock_current_app._get_current_object = Mock(return_value=Mock()) +def _session_proxy(session: Session) -> MagicMock: + """Emulate Flask-SQLAlchemy's callable scoped-session proxy around a test session.""" + proxy = MagicMock(wraps=session) + proxy.return_value = session + return proxy + + +def _api_token(*, tenant_id: str, app_id: str | None = None, token_type: ApiTokenType) -> ApiToken: + return ApiToken( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + app_id=app_id, + type=token_type, + token="test_token", + ) + + +def _persist_workspace(session: Session) -> tuple[Tenant, Account, TenantAccountJoin]: + tenant = Tenant(name="Workspace") + account = Account(name="Owner", email=f"owner-{uuid.uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + ) + session.add_all([tenant, account, membership]) + session.commit() + return tenant, account, membership + + +def _app_model(*, tenant_id: str, enable_api: bool = True) -> App: + return App( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Service API App", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#FFFFFF", + enable_site=False, + enable_api=enable_api, + ) + + class TestValidateAndGetApiToken: """Test suite for validate_and_get_api_token function""" @@ -70,21 +117,24 @@ class TestValidateAndGetApiToken: def test_valid_token_returns_api_token(self, mock_fetch_token, mock_cache_cls, mock_record_usage, app: Flask): """Test that valid token returns the ApiToken object.""" # Arrange - mock_api_token = Mock(spec=ApiToken) - mock_api_token.token = "valid_token_123" - mock_api_token.type = "app" + api_token = _api_token( + tenant_id=str(uuid.uuid4()), + app_id=str(uuid.uuid4()), + token_type=ApiTokenType.APP, + ) + api_token.token = "valid_token_123" mock_cache_instance = Mock() mock_cache_instance.get.return_value = None # Cache miss mock_cache_cls.get = mock_cache_instance.get - mock_fetch_token.return_value = mock_api_token + mock_fetch_token.return_value = api_token # Act with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer valid_token_123"}): result = validate_and_get_api_token("app") # Assert - assert result == mock_api_token + assert result == api_token @patch("controllers.service_api.wraps.record_token_usage") @patch("controllers.service_api.wraps.ApiTokenCache") @@ -117,116 +167,124 @@ class TestValidateAppToken: return app @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.current_app") + @pytest.mark.parametrize( + "sqlite_session", + [(App, ApiToken, Tenant, Account, TenantAccountJoin)], + indirect=True, + ) def test_valid_app_token_allows_access( - self, mock_current_app, mock_validate_token, mock_db, mock_user_logged_in, app + self, + mock_current_app, + mock_validate_token, + mock_user_logged_in, + app: Flask, + sqlite_session: Session, ): """Test that valid app token allows access to decorated view.""" # Arrange _configure_current_app_mock(mock_current_app) - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_api_token.tenant_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_app = Mock() - mock_app.id = mock_api_token.app_id - mock_app.status = "normal" - mock_app.enable_api = True - mock_app.tenant_id = mock_api_token.tenant_id - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - mock_tenant.id = mock_api_token.tenant_id - - mock_account = Mock() - mock_account.id = str(uuid.uuid4()) - - # Use side_effect to return app first, then tenant via session.get() - mock_db.session.get.side_effect = [mock_app, mock_tenant] - - # Mock the tenant owner execute result (execute(select(...)).one_or_none()) - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) + tenant, account, _ = _persist_workspace(sqlite_session) + app_model = _app_model(tenant_id=tenant.id) + api_token = _api_token(tenant_id=tenant.id, app_id=app_model.id, token_type=ApiTokenType.APP) + sqlite_session.add_all([app_model, api_token]) + sqlite_session.commit() + mock_validate_token.return_value = api_token @validate_app_token def protected_view(app_model): return {"success": True, "app_id": app_model.id} # Act - with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}): + with ( + app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}), + patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)), + ): result = protected_view() # Assert assert result["success"] is True - assert result["app_id"] == mock_app.id + assert result["app_id"] == app_model.id + assert account.current_tenant_id == tenant.id - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_app_not_found_raises_forbidden(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) + def test_app_not_found_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that Forbidden is raised when app no longer exists.""" # Arrange - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_db.session.get.return_value = None + api_token = _api_token( + tenant_id=str(uuid.uuid4()), + app_id=str(uuid.uuid4()), + token_type=ApiTokenType.APP, + ) + mock_validate_token.return_value = api_token @validate_app_token def protected_view(**kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(Forbidden) as exc_info: protected_view() assert "no longer exists" in str(exc_info.value) - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_app_status_abnormal_raises_forbidden(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) + def test_app_status_abnormal_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that Forbidden is raised when app status is abnormal.""" # Arrange - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_app = Mock() - mock_app.status = "abnormal" - mock_db.session.get.return_value = mock_app + app_model = _app_model(tenant_id=str(uuid.uuid4())) + sqlite_session.add(app_model) + sqlite_session.commit() + app_model.status = "abnormal" + mock_validate_token.return_value = _api_token( + tenant_id=app_model.tenant_id, + app_id=app_model.id, + token_type=ApiTokenType.APP, + ) @validate_app_token def protected_view(**kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(Forbidden) as exc_info: protected_view() assert "status is abnormal" in str(exc_info.value) - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_app_api_disabled_raises_forbidden(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) + def test_app_api_disabled_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that Forbidden is raised when app API is disabled.""" # Arrange - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_app = Mock() - mock_app.status = "normal" - mock_app.enable_api = False - mock_db.session.get.return_value = mock_app + app_model = _app_model(tenant_id=str(uuid.uuid4()), enable_api=False) + sqlite_session.add(app_model) + sqlite_session.commit() + mock_validate_token.return_value = _api_token( + tenant_id=app_model.tenant_id, + app_id=app_model.id, + token_type=ApiTokenType.APP, + ) @validate_app_token def protected_view(**kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(Forbidden) as exc_info: protected_view() assert "API service has been disabled" in str(exc_info.value) @@ -468,26 +526,35 @@ class TestCloudEditionBillingRateLimitCheck: @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.FeatureService.get_knowledge_rate_limit") - @patch("controllers.service_api.wraps.db") - @patch("controllers.service_api.wraps.sessionmaker") + @pytest.mark.parametrize("sqlite_session", [(RateLimitLog,)], indirect=True) def test_rejects_over_rate_limit( - self, mock_sessionmaker, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask + self, + mock_get_rate_limit, + mock_validate_token, + app: Flask, + sqlite_session: Session, ): """Test that Forbidden is raised when over rate limit.""" # Arrange - mock_validate_token.return_value = Mock(tenant_id="tenant123") + tenant_id = str(uuid.uuid4()) + mock_validate_token.return_value = _api_token( + tenant_id=tenant_id, + token_type=ApiTokenType.DATASET, + ) mock_rate_limit = Mock() mock_rate_limit.enabled = True mock_rate_limit.limit = 10 mock_rate_limit.subscription_plan = "pro" mock_get_rate_limit.return_value = mock_rate_limit - rate_limit_log_session = MagicMock() - session_factory = MagicMock() - session_factory.begin.return_value.__enter__.return_value = rate_limit_log_session - mock_sessionmaker.return_value = session_factory - with patch("controllers.service_api.wraps.redis_client") as mock_redis: + with ( + patch("controllers.service_api.wraps.redis_client") as mock_redis, + patch( + "controllers.service_api.wraps.db", + SimpleNamespace(engine=sqlite_session.get_bind()), + ), + ): mock_redis.zcard.return_value = 15 # Over limit @cloud_edition_billing_rate_limit_check("knowledge", "dataset") @@ -499,9 +566,12 @@ class TestCloudEditionBillingRateLimitCheck: with pytest.raises(Forbidden) as exc_info: knowledge_request() assert "rate limit" in str(exc_info.value) - mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) - rate_limit_log_session.add.assert_called_once() - mock_db.session.commit.assert_not_called() + + persisted_logs = sqlite_session.scalars(select(RateLimitLog)).all() + assert len(persisted_logs) == 1 + assert persisted_logs[0].tenant_id == tenant_id + assert persisted_logs[0].subscription_plan == "pro" + assert persisted_logs[0].operation == "knowledge" class TestValidateDatasetToken: @@ -515,65 +585,62 @@ class TestValidateDatasetToken: return app @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.current_app") - def test_valid_dataset_token(self, mock_current_app, mock_validate_token, mock_db, mock_user_logged_in, app: Flask): + @pytest.mark.parametrize( + "sqlite_session", + [(Tenant, Account, TenantAccountJoin)], + indirect=True, + ) + def test_valid_dataset_token( + self, + mock_current_app, + mock_validate_token, + mock_user_logged_in, + app: Flask, + sqlite_session: Session, + ): """Test that valid dataset token allows access.""" # Arrange _configure_current_app_mock(mock_current_app) - tenant_id = str(uuid.uuid4()) - mock_api_token = Mock() - mock_api_token.tenant_id = tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.id = tenant_id - mock_tenant.status = TenantStatus.NORMAL - - mock_ta = Mock() - mock_ta.account_id = str(uuid.uuid4()) - - mock_account = Mock() - mock_account.id = mock_ta.account_id - mock_account.current_tenant = mock_tenant - - # Mock the tenant account join query (execute(select(...)).one_or_none()) - setup_mock_dataset_owner_execute_result(mock_db, mock_tenant, mock_ta) - - # Mock the account lookup via session.get() - mock_db.session.get.return_value = mock_account + tenant, account, _ = _persist_workspace(sqlite_session) + api_token = _api_token(tenant_id=tenant.id, token_type=ApiTokenType.DATASET) + mock_validate_token.return_value = api_token @validate_dataset_token def protected_view(tenant_id): return {"success": True, "tenant_id": tenant_id} # Act - with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}): + with ( + app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}), + patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)), + ): result = protected_view() # Assert assert result["success"] is True - assert result["tenant_id"] == tenant_id + assert result["tenant_id"] == tenant.id + assert account.current_tenant_id == tenant.id - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_dataset_not_found_raises_not_found(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) + def test_dataset_not_found_raises_not_found(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that NotFound is raised when dataset doesn't exist.""" # Arrange - mock_api_token = Mock() - mock_api_token.tenant_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_db.session.scalar.return_value = None + api_token = _api_token(tenant_id=str(uuid.uuid4()), token_type=ApiTokenType.DATASET) + mock_validate_token.return_value = api_token @validate_dataset_token def protected_view(dataset_id=None, **kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(NotFound) as exc_info: protected_view(dataset_id=str(uuid.uuid4())) assert "Dataset not found" in str(exc_info.value)