mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
test: move API key authentication coverage to unit tests (#38919)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
autofix-ci[bot]
parent
557d92735d
commit
0c15a71cc4
+62
-77
@@ -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):
|
||||
Reference in New Issue
Block a user