mirror of
https://github.com/fastapi-users/fastapi-users.git
synced 2026-03-13 07:49:55 +08:00
Improve Strategy typing
This commit is contained in:
@@ -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
|
||||
):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user