diff --git a/fastapi_users/authentication/strategy/db/strategy.py b/fastapi_users/authentication/strategy/db/strategy.py index e3cb93df..c9e40ebc 100644 --- a/fastapi_users/authentication/strategy/db/strategy.py +++ b/fastapi_users/authentication/strategy/db/strategy.py @@ -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 ): diff --git a/fastapi_users/authentication/strategy/jwt.py b/fastapi_users/authentication/strategy/jwt.py index a2fbf58e..4bad548d 100644 --- a/fastapi_users/authentication/strategy/jwt.py +++ b/fastapi_users/authentication/strategy/jwt.py @@ -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, diff --git a/fastapi_users/authentication/strategy/redis.py b/fastapi_users/authentication/strategy/redis.py index d6527892..616ac7e5 100644 --- a/fastapi_users/authentication/strategy/redis.py +++ b/fastapi_users/authentication/strategy/redis.py @@ -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 diff --git a/tests/test_authentication_strategy_db.py b/tests/test_authentication_strategy_db.py index 00b7047f..201f8949 100644 --- a/tests/test_authentication_strategy_db.py +++ b/tests/test_authentication_strategy_db.py @@ -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, ): diff --git a/tests/test_authentication_strategy_jwt.py b/tests/test_authentication_strategy_jwt.py index 08e5fd7a..45aea5a5 100644 --- a/tests/test_authentication_strategy_jwt.py +++ b/tests/test_authentication_strategy_jwt.py @@ -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) diff --git a/tests/test_authentication_strategy_redis.py b/tests/test_authentication_strategy_redis.py index cdddde1a..61fe250f 100644 --- a/tests/test_authentication_strategy_redis.py +++ b/tests/test_authentication_strategy_redis.py @@ -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)