Add stateless JWT refresh token mechanism

- New config: Controller.jwt_refresh_token_expire_minutes (default 30 days)
- New endpoint: POST /v3/access/users/refresh (public, unauthenticated)
- Login/authenticate responses now include refresh_token
- AuthService: _create_token helper, create_refresh_token, get_token_data
  now parses type claim (token_use) for token classification
- Security: refresh tokens rejected on HTTP + WebSocket access paths;
  /refresh strictly requires type=='refresh'
- Logout works for free via existing token_version mechanism
- Tests: 9 new TestRefreshToken cases, all passing; 34 existing tests
  still pass (no regressions)
This commit is contained in:
YueGuobin 2026-06-23 22:39:11 +08:00
parent 4ffab7abee
commit 0e8d0cb87b
No known key found for this signature in database
7 changed files with 262 additions and 8 deletions

View File

@ -35,6 +35,17 @@ log = logging.getLogger(__name__)
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/v3/access/users/login", auto_error=False)
def _reject_refresh_token(token_data) -> None:
"""Reject tokens with type == 'refresh' — they must not grant API access."""
if token_data.token_use == "refresh":
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Refresh tokens cannot be used for API access",
headers={"WWW-Authenticate": "Bearer"},
)
async def get_user_from_token(
bearer_token: str = Depends(oauth2_scheme),
user_repo: UsersRepository = Depends(get_repository(UsersRepository)),
@ -82,6 +93,7 @@ async def get_user_from_token(
# JWT authentication
token_data = auth_service.get_token_data(token)
_reject_refresh_token(token_data)
user = await user_repo.get_user_by_username(token_data.username)
if user is None:
raise HTTPException(
@ -137,6 +149,7 @@ async def get_current_active_user_from_websocket(
try:
token_data = auth_service.get_token_data(token)
_reject_refresh_token(token_data)
user = await user_repo.get_user_by_username(token_data.username)
if user is None:

View File

@ -67,7 +67,8 @@ async def login(
token = schemas.Token(
access_token=auth_service.create_access_token(user.username, token_version=user.token_version),
token_type="bearer"
token_type="bearer",
refresh_token=auth_service.create_refresh_token(user.username, token_version=user.token_version),
)
return token
@ -92,11 +93,55 @@ async def authenticate(
token = schemas.Token(
access_token=auth_service.create_access_token(user.username, token_version=user.token_version),
token_type="bearer"
token_type="bearer",
refresh_token=auth_service.create_refresh_token(user.username, token_version=user.token_version),
)
return token
@router.post("/refresh", response_model=schemas.Token)
async def refresh_access_token(
request: schemas.RefreshTokenRequest,
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> schemas.Token:
"""
Exchange a refresh token for a new access token.
Public endpoint the refresh token itself proves identity. Respects the
user's token_version, so logout (which increments it) invalidates all
outstanding refresh tokens. Refresh tokens are stateless JWTs with a
longer expiry (default 30 days). Stolen tokens remain valid until their
`exp` or until logout no replay protection without a server-side table.
"""
token_data = auth_service.get_token_data(request.refresh_token)
if token_data.token_use != "refresh":
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid refresh token",
headers={"WWW-Authenticate": "Bearer"},
)
user = await users_repo.get_user_by_username(token_data.username)
if user is None or not user.is_active:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Could not validate credentials",
headers={"WWW-Authenticate": "Bearer"},
)
if token_data.token_version != user.token_version:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=f"Token has been revoked for '{token_data.username}'",
headers={"WWW-Authenticate": "Bearer"},
)
return schemas.Token(
access_token=auth_service.create_access_token(user.username, token_version=user.token_version),
token_type="bearer",
refresh_token=auth_service.create_refresh_token(user.username, token_version=user.token_version),
)
@router.post("/logout", status_code=status.HTTP_204_NO_CONTENT)
async def logout(
current_user: schemas.User = Depends(get_current_active_user),

View File

@ -57,7 +57,7 @@ except ImportError:
from .controller.rbac import RoleCreate, RoleUpdate, Role, Privilege, ACECreate, ACEUpdate, ACE
from .controller.pools import Resource, ResourceCreate, ResourcePoolCreate, ResourcePoolUpdate, ResourcePool
from .controller.tokens import Token, ApiKeyCreate
from .controller.tokens import Token, ApiKeyCreate, RefreshTokenRequest
from .controller.snapshots import SnapshotCreate, Snapshot
from .controller.iou_license import IOULicense
from .controller.capabilities import Capabilities

View File

@ -35,6 +35,7 @@ class ControllerSettings(BaseModel):
jwt_secret_key: str = None
jwt_algorithm: str = "HS256"
jwt_access_token_expire_minutes: int = 1440 # 24 hours
jwt_refresh_token_expire_minutes: int = 43200 # 30 days
default_admin_username: str = "admin"
default_admin_password: SecretStr = SecretStr("admin")
model_config = ConfigDict(validate_assignment=True, str_strip_whitespace=True)

View File

@ -22,12 +22,20 @@ class Token(BaseModel):
access_token: str
token_type: str
refresh_token: Optional[str] = None
class TokenData(BaseModel):
username: Optional[str] = None
token_version: int = 0
token_use: str = "access"
class RefreshTokenRequest(BaseModel):
"""Schema for requesting a token refresh."""
refresh_token: str
class ApiKeyCreate(BaseModel):

View File

@ -46,12 +46,11 @@ class AuthService:
return bcrypt.checkpw(password=password.encode('utf-8'), hashed_password=hashed_password.encode('utf-8'))
def create_access_token(self, username, token_version: int = 0, secret_key: str = None, expires_in: int = 0) -> str:
def _create_token(self, username, token_version, token_type, expires_in, secret_key=None) -> str:
"""Shared helper to create any kind of signed JWT token."""
if not expires_in:
expires_in = Config.instance().settings.Controller.jwt_access_token_expire_minutes
expire = datetime.now(timezone.utc) + timedelta(minutes=expires_in)
to_encode = {"sub": username, "exp": expire, "ver": token_version}
to_encode = {"sub": username, "exp": expire, "ver": token_version, "type": token_type}
if secret_key is None:
secret_key = Config.instance().settings.Controller.jwt_secret_key
if secret_key is None:
@ -62,6 +61,18 @@ class AuthService:
encoded_jwt = jwt.encode({"alg": algorithm}, to_encode, key)
return encoded_jwt
def create_access_token(self, username, token_version: int = 0, secret_key: str = None, expires_in: int = 0) -> str:
if not expires_in:
expires_in = Config.instance().settings.Controller.jwt_access_token_expire_minutes
return self._create_token(username, token_version, "access", expires_in, secret_key)
def create_refresh_token(self, username, token_version: int = 0, secret_key: str = None, expires_in: int = 0) -> str:
if not expires_in:
expires_in = Config.instance().settings.Controller.jwt_refresh_token_expire_minutes
return self._create_token(username, token_version, "refresh", expires_in, secret_key)
def get_token_data(self, token: str, secret_key: str = None) -> TokenData:
credentials_exception = HTTPException(
@ -86,7 +97,8 @@ class AuthService:
if token_exp and time.time() > token_exp:
raise credentials_exception
token_version: int = payload.claims.get("ver", 0)
token_data = TokenData(username=username, token_version=token_version)
token_use: str = payload.claims.get("type", "access")
token_data = TokenData(username=username, token_version=token_version, token_use=token_use)
except (JoseError, ValidationError, ValueError):
raise credentials_exception
return token_data

View File

@ -255,6 +255,10 @@ class TestUserLogin:
assert "token_type" in response.json()
assert response.json().get("token_type") == "bearer"
# check that refresh token is returned
assert "refresh_token" in response.json()
assert response.json().get("refresh_token") is not None
@pytest.mark.parametrize(
"username, password, status_code",
(
@ -311,6 +315,7 @@ class TestUnauthorizedUser:
response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
assert response.status_code == status.HTTP_200_OK
assert response.json().get("access_token")
assert response.json().get("refresh_token") is not None
token = response.json().get("access_token")
response = await unauthorized_client.get(app.url_path_for("statistics"), params={"token": token})
@ -485,6 +490,176 @@ class TestLogout:
assert response.status_code == status.HTTP_401_UNAUTHORIZED
class TestRefreshToken:
async def test_login_returns_refresh_token(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
credentials = {"username": test_user.username, "password": "user1_password"}
response = await unauthorized_client.post(app.url_path_for("login"),
data=credentials,
headers={"content-type": "application/x-www-form-urlencoded"})
assert response.status_code == status.HTTP_200_OK
assert "refresh_token" in response.json()
assert response.json().get("refresh_token") is not None
async def test_authenticate_returns_refresh_token(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
credentials = {"username": test_user.username, "password": "user1_password"}
response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
assert response.status_code == status.HTTP_200_OK
assert "refresh_token" in response.json()
assert response.json().get("refresh_token") is not None
async def test_refresh_endpoint_returns_new_tokens(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
# authenticate to get a refresh token
credentials = {"username": test_user.username, "password": "user1_password"}
auth_response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
refresh_token = auth_response.json()["refresh_token"]
# use the refresh token at /refresh
response = await unauthorized_client.post(
app.url_path_for("refresh_access_token"),
json={"refresh_token": refresh_token},
)
assert response.status_code == status.HTTP_200_OK
assert "access_token" in response.json()
assert response.json().get("token_type") == "bearer"
assert "refresh_token" in response.json()
assert response.json().get("refresh_token") is not None
async def test_new_access_token_from_refresh_works_on_protected_route(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
credentials = {"username": test_user.username, "password": "user1_password"}
auth_response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
refresh_token = auth_response.json()["refresh_token"]
# refresh to get a new access token
refresh_response = await unauthorized_client.post(
app.url_path_for("refresh_access_token"),
json={"refresh_token": refresh_token},
)
new_access_token = refresh_response.json()["access_token"]
# the new access token must work on a protected route
response = await unauthorized_client.get(
app.url_path_for("get_logged_in_user"),
headers={"Authorization": f"Bearer {new_access_token}"},
)
assert response.status_code == status.HTTP_200_OK
assert response.json()["username"] == test_user.username
async def test_refresh_with_stale_token_after_logout(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
# authenticate and get a refresh token
credentials = {"username": test_user.username, "password": "user1_password"}
auth_response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
refresh_token = auth_response.json()["refresh_token"]
# logout — bumps token_version, invalidating the refresh token
access_token = auth_response.json()["access_token"]
await unauthorized_client.post(
app.url_path_for("logout"),
headers={"Authorization": f"Bearer {access_token}"}
)
# /refresh must now reject the stale refresh token
response = await unauthorized_client.post(
app.url_path_for("refresh_access_token"),
json={"refresh_token": refresh_token},
)
assert response.status_code == status.HTTP_401_UNAUTHORIZED
async def test_refresh_rejects_access_token(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
# an access token presented at /refresh must be rejected
credentials = {"username": test_user.username, "password": "user1_password"}
auth_response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
access_token = auth_response.json()["access_token"]
response = await unauthorized_client.post(
app.url_path_for("refresh_access_token"),
json={"refresh_token": access_token},
)
assert response.status_code == status.HTTP_401_UNAUTHORIZED
async def test_refresh_rejects_expired_token(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
# a refresh token with an already-expired timestamp
expired_refresh = auth_service.create_refresh_token(test_user.username, expires_in=-1)
response = await unauthorized_client.post(
app.url_path_for("refresh_access_token"),
json={"refresh_token": expired_refresh},
)
assert response.status_code == status.HTTP_401_UNAUTHORIZED
async def test_refresh_rejects_invalid_token(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
) -> None:
response = await unauthorized_client.post(
app.url_path_for("refresh_access_token"),
json={"refresh_token": "not-a-valid-token"},
)
assert response.status_code == status.HTTP_401_UNAUTHORIZED
async def test_refresh_token_rejected_as_bearer(
self,
app: FastAPI,
unauthorized_client: AsyncClient,
test_user: User,
) -> None:
# a refresh token used as a bearer access token must be rejected
credentials = {"username": test_user.username, "password": "user1_password"}
auth_response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
refresh_token = auth_response.json()["refresh_token"]
response = await unauthorized_client.get(
app.url_path_for("get_logged_in_user"),
headers={"Authorization": f"Bearer {refresh_token}"},
)
assert response.status_code == status.HTTP_401_UNAUTHORIZED
class TestSuperAdmin:
async def test_super_admin_exists(