Improve Strategy typing

This commit is contained in:
François Voron
2022-05-01 13:57:26 +02:00
parent b6d7c6a621
commit 2cf0ebcdaa
6 changed files with 42 additions and 26 deletions

View File

@@ -9,7 +9,9 @@ from fastapi_users.authentication.strategy.db.models import AP
from fastapi_users.manager import BaseUserManager, InvalidID, UserNotExists
class DatabaseStrategy(Strategy, Generic[models.UP, AP]):
class DatabaseStrategy(
Strategy[models.UP, models.ID], Generic[models.UP, models.ID, AP]
):
def __init__(
self, database: AccessTokenDatabase[AP], lifetime_seconds: Optional[int] = None
):

View File

@@ -11,7 +11,7 @@ from fastapi_users.jwt import SecretType, decode_jwt, generate_jwt
from fastapi_users.manager import BaseUserManager, InvalidID, UserNotExists
class JWTStrategy(Strategy, Generic[models.UP]):
class JWTStrategy(Strategy[models.UP, models.ID], Generic[models.UP, models.ID]):
def __init__(
self,
secret: SecretType,

View File

@@ -8,7 +8,7 @@ from fastapi_users.authentication.strategy.base import Strategy
from fastapi_users.manager import BaseUserManager, InvalidID, UserNotExists
class RedisStrategy(Strategy, Generic[models.UP]):
class RedisStrategy(Strategy[models.UP, models.ID], Generic[models.UP, models.ID]):
def __init__(self, redis: aioredis.Redis, lifetime_seconds: Optional[int] = None):
self.redis = redis
self.lifetime_seconds = lifetime_seconds

View File

@@ -10,7 +10,7 @@ from fastapi_users.authentication.strategy import (
AccessTokenProtocol,
DatabaseStrategy,
)
from tests.conftest import UserModel, IDType
from tests.conftest import IDType, UserModel
@dataclasses.dataclass
@@ -75,7 +75,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_missing_token(
self,
database_strategy: DatabaseStrategy[UserModel, AccessTokenModel],
database_strategy: DatabaseStrategy[UserModel, IDType, AccessTokenModel],
user_manager,
):
authenticated_user = await database_strategy.read_token(None, user_manager)
@@ -84,7 +84,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_invalid_token(
self,
database_strategy: DatabaseStrategy[UserModel, AccessTokenModel],
database_strategy: DatabaseStrategy[UserModel, IDType, AccessTokenModel],
user_manager,
):
authenticated_user = await database_strategy.read_token("TOKEN", user_manager)
@@ -93,7 +93,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_valid_token_not_existing_user(
self,
database_strategy: DatabaseStrategy[UserModel, AccessTokenModel],
database_strategy: DatabaseStrategy[UserModel, IDType, AccessTokenModel],
access_token_database: AccessTokenDatabaseMock,
user_manager,
):
@@ -109,7 +109,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_valid_token(
self,
database_strategy: DatabaseStrategy[UserModel, AccessTokenModel],
database_strategy: DatabaseStrategy[UserModel, IDType, AccessTokenModel],
access_token_database: AccessTokenDatabaseMock,
user_manager,
user: UserModel,
@@ -123,7 +123,7 @@ class TestReadToken:
@pytest.mark.authentication
@pytest.mark.asyncio
async def test_write_token(
database_strategy: DatabaseStrategy[UserModel, AccessTokenModel],
database_strategy: DatabaseStrategy[UserModel, IDType, AccessTokenModel],
access_token_database: AccessTokenDatabaseMock,
user: UserModel,
):
@@ -137,7 +137,7 @@ async def test_write_token(
@pytest.mark.authentication
@pytest.mark.asyncio
async def test_destroy_token(
database_strategy: DatabaseStrategy[UserModel, AccessTokenModel],
database_strategy: DatabaseStrategy[UserModel, IDType, AccessTokenModel],
access_token_database: AccessTokenDatabaseMock,
user: UserModel,
):

View File

@@ -5,6 +5,7 @@ from fastapi_users.authentication.strategy import (
StrategyDestroyNotSupportedError,
)
from fastapi_users.jwt import SecretType, decode_jwt, generate_jwt
from tests.conftest import IDType, UserModel
LIFETIME = 3600
@@ -74,7 +75,7 @@ def jwt_strategy(request, secret: SecretType):
@pytest.fixture
def token(jwt_strategy: JWTStrategy):
def token(jwt_strategy: JWTStrategy[UserModel, IDType]):
def _token(user_id=None, lifetime=LIFETIME):
data = {"aud": "fastapi-users:auth"}
if user_id is not None:
@@ -90,32 +91,36 @@ def token(jwt_strategy: JWTStrategy):
@pytest.mark.authentication
class TestReadToken:
@pytest.mark.asyncio
async def test_missing_token(self, jwt_strategy: JWTStrategy, user_manager):
async def test_missing_token(
self, jwt_strategy: JWTStrategy[UserModel, IDType], user_manager
):
authenticated_user = await jwt_strategy.read_token(None, user_manager)
assert authenticated_user is None
@pytest.mark.asyncio
async def test_invalid_token(self, jwt_strategy: JWTStrategy, user_manager):
async def test_invalid_token(
self, jwt_strategy: JWTStrategy[UserModel, IDType], user_manager
):
authenticated_user = await jwt_strategy.read_token("foo", user_manager)
assert authenticated_user is None
@pytest.mark.asyncio
async def test_valid_token_missing_user_payload(
self, jwt_strategy: JWTStrategy, user_manager, token
self, jwt_strategy: JWTStrategy[UserModel, IDType], user_manager, token
):
authenticated_user = await jwt_strategy.read_token(token(), user_manager)
assert authenticated_user is None
@pytest.mark.asyncio
async def test_valid_token_invalid_uuid(
self, jwt_strategy: JWTStrategy, user_manager, token
self, jwt_strategy: JWTStrategy[UserModel, IDType], user_manager, token
):
authenticated_user = await jwt_strategy.read_token(token("foo"), user_manager)
assert authenticated_user is None
@pytest.mark.asyncio
async def test_valid_token_not_existing_user(
self, jwt_strategy: JWTStrategy, user_manager, token
self, jwt_strategy: JWTStrategy[UserModel, IDType], user_manager, token
):
authenticated_user = await jwt_strategy.read_token(
token("d35d213e-f3d8-4f08-954a-7e0d1bea286f"), user_manager
@@ -124,7 +129,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_valid_token(
self, jwt_strategy: JWTStrategy, user_manager, token, user
self, jwt_strategy: JWTStrategy[UserModel, IDType], user_manager, token, user
):
authenticated_user = await jwt_strategy.read_token(token(user.id), user_manager)
assert authenticated_user is not None
@@ -134,7 +139,7 @@ class TestReadToken:
@pytest.mark.parametrize("jwt_strategy", ["HS256", "RS256", "ES256"], indirect=True)
@pytest.mark.authentication
@pytest.mark.asyncio
async def test_write_token(jwt_strategy: JWTStrategy, user):
async def test_write_token(jwt_strategy: JWTStrategy[UserModel, IDType], user):
token = await jwt_strategy.write_token(user)
decoded = decode_jwt(
@@ -149,6 +154,6 @@ async def test_write_token(jwt_strategy: JWTStrategy, user):
@pytest.mark.parametrize("jwt_strategy", ["HS256", "RS256", "ES256"], indirect=True)
@pytest.mark.authentication
@pytest.mark.asyncio
async def test_destroy_token(jwt_strategy: JWTStrategy, user):
async def test_destroy_token(jwt_strategy: JWTStrategy[UserModel, IDType], user):
with pytest.raises(StrategyDestroyNotSupportedError):
await jwt_strategy.destroy_token("TOKEN", user)

View File

@@ -4,6 +4,7 @@ from typing import Dict, Optional, Tuple
import pytest
from fastapi_users.authentication.strategy import RedisStrategy
from tests.conftest import IDType, UserModel
class RedisMock:
@@ -47,19 +48,23 @@ def redis_strategy(redis):
@pytest.mark.authentication
class TestReadToken:
@pytest.mark.asyncio
async def test_missing_token(self, redis_strategy: RedisStrategy, user_manager):
async def test_missing_token(
self, redis_strategy: RedisStrategy[UserModel, IDType], user_manager
):
authenticated_user = await redis_strategy.read_token(None, user_manager)
assert authenticated_user is None
@pytest.mark.asyncio
async def test_invalid_token(self, redis_strategy: RedisStrategy, user_manager):
async def test_invalid_token(
self, redis_strategy: RedisStrategy[UserModel, IDType], user_manager
):
authenticated_user = await redis_strategy.read_token("TOKEN", user_manager)
assert authenticated_user is None
@pytest.mark.asyncio
async def test_valid_token_invalid_uuid(
self,
redis_strategy: RedisStrategy,
redis_strategy: RedisStrategy[UserModel, IDType],
redis: RedisMock,
user_manager,
):
@@ -70,7 +75,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_valid_token_not_existing_user(
self,
redis_strategy: RedisStrategy,
redis_strategy: RedisStrategy[UserModel, IDType],
redis: RedisMock,
user_manager,
):
@@ -81,7 +86,7 @@ class TestReadToken:
@pytest.mark.asyncio
async def test_valid_token(
self,
redis_strategy: RedisStrategy,
redis_strategy: RedisStrategy[UserModel, IDType],
redis: RedisMock,
user_manager,
user,
@@ -94,7 +99,9 @@ class TestReadToken:
@pytest.mark.authentication
@pytest.mark.asyncio
async def test_write_token(redis_strategy: RedisStrategy, redis: RedisMock, user):
async def test_write_token(
redis_strategy: RedisStrategy[UserModel, IDType], redis: RedisMock, user
):
token = await redis_strategy.write_token(user)
value = await redis.get(token)
@@ -103,7 +110,9 @@ async def test_write_token(redis_strategy: RedisStrategy, redis: RedisMock, user
@pytest.mark.authentication
@pytest.mark.asyncio
async def test_destroy_token(redis_strategy: RedisStrategy, redis: RedisMock, user):
async def test_destroy_token(
redis_strategy: RedisStrategy[UserModel, IDType], redis: RedisMock, user
):
await redis.set("TOKEN", str(user.id))
await redis_strategy.destroy_token("TOKEN", user)