Merge branch 'main' into copilot/fix-password-length-issue
This commit is contained in:
@@ -0,0 +1,108 @@
|
||||
"""
|
||||
Unit tests for the FastAPI application factory and core endpoints.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from httpx import AsyncClient, ASGITransport
|
||||
from app.main import create_application
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app():
|
||||
"""Create a fresh application instance for testing."""
|
||||
return create_application()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestRootEndpoint:
|
||||
"""Test root endpoint"""
|
||||
|
||||
async def test_root_returns_200(self, app):
|
||||
"""Test that root endpoint returns 200"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/")
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
async def test_root_returns_api_info(self, app):
|
||||
"""Test that root returns API information"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/")
|
||||
|
||||
data = response.json()
|
||||
assert "message" in data
|
||||
assert "version" in data
|
||||
assert "docs" in data
|
||||
assert data["docs"] == "/api/docs"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestHealthEndpoint:
|
||||
"""Test health check endpoint"""
|
||||
|
||||
async def test_health_returns_200(self, app):
|
||||
"""Test that health endpoint returns 200"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/health")
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
async def test_health_returns_healthy(self, app):
|
||||
"""Test that health endpoint returns healthy status"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/health")
|
||||
|
||||
data = response.json()
|
||||
assert data["status"] == "healthy"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestSecurityHeaders:
|
||||
"""Test that security headers are present in responses"""
|
||||
|
||||
async def test_security_headers_on_root(self, app):
|
||||
"""Test security headers on root endpoint"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/")
|
||||
|
||||
assert response.headers["X-Frame-Options"] == "DENY"
|
||||
assert response.headers["X-Content-Type-Options"] == "nosniff"
|
||||
assert response.headers["X-XSS-Protection"] == "1; mode=block"
|
||||
|
||||
async def test_security_headers_on_health(self, app):
|
||||
"""Test security headers on health endpoint"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/health")
|
||||
|
||||
assert response.headers["X-Frame-Options"] == "DENY"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestApplicationFactory:
|
||||
"""Test the create_application factory"""
|
||||
|
||||
async def test_app_title(self, app):
|
||||
"""Test that app has correct title"""
|
||||
assert app.title == "POP3 Forwarder SaaS"
|
||||
|
||||
async def test_app_version(self, app):
|
||||
"""Test that app has a version"""
|
||||
assert app.version is not None
|
||||
assert len(app.version) > 0
|
||||
|
||||
async def test_openapi_endpoint(self, app):
|
||||
"""Test that OpenAPI schema is available"""
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/api/openapi.json")
|
||||
|
||||
assert response.status_code == 200
|
||||
schema = response.json()
|
||||
assert "openapi" in schema
|
||||
assert "info" in schema
|
||||
@@ -0,0 +1,127 @@
|
||||
"""
|
||||
Unit tests for security middleware.
|
||||
"""
|
||||
|
||||
from starlette.testclient import TestClient
|
||||
from starlette.applications import Starlette
|
||||
from starlette.responses import PlainTextResponse
|
||||
from starlette.routing import Route
|
||||
|
||||
from app.core.middleware import SecurityHeadersMiddleware, CSRFProtectionMiddleware
|
||||
|
||||
|
||||
def _make_app(middleware_classes):
|
||||
"""Helper to build a Starlette app with given middleware."""
|
||||
|
||||
async def homepage(request):
|
||||
return PlainTextResponse("OK")
|
||||
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Route("/", homepage, methods=["GET", "HEAD", "POST", "OPTIONS"]),
|
||||
Route("/api/v1/auth/login", homepage, methods=["GET", "POST"]),
|
||||
]
|
||||
)
|
||||
for cls in middleware_classes:
|
||||
app.add_middleware(cls)
|
||||
return app
|
||||
|
||||
|
||||
class TestSecurityHeadersMiddleware:
|
||||
"""Test security headers added to all responses"""
|
||||
|
||||
def setup_method(self):
|
||||
app = _make_app([SecurityHeadersMiddleware])
|
||||
self.client = TestClient(app)
|
||||
|
||||
def test_x_frame_options_header(self):
|
||||
"""Test X-Frame-Options is set to DENY"""
|
||||
response = self.client.get("/")
|
||||
assert response.headers["X-Frame-Options"] == "DENY"
|
||||
|
||||
def test_x_content_type_options_header(self):
|
||||
"""Test X-Content-Type-Options is set to nosniff"""
|
||||
response = self.client.get("/")
|
||||
assert response.headers["X-Content-Type-Options"] == "nosniff"
|
||||
|
||||
def test_x_xss_protection_header(self):
|
||||
"""Test X-XSS-Protection header is set"""
|
||||
response = self.client.get("/")
|
||||
assert response.headers["X-XSS-Protection"] == "1; mode=block"
|
||||
|
||||
def test_content_security_policy_header(self):
|
||||
"""Test Content-Security-Policy header is present"""
|
||||
response = self.client.get("/")
|
||||
csp = response.headers["Content-Security-Policy"]
|
||||
assert "default-src 'self'" in csp
|
||||
assert "script-src" in csp
|
||||
|
||||
def test_referrer_policy_header(self):
|
||||
"""Test Referrer-Policy header"""
|
||||
response = self.client.get("/")
|
||||
assert response.headers["Referrer-Policy"] == "strict-origin-when-cross-origin"
|
||||
|
||||
def test_permissions_policy_header(self):
|
||||
"""Test Permissions-Policy header"""
|
||||
response = self.client.get("/")
|
||||
policy = response.headers["Permissions-Policy"]
|
||||
assert "geolocation=()" in policy
|
||||
assert "microphone=()" in policy
|
||||
assert "camera=()" in policy
|
||||
|
||||
def test_no_hsts_for_localhost(self):
|
||||
"""Test that HSTS header check depends on hostname"""
|
||||
# The HSTS header is only skipped when hostname is localhost or 127.0.0.1.
|
||||
# TestClient uses 'testserver' as hostname, which is not in the skip list,
|
||||
# so HSTS will be set. Verify the logic works with a direct check.
|
||||
response = self.client.get("/")
|
||||
# TestClient hostname is 'testserver', not localhost, so HSTS IS set
|
||||
assert "Strict-Transport-Security" in response.headers
|
||||
|
||||
|
||||
class TestCSRFProtectionMiddleware:
|
||||
"""Test CSRF protection middleware"""
|
||||
|
||||
def setup_method(self):
|
||||
app = _make_app([CSRFProtectionMiddleware])
|
||||
self.client = TestClient(app)
|
||||
|
||||
def test_get_requests_pass_through(self):
|
||||
"""Test that GET requests are not blocked"""
|
||||
response = self.client.get("/")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_head_requests_pass_through(self):
|
||||
"""Test that HEAD requests are not blocked"""
|
||||
response = self.client.head("/")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_options_requests_pass_through(self):
|
||||
"""Test that OPTIONS requests are not blocked"""
|
||||
response = self.client.options("/")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_exempt_paths_pass_through(self):
|
||||
"""Test that exempt paths are not CSRF-checked for POST"""
|
||||
response = self.client.post("/api/v1/auth/login")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_post_to_non_exempt_path_passes(self):
|
||||
"""Test that POST to non-exempt path also passes (JWT provides CSRF protection)"""
|
||||
response = self.client.post("/")
|
||||
assert response.status_code == 200
|
||||
|
||||
def test_generate_csrf_token(self):
|
||||
"""Test CSRF token generation produces valid token"""
|
||||
token = CSRFProtectionMiddleware._generate_csrf_token()
|
||||
assert isinstance(token, str)
|
||||
assert len(token) == 43 # token_urlsafe(32) produces 43 chars
|
||||
|
||||
def test_validate_csrf_token_valid(self):
|
||||
"""Test CSRF token validation with valid token"""
|
||||
token = CSRFProtectionMiddleware._generate_csrf_token()
|
||||
assert CSRFProtectionMiddleware._validate_csrf_token(token) is True
|
||||
|
||||
def test_validate_csrf_token_invalid(self):
|
||||
"""Test CSRF token validation with invalid token"""
|
||||
assert CSRFProtectionMiddleware._validate_csrf_token("short") is False
|
||||
@@ -0,0 +1,277 @@
|
||||
"""
|
||||
Unit tests for Pydantic schema validation.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
from app.models.schemas import (
|
||||
UserCreate,
|
||||
UserUpdate,
|
||||
MailAccountCreate,
|
||||
MailAccountUpdate,
|
||||
MailAccountTestRequest,
|
||||
MailAccountAutoDetectRequest,
|
||||
MailProtocol,
|
||||
DeliveryMethod,
|
||||
SubscriptionTier,
|
||||
AccountStatus,
|
||||
NotificationChannel,
|
||||
Token,
|
||||
TokenPayload,
|
||||
GoogleAuthRequest,
|
||||
NotificationConfigCreate,
|
||||
SubscriptionCheckoutRequest,
|
||||
ProviderPreset,
|
||||
)
|
||||
|
||||
|
||||
class TestEnums:
|
||||
"""Test enum values"""
|
||||
|
||||
def test_subscription_tiers(self):
|
||||
"""Test all subscription tier values"""
|
||||
assert SubscriptionTier.FREE == "free"
|
||||
assert SubscriptionTier.BASIC == "basic"
|
||||
assert SubscriptionTier.PRO == "pro"
|
||||
assert SubscriptionTier.ENTERPRISE == "enterprise"
|
||||
|
||||
def test_mail_protocols(self):
|
||||
"""Test all mail protocol values"""
|
||||
assert MailProtocol.POP3 == "pop3"
|
||||
assert MailProtocol.POP3_SSL == "pop3_ssl"
|
||||
assert MailProtocol.IMAP == "imap"
|
||||
assert MailProtocol.IMAP_SSL == "imap_ssl"
|
||||
|
||||
def test_account_status(self):
|
||||
"""Test all account status values"""
|
||||
assert AccountStatus.ACTIVE == "active"
|
||||
assert AccountStatus.INACTIVE == "inactive"
|
||||
assert AccountStatus.ERROR == "error"
|
||||
assert AccountStatus.TESTING == "testing"
|
||||
|
||||
def test_delivery_method(self):
|
||||
"""Test all delivery method values"""
|
||||
assert DeliveryMethod.SMTP == "smtp"
|
||||
assert DeliveryMethod.GMAIL_API == "gmail_api"
|
||||
|
||||
def test_notification_channels(self):
|
||||
"""Test all notification channel values"""
|
||||
assert NotificationChannel.EMAIL == "email"
|
||||
assert NotificationChannel.TELEGRAM == "telegram"
|
||||
assert NotificationChannel.WEBHOOK == "webhook"
|
||||
assert NotificationChannel.SLACK == "slack"
|
||||
assert NotificationChannel.DISCORD == "discord"
|
||||
|
||||
|
||||
class TestUserSchemas:
|
||||
"""Test user-related schemas"""
|
||||
|
||||
def test_user_create_with_email(self):
|
||||
"""Test UserCreate with valid email"""
|
||||
user = UserCreate(email="test@example.com", password="password123")
|
||||
assert user.email == "test@example.com"
|
||||
assert user.password == "password123"
|
||||
|
||||
def test_user_create_without_password(self):
|
||||
"""Test UserCreate without password (OAuth users)"""
|
||||
user = UserCreate(email="test@example.com")
|
||||
assert user.password is None
|
||||
|
||||
def test_user_create_with_full_name(self):
|
||||
"""Test UserCreate with full name"""
|
||||
user = UserCreate(
|
||||
email="test@example.com", full_name="Test User", password="pass"
|
||||
)
|
||||
assert user.full_name == "Test User"
|
||||
|
||||
def test_user_create_invalid_email(self):
|
||||
"""Test UserCreate rejects invalid email"""
|
||||
with pytest.raises(ValidationError):
|
||||
UserCreate(email="not-an-email", password="pass")
|
||||
|
||||
def test_user_update_partial(self):
|
||||
"""Test UserUpdate with partial data"""
|
||||
update = UserUpdate(full_name="New Name")
|
||||
assert update.full_name == "New Name"
|
||||
assert update.email is None
|
||||
|
||||
|
||||
class TestTokenSchemas:
|
||||
"""Test token schemas"""
|
||||
|
||||
def test_token_schema(self):
|
||||
"""Test Token schema"""
|
||||
token = Token(
|
||||
access_token="abc123", refresh_token="def456", token_type="bearer"
|
||||
)
|
||||
assert token.access_token == "abc123"
|
||||
assert token.token_type == "bearer"
|
||||
|
||||
def test_token_payload_schema(self):
|
||||
"""Test TokenPayload schema"""
|
||||
payload = TokenPayload(sub=42, type="access")
|
||||
assert payload.sub == 42
|
||||
assert payload.type == "access"
|
||||
|
||||
def test_google_auth_request(self):
|
||||
"""Test GoogleAuthRequest schema"""
|
||||
req = GoogleAuthRequest(
|
||||
code="auth-code-123", redirect_uri="http://localhost:3000/callback"
|
||||
)
|
||||
assert req.code == "auth-code-123"
|
||||
|
||||
|
||||
class TestMailAccountSchemas:
|
||||
"""Test mail account schemas"""
|
||||
|
||||
def test_mail_account_create_valid(self):
|
||||
"""Test creating a valid mail account"""
|
||||
account = MailAccountCreate(
|
||||
name="Test Account",
|
||||
email_address="user@example.com",
|
||||
host="imap.example.com",
|
||||
port=993,
|
||||
username="user@example.com",
|
||||
password="secret",
|
||||
forward_to="me@gmail.com",
|
||||
)
|
||||
assert account.name == "Test Account"
|
||||
assert account.protocol == MailProtocol.POP3_SSL # default
|
||||
assert account.use_ssl is True
|
||||
|
||||
def test_mail_account_create_invalid_port(self):
|
||||
"""Test that invalid port is rejected"""
|
||||
with pytest.raises(ValidationError):
|
||||
MailAccountCreate(
|
||||
name="Test",
|
||||
email_address="user@example.com",
|
||||
host="imap.example.com",
|
||||
port=0, # invalid
|
||||
username="user@example.com",
|
||||
password="secret",
|
||||
forward_to="me@gmail.com",
|
||||
)
|
||||
|
||||
def test_mail_account_create_port_too_high(self):
|
||||
"""Test that port above 65535 is rejected"""
|
||||
with pytest.raises(ValidationError):
|
||||
MailAccountCreate(
|
||||
name="Test",
|
||||
email_address="user@example.com",
|
||||
host="imap.example.com",
|
||||
port=70000, # invalid
|
||||
username="user@example.com",
|
||||
password="secret",
|
||||
forward_to="me@gmail.com",
|
||||
)
|
||||
|
||||
def test_mail_account_update_partial(self):
|
||||
"""Test partial mail account update"""
|
||||
update = MailAccountUpdate(is_enabled=False)
|
||||
assert update.is_enabled is False
|
||||
assert update.name is None
|
||||
assert update.password is None
|
||||
|
||||
def test_mail_account_test_request(self):
|
||||
"""Test mail account test connection schema"""
|
||||
req = MailAccountTestRequest(
|
||||
host="imap.gmail.com",
|
||||
port=993,
|
||||
protocol=MailProtocol.IMAP_SSL,
|
||||
username="user@gmail.com",
|
||||
password="app-password",
|
||||
)
|
||||
assert req.host == "imap.gmail.com"
|
||||
|
||||
def test_auto_detect_request(self):
|
||||
"""Test auto-detect request schema"""
|
||||
req = MailAccountAutoDetectRequest(email_address="user@gmail.com")
|
||||
assert req.email_address == "user@gmail.com"
|
||||
|
||||
def test_auto_detect_invalid_email(self):
|
||||
"""Test auto-detect rejects invalid email"""
|
||||
with pytest.raises(ValidationError):
|
||||
MailAccountAutoDetectRequest(email_address="not-email")
|
||||
|
||||
|
||||
class TestNotificationSchemas:
|
||||
"""Test notification schemas"""
|
||||
|
||||
def test_notification_config_create(self):
|
||||
"""Test creating notification config"""
|
||||
config = NotificationConfigCreate(
|
||||
channel=NotificationChannel.TELEGRAM,
|
||||
config={"bot_token": "123:abc", "chat_id": "456"},
|
||||
)
|
||||
assert config.channel == NotificationChannel.TELEGRAM
|
||||
assert config.notify_on_errors is True # default
|
||||
assert config.notify_on_success is False # default
|
||||
|
||||
def test_notification_config_threshold_validation(self):
|
||||
"""Test notification threshold validation"""
|
||||
with pytest.raises(ValidationError):
|
||||
NotificationConfigCreate(
|
||||
channel=NotificationChannel.EMAIL,
|
||||
config={},
|
||||
notify_threshold=0, # must be > 0
|
||||
)
|
||||
|
||||
|
||||
class TestSubscriptionSchemas:
|
||||
"""Test subscription schemas"""
|
||||
|
||||
def test_subscription_checkout_request_monthly(self):
|
||||
"""Test subscription checkout with monthly billing"""
|
||||
req = SubscriptionCheckoutRequest(
|
||||
tier=SubscriptionTier.PRO,
|
||||
billing_period="monthly",
|
||||
success_url="https://example.com/success",
|
||||
cancel_url="https://example.com/cancel",
|
||||
)
|
||||
assert req.tier == SubscriptionTier.PRO
|
||||
assert req.billing_period == "monthly"
|
||||
|
||||
def test_subscription_checkout_request_yearly(self):
|
||||
"""Test subscription checkout with yearly billing"""
|
||||
req = SubscriptionCheckoutRequest(
|
||||
tier=SubscriptionTier.BASIC,
|
||||
billing_period="yearly",
|
||||
success_url="https://example.com/success",
|
||||
cancel_url="https://example.com/cancel",
|
||||
)
|
||||
assert req.billing_period == "yearly"
|
||||
|
||||
def test_subscription_checkout_invalid_period(self):
|
||||
"""Test that invalid billing period is rejected"""
|
||||
with pytest.raises(ValidationError):
|
||||
SubscriptionCheckoutRequest(
|
||||
tier=SubscriptionTier.BASIC,
|
||||
billing_period="quarterly", # invalid
|
||||
success_url="https://example.com/success",
|
||||
cancel_url="https://example.com/cancel",
|
||||
)
|
||||
|
||||
|
||||
class TestProviderPresetSchema:
|
||||
"""Test provider preset schema"""
|
||||
|
||||
def test_provider_preset_with_imap(self):
|
||||
"""Test provider preset with IMAP config"""
|
||||
preset = ProviderPreset(
|
||||
id="gmail",
|
||||
name="Gmail",
|
||||
domains=["gmail.com", "googlemail.com"],
|
||||
imap_ssl={"host": "imap.gmail.com", "port": 993},
|
||||
)
|
||||
assert preset.id == "gmail"
|
||||
assert "gmail.com" in preset.domains
|
||||
|
||||
def test_provider_preset_without_pop3(self):
|
||||
"""Test provider preset without POP3 (IMAP only)"""
|
||||
preset = ProviderPreset(
|
||||
id="posteo",
|
||||
name="Posteo",
|
||||
domains=["posteo.de"],
|
||||
imap_ssl={"host": "posteo.de", "port": 993},
|
||||
)
|
||||
assert preset.pop3_ssl is None
|
||||
@@ -0,0 +1,188 @@
|
||||
"""
|
||||
Unit tests for extended security module functionality.
|
||||
"""
|
||||
|
||||
from datetime import timedelta
|
||||
from app.core.security import (
|
||||
create_access_token,
|
||||
create_refresh_token,
|
||||
decode_token,
|
||||
generate_random_token,
|
||||
encrypt_credential,
|
||||
decrypt_credential,
|
||||
CredentialEncryption,
|
||||
)
|
||||
|
||||
|
||||
class TestAccessToken:
|
||||
"""Test JWT access token creation and decoding"""
|
||||
|
||||
def test_create_access_token_with_custom_expiry(self):
|
||||
"""Test creating an access token with custom expiry"""
|
||||
data = {"sub": "test@example.com"}
|
||||
token = create_access_token(data, expires_delta=timedelta(hours=1))
|
||||
|
||||
assert isinstance(token, str)
|
||||
assert token.count(".") == 2
|
||||
|
||||
def test_decode_valid_access_token(self):
|
||||
"""Test decoding a valid access token"""
|
||||
data = {"sub": "user123"}
|
||||
token = create_access_token(data)
|
||||
|
||||
payload = decode_token(token)
|
||||
assert payload is not None
|
||||
assert payload["sub"] == "user123"
|
||||
assert payload["type"] == "access"
|
||||
|
||||
def test_decode_invalid_token_returns_none(self):
|
||||
"""Test that decoding an invalid token returns None"""
|
||||
result = decode_token("invalid.token.string")
|
||||
assert result is None
|
||||
|
||||
def test_decode_empty_token_returns_none(self):
|
||||
"""Test that decoding an empty string returns None"""
|
||||
result = decode_token("")
|
||||
assert result is None
|
||||
|
||||
def test_access_token_contains_type(self):
|
||||
"""Test that access token payload contains type 'access'"""
|
||||
data = {"sub": "user@example.com"}
|
||||
token = create_access_token(data)
|
||||
payload = decode_token(token)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["type"] == "access"
|
||||
|
||||
def test_access_token_contains_expiry(self):
|
||||
"""Test that access token payload contains expiry"""
|
||||
data = {"sub": "user@example.com"}
|
||||
token = create_access_token(data)
|
||||
payload = decode_token(token)
|
||||
|
||||
assert payload is not None
|
||||
assert "exp" in payload
|
||||
|
||||
|
||||
class TestRefreshToken:
|
||||
"""Test JWT refresh token creation and decoding"""
|
||||
|
||||
def test_create_refresh_token(self):
|
||||
"""Test refresh token creation"""
|
||||
data = {"sub": "test@example.com"}
|
||||
token = create_refresh_token(data)
|
||||
|
||||
assert isinstance(token, str)
|
||||
assert token.count(".") == 2
|
||||
|
||||
def test_decode_refresh_token(self):
|
||||
"""Test decoding a valid refresh token"""
|
||||
data = {"sub": "user456"}
|
||||
token = create_refresh_token(data)
|
||||
|
||||
payload = decode_token(token)
|
||||
assert payload is not None
|
||||
assert payload["sub"] == "user456"
|
||||
assert payload["type"] == "refresh"
|
||||
|
||||
def test_refresh_token_different_from_access(self):
|
||||
"""Test that refresh and access tokens are different"""
|
||||
data = {"sub": "test@example.com"}
|
||||
access = create_access_token(data)
|
||||
refresh = create_refresh_token(data)
|
||||
|
||||
assert access != refresh
|
||||
|
||||
|
||||
class TestRandomToken:
|
||||
"""Test random token generation"""
|
||||
|
||||
def test_generate_random_token_default_length(self):
|
||||
"""Test generating a random token with default length"""
|
||||
token = generate_random_token()
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_generate_random_token_custom_length(self):
|
||||
"""Test generating a random token with custom length"""
|
||||
token = generate_random_token(64)
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_random_tokens_are_unique(self):
|
||||
"""Test that generated tokens are unique"""
|
||||
tokens = {generate_random_token() for _ in range(10)}
|
||||
assert len(tokens) == 10
|
||||
|
||||
|
||||
class TestGlobalEncryptionFunctions:
|
||||
"""Test global encryption convenience functions"""
|
||||
|
||||
def test_encrypt_credential_returns_string(self):
|
||||
"""Test that encrypt_credential returns a non-empty string"""
|
||||
encrypted = encrypt_credential("my-password")
|
||||
assert isinstance(encrypted, str)
|
||||
assert len(encrypted) > 0
|
||||
|
||||
def test_decrypt_credential_roundtrip(self):
|
||||
"""Test encrypt/decrypt roundtrip with global functions"""
|
||||
original = "super-secret-password-123"
|
||||
encrypted = encrypt_credential(original)
|
||||
decrypted = decrypt_credential(encrypted)
|
||||
assert decrypted == original
|
||||
|
||||
def test_encrypt_credential_is_not_plaintext(self):
|
||||
"""Test that encrypted credential differs from plaintext"""
|
||||
password = "my-password"
|
||||
encrypted = encrypt_credential(password)
|
||||
assert encrypted != password
|
||||
|
||||
|
||||
class TestCredentialEncryptionEdgeCases:
|
||||
"""Test edge cases in credential encryption"""
|
||||
|
||||
def test_encrypt_empty_string(self):
|
||||
"""Test encrypting an empty string"""
|
||||
encryptor = CredentialEncryption(user_id=1)
|
||||
encrypted = encryptor.encrypt("")
|
||||
decrypted = encryptor.decrypt(encrypted)
|
||||
assert decrypted == ""
|
||||
|
||||
def test_encrypt_long_string(self):
|
||||
"""Test encrypting a very long string"""
|
||||
long_password = "a" * 10000
|
||||
encryptor = CredentialEncryption(user_id=1)
|
||||
encrypted = encryptor.encrypt(long_password)
|
||||
decrypted = encryptor.decrypt(encrypted)
|
||||
assert decrypted == long_password
|
||||
|
||||
def test_encrypt_special_characters(self):
|
||||
"""Test encrypting a string with special characters"""
|
||||
special = "p@$$w0rd!#%^&*()_+-=[]{}|;':\",./<>?"
|
||||
encryptor = CredentialEncryption(user_id=1)
|
||||
encrypted = encryptor.encrypt(special)
|
||||
decrypted = encryptor.decrypt(encrypted)
|
||||
assert decrypted == special
|
||||
|
||||
def test_encrypt_unicode(self):
|
||||
"""Test encrypting unicode characters"""
|
||||
unicode_str = "密码テスト🔒"
|
||||
encryptor = CredentialEncryption(user_id=1)
|
||||
encrypted = encryptor.encrypt(unicode_str)
|
||||
decrypted = encryptor.decrypt(encrypted)
|
||||
assert decrypted == unicode_str
|
||||
|
||||
def test_custom_key(self):
|
||||
"""Test encryption with a custom key"""
|
||||
key = "custom-encryption-key-that-is-at-least-32-chars-long"
|
||||
encryptor = CredentialEncryption(key=key, user_id=1)
|
||||
encrypted = encryptor.encrypt("test-data")
|
||||
decrypted = encryptor.decrypt(encrypted)
|
||||
assert decrypted == "test-data"
|
||||
|
||||
def test_system_salt_without_user_id(self):
|
||||
"""Test encryption with system salt (no user_id)"""
|
||||
encryptor = CredentialEncryption()
|
||||
encrypted = encryptor.encrypt("system-data")
|
||||
decrypted = encryptor.decrypt(encrypted)
|
||||
assert decrypted == "system-data"
|
||||
Reference in New Issue
Block a user