mirror of
https://github.com/GNS3/gns3-server.git
synced 2026-08-27 12:30:13 +03:00
Fix addressing issue #2005 including tests
This commit is contained in:
parent
1d10bda22e
commit
4cf8b79692
@ -171,11 +171,13 @@ router.include_router(
|
||||
router.include_router(
|
||||
_llm_router,
|
||||
prefix="/access",
|
||||
dependencies=[Depends(get_current_active_user)],
|
||||
tags=["LLM Model Configurations"]
|
||||
)
|
||||
|
||||
router.include_router(
|
||||
_chat_router,
|
||||
prefix="/projects/{project_id}/chat",
|
||||
dependencies=[Depends(get_current_active_user)],
|
||||
tags=["Chat"]
|
||||
)
|
||||
|
||||
@ -47,14 +47,20 @@ async def get_user_from_token(
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
username = auth_service.get_username_from_token(token)
|
||||
user = await user_repo.get_user_by_username(username)
|
||||
token_data = auth_service.get_token_data(token)
|
||||
user = await user_repo.get_user_by_username(token_data.username)
|
||||
if user is None:
|
||||
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="Token has been revoked",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
return user
|
||||
|
||||
|
||||
@ -87,13 +93,18 @@ async def get_current_active_user_from_websocket(
|
||||
await websocket.accept()
|
||||
|
||||
try:
|
||||
username = auth_service.get_username_from_token(token)
|
||||
user = await user_repo.get_user_by_username(username)
|
||||
token_data = auth_service.get_token_data(token)
|
||||
user = await user_repo.get_user_by_username(token_data.username)
|
||||
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=f"Could not validate credentials for '{username}'"
|
||||
detail=f"Could not validate credentials for '{token_data.username}'"
|
||||
)
|
||||
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}'"
|
||||
)
|
||||
|
||||
# Super admin is always authorized
|
||||
|
||||
@ -65,7 +65,10 @@ async def login(
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
token = schemas.Token(access_token=auth_service.create_access_token(user.username), token_type="bearer")
|
||||
token = schemas.Token(
|
||||
access_token=auth_service.create_access_token(user.username, token_version=user.token_version),
|
||||
token_type="bearer"
|
||||
)
|
||||
return token
|
||||
|
||||
|
||||
@ -87,10 +90,25 @@ async def authenticate(
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
token = schemas.Token(access_token=auth_service.create_access_token(user.username), token_type="bearer")
|
||||
token = schemas.Token(
|
||||
access_token=auth_service.create_access_token(user.username, token_version=user.token_version),
|
||||
token_type="bearer"
|
||||
)
|
||||
return token
|
||||
|
||||
|
||||
@router.post("/logout", status_code=status.HTTP_204_NO_CONTENT)
|
||||
async def logout(
|
||||
current_user: schemas.User = Depends(get_current_active_user),
|
||||
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
||||
) -> None:
|
||||
"""
|
||||
Logout the current user by revoking all existing tokens.
|
||||
"""
|
||||
|
||||
await users_repo.logout_user(current_user.user_id)
|
||||
|
||||
|
||||
@router.get("/me", response_model=schemas.User)
|
||||
async def get_logged_in_user(current_user: schemas.User = Depends(get_current_active_user)) -> schemas.User:
|
||||
"""
|
||||
|
||||
@ -15,7 +15,7 @@
|
||||
# You should have received a copy of the GNU General Public License
|
||||
# along with this program. If not, see <http://www.gnu.org/licenses/>.
|
||||
|
||||
from sqlalchemy import Table, Boolean, Column, String, DateTime, ForeignKey, event
|
||||
from sqlalchemy import Table, Boolean, Column, Integer, String, DateTime, ForeignKey, event
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from .base import Base, BaseTable, generate_uuid, GUID
|
||||
@ -45,6 +45,7 @@ class User(BaseTable):
|
||||
full_name = Column(String)
|
||||
hashed_password = Column(String)
|
||||
last_login = Column(DateTime)
|
||||
token_version = Column(Integer, default=0, nullable=False, server_default="0")
|
||||
is_active = Column(Boolean, default=True)
|
||||
is_superadmin = Column(Boolean, default=False)
|
||||
groups = relationship("UserGroup", secondary=user_group_map, back_populates="users")
|
||||
|
||||
@ -113,6 +113,18 @@ class UsersRepository(BaseRepository):
|
||||
await self._db_session.refresh(user_db) # force refresh of updated_at value
|
||||
return user_db
|
||||
|
||||
async def logout_user(self, user_id: UUID) -> None:
|
||||
"""
|
||||
Increment token_version to invalidate all existing tokens for the user.
|
||||
"""
|
||||
|
||||
query = update(models.User).\
|
||||
where(models.User.user_id == user_id).\
|
||||
values(token_version=models.User.token_version + 1)
|
||||
|
||||
await self._db_session.execute(query)
|
||||
await self._db_session.commit()
|
||||
|
||||
async def delete_user(self, user_id: UUID) -> bool:
|
||||
"""
|
||||
Delete a user.
|
||||
|
||||
@ -0,0 +1,28 @@
|
||||
"""add token_version to users
|
||||
|
||||
Revision ID: 20260405_add_token_version_to_users
|
||||
Revises: 20260303_create_llm_model_configs
|
||||
Create Date: 2026-04-05
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '20260405_add_token_version_to_users'
|
||||
down_revision = 'ec4b7b198555'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = inspect(conn)
|
||||
columns = [col['name'] for col in inspector.get_columns('users')]
|
||||
if 'token_version' not in columns:
|
||||
op.add_column('users', sa.Column('token_version', sa.Integer(), nullable=False, server_default='0'))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column('users', 'token_version')
|
||||
@ -27,3 +27,4 @@ class Token(BaseModel):
|
||||
class TokenData(BaseModel):
|
||||
|
||||
username: Optional[str] = None
|
||||
token_version: int = 0
|
||||
|
||||
@ -45,12 +45,12 @@ class AuthService:
|
||||
|
||||
return bcrypt.checkpw(password=password.encode('utf-8'), hashed_password=hashed_password.encode('utf-8'))
|
||||
|
||||
def create_access_token(self, username, secret_key: str = None, expires_in: int = 0) -> str:
|
||||
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
|
||||
expire = datetime.now(timezone.utc) + timedelta(minutes=expires_in)
|
||||
to_encode = {"sub": username, "exp": expire}
|
||||
to_encode = {"sub": username, "exp": expire, "ver": token_version}
|
||||
if secret_key is None:
|
||||
secret_key = Config.instance().settings.Controller.jwt_secret_key
|
||||
if secret_key is None:
|
||||
@ -61,7 +61,7 @@ class AuthService:
|
||||
encoded_jwt = jwt.encode({"alg": algorithm}, to_encode, key)
|
||||
return encoded_jwt
|
||||
|
||||
def get_username_from_token(self, token: str, secret_key: str = None) -> Optional[str]:
|
||||
def get_token_data(self, token: str, secret_key: str = None) -> TokenData:
|
||||
|
||||
credentials_exception = HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@ -80,7 +80,11 @@ class AuthService:
|
||||
username: str = payload.claims.get("sub")
|
||||
if username is None:
|
||||
raise credentials_exception
|
||||
token_data = TokenData(username=username)
|
||||
token_version: int = payload.claims.get("ver", 0)
|
||||
token_data = TokenData(username=username, token_version=token_version)
|
||||
except (JoseError, ValidationError, ValueError):
|
||||
raise credentials_exception
|
||||
return token_data.username
|
||||
return token_data
|
||||
|
||||
def get_username_from_token(self, token: str, secret_key: str = None) -> Optional[str]:
|
||||
return self.get_token_data(token, secret_key).username
|
||||
|
||||
@ -247,6 +247,7 @@ class TestUserLogin:
|
||||
key = OctKey.import_key(jwt_secret)
|
||||
payload = jwt.decode(token, key, algorithms=["HS256"])
|
||||
assert "sub" in payload.claims
|
||||
assert "ver" in payload.claims
|
||||
username = payload.claims.get("sub")
|
||||
assert username == test_user.username
|
||||
|
||||
@ -373,6 +374,117 @@ class TestUserMe:
|
||||
assert response.status_code == status_code
|
||||
|
||||
|
||||
class TestLogout:
|
||||
|
||||
async def test_logout_returns_no_content(
|
||||
self,
|
||||
app: FastAPI,
|
||||
unauthorized_client: AsyncClient,
|
||||
test_user: User,
|
||||
) -> None:
|
||||
|
||||
# login to get a fresh token that includes token_version
|
||||
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
|
||||
token = response.json()["access_token"]
|
||||
|
||||
response = await unauthorized_client.post(
|
||||
app.url_path_for("logout"),
|
||||
headers={"Authorization": f"Bearer {token}"}
|
||||
)
|
||||
assert response.status_code == status.HTTP_204_NO_CONTENT
|
||||
|
||||
async def test_token_is_rejected_after_logout(
|
||||
self,
|
||||
app: FastAPI,
|
||||
unauthorized_client: AsyncClient,
|
||||
test_user: User,
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
|
||||
# login and get a token
|
||||
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
|
||||
token = response.json()["access_token"]
|
||||
|
||||
# logout — increments token_version
|
||||
await unauthorized_client.post(
|
||||
app.url_path_for("logout"),
|
||||
headers={"Authorization": f"Bearer {token}"}
|
||||
)
|
||||
|
||||
# old token must now be rejected
|
||||
response = await unauthorized_client.get(
|
||||
app.url_path_for("get_logged_in_user"),
|
||||
headers={"Authorization": f"Bearer {token}"}
|
||||
)
|
||||
assert response.status_code == status.HTTP_401_UNAUTHORIZED
|
||||
assert response.json()["message"] == "Token has been revoked"
|
||||
|
||||
async def test_new_token_works_after_logout_and_relogin(
|
||||
self,
|
||||
app: FastAPI,
|
||||
unauthorized_client: AsyncClient,
|
||||
test_user: User,
|
||||
) -> None:
|
||||
|
||||
credentials = {"username": test_user.username, "password": "user1_password"}
|
||||
|
||||
# login, then logout
|
||||
response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
|
||||
old_token = response.json()["access_token"]
|
||||
await unauthorized_client.post(
|
||||
app.url_path_for("logout"),
|
||||
headers={"Authorization": f"Bearer {old_token}"}
|
||||
)
|
||||
|
||||
# login again to get a fresh token
|
||||
response = await unauthorized_client.post(app.url_path_for("authenticate"), json=credentials)
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
new_token = response.json()["access_token"]
|
||||
|
||||
# new token must work
|
||||
response = await unauthorized_client.get(
|
||||
app.url_path_for("get_logged_in_user"),
|
||||
headers={"Authorization": f"Bearer {new_token}"}
|
||||
)
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
assert response.json()["username"] == test_user.username
|
||||
|
||||
async def test_stale_version_token_is_rejected(
|
||||
self,
|
||||
app: FastAPI,
|
||||
unauthorized_client: AsyncClient,
|
||||
test_user: User,
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
|
||||
# craft a token with ver=0 while the user's token_version is already higher
|
||||
user_repo = UsersRepository(db_session)
|
||||
user_in_db = await user_repo.get_user_by_username(test_user.username)
|
||||
|
||||
# force token_version ahead so any ver=0 token is stale
|
||||
await user_repo.logout_user(user_in_db.user_id)
|
||||
|
||||
stale_token = auth_service.create_access_token(test_user.username, token_version=0)
|
||||
response = await unauthorized_client.get(
|
||||
app.url_path_for("get_logged_in_user"),
|
||||
headers={"Authorization": f"Bearer {stale_token}"}
|
||||
)
|
||||
assert response.status_code == status.HTTP_401_UNAUTHORIZED
|
||||
|
||||
async def test_logout_without_token_returns_unauthorized(
|
||||
self,
|
||||
app: FastAPI,
|
||||
unauthorized_client: AsyncClient,
|
||||
) -> None:
|
||||
|
||||
response = await unauthorized_client.post(app.url_path_for("logout"))
|
||||
assert response.status_code == status.HTTP_401_UNAUTHORIZED
|
||||
|
||||
|
||||
class TestSuperAdmin:
|
||||
|
||||
async def test_super_admin_exists(
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user