feat(auth): migrate to Keycloak metadata auth
This commit is contained in:
@@ -0,0 +1,48 @@
|
|||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from app.auth.keycloak_dependencies import get_current_keycloak_payload
|
||||||
|
from app.auth.metadata_dependencies import get_current_metadata_user
|
||||||
|
from app.auth.project_dependencies import (
|
||||||
|
ProjectContext,
|
||||||
|
get_project_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
|
class AgentAuthContextResponse(BaseModel):
|
||||||
|
user_id: str
|
||||||
|
keycloak_sub: str
|
||||||
|
username: str
|
||||||
|
role: str
|
||||||
|
is_superuser: bool
|
||||||
|
project_id: str
|
||||||
|
project_role: str
|
||||||
|
token_expires_at: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/agent/auth/context", response_model=AgentAuthContextResponse)
|
||||||
|
async def get_agent_auth_context(
|
||||||
|
ctx: ProjectContext = Depends(get_project_context),
|
||||||
|
current_user=Depends(get_current_metadata_user),
|
||||||
|
keycloak_payload: dict = Depends(get_current_keycloak_payload),
|
||||||
|
) -> AgentAuthContextResponse:
|
||||||
|
exp = keycloak_payload.get("exp")
|
||||||
|
token_expires_at = (
|
||||||
|
datetime.fromtimestamp(exp, tz=timezone.utc).isoformat()
|
||||||
|
if isinstance(exp, int)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
return AgentAuthContextResponse(
|
||||||
|
user_id=str(current_user.id),
|
||||||
|
keycloak_sub=str(current_user.keycloak_id),
|
||||||
|
username=current_user.username,
|
||||||
|
role=current_user.role,
|
||||||
|
is_superuser=current_user.is_superuser,
|
||||||
|
project_id=str(ctx.project_id),
|
||||||
|
project_role=ctx.project_role,
|
||||||
|
token_expires_at=token_expires_at,
|
||||||
|
)
|
||||||
@@ -1,190 +0,0 @@
|
|||||||
from typing import Annotated
|
|
||||||
from datetime import timedelta
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, status
|
|
||||||
from fastapi.security import OAuth2PasswordRequestForm
|
|
||||||
from app.core.config import settings
|
|
||||||
from app.core.security import create_access_token, create_refresh_token, verify_password
|
|
||||||
from app.domain.schemas.user import UserCreate, UserResponse, UserLogin, Token
|
|
||||||
from app.infra.db.metadb.repositories.user_repository import UserRepository
|
|
||||||
from app.auth.dependencies import get_user_repository, get_current_active_user
|
|
||||||
from app.domain.schemas.user import UserInDB
|
|
||||||
import logging
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
router = APIRouter()
|
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
|
||||||
"/register", response_model=UserResponse, status_code=status.HTTP_201_CREATED
|
|
||||||
)
|
|
||||||
async def register(
|
|
||||||
user_data: UserCreate, user_repo: UserRepository = Depends(get_user_repository)
|
|
||||||
) -> UserResponse:
|
|
||||||
"""
|
|
||||||
用户注册
|
|
||||||
|
|
||||||
创建新用户账号
|
|
||||||
"""
|
|
||||||
# 检查用户名和邮箱是否已存在
|
|
||||||
if await user_repo.user_exists(username=user_data.username):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
|
||||||
detail="Username already registered",
|
|
||||||
)
|
|
||||||
|
|
||||||
if await user_repo.user_exists(email=user_data.email):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Email already registered"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 创建用户
|
|
||||||
try:
|
|
||||||
user = await user_repo.create_user(user_data)
|
|
||||||
if not user:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
||||||
detail="Failed to create user",
|
|
||||||
)
|
|
||||||
return UserResponse.model_validate(user)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error during user registration: {e}")
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
||||||
detail="Registration failed",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/login", response_model=Token)
|
|
||||||
async def login(
|
|
||||||
form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> Token:
|
|
||||||
"""
|
|
||||||
用户登录(OAuth2 标准格式)
|
|
||||||
|
|
||||||
返回 JWT Access Token 和 Refresh Token
|
|
||||||
"""
|
|
||||||
# 验证用户(支持用户名或邮箱登录)
|
|
||||||
user = await user_repo.get_user_by_username(form_data.username)
|
|
||||||
if not user:
|
|
||||||
# 尝试用邮箱登录
|
|
||||||
user = await user_repo.get_user_by_email(form_data.username)
|
|
||||||
|
|
||||||
if not user or not verify_password(form_data.password, user.hashed_password):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Incorrect username or password",
|
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
|
||||||
)
|
|
||||||
|
|
||||||
if not user.is_active:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user account"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 生成 Token
|
|
||||||
access_token = create_access_token(subject=user.username)
|
|
||||||
refresh_token = create_refresh_token(subject=user.username)
|
|
||||||
|
|
||||||
return Token(
|
|
||||||
access_token=access_token,
|
|
||||||
refresh_token=refresh_token,
|
|
||||||
token_type="bearer",
|
|
||||||
expires_in=settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/login/simple", response_model=Token)
|
|
||||||
async def login_simple(
|
|
||||||
username: str,
|
|
||||||
password: str,
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> Token:
|
|
||||||
"""
|
|
||||||
简化版登录接口(保持向后兼容)
|
|
||||||
|
|
||||||
直接使用 username 和 password 参数
|
|
||||||
"""
|
|
||||||
# 验证用户
|
|
||||||
user = await user_repo.get_user_by_username(username)
|
|
||||||
if not user:
|
|
||||||
user = await user_repo.get_user_by_email(username)
|
|
||||||
|
|
||||||
if not user or not verify_password(password, user.hashed_password):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Incorrect username or password",
|
|
||||||
)
|
|
||||||
|
|
||||||
if not user.is_active:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user account"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 生成 Token
|
|
||||||
access_token = create_access_token(subject=user.username)
|
|
||||||
refresh_token = create_refresh_token(subject=user.username)
|
|
||||||
|
|
||||||
return Token(
|
|
||||||
access_token=access_token,
|
|
||||||
refresh_token=refresh_token,
|
|
||||||
token_type="bearer",
|
|
||||||
expires_in=settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@router.get("/me", response_model=UserResponse)
|
|
||||||
async def get_current_user_info(
|
|
||||||
current_user: UserInDB = Depends(get_current_active_user),
|
|
||||||
) -> UserResponse:
|
|
||||||
"""
|
|
||||||
获取当前登录用户信息
|
|
||||||
"""
|
|
||||||
return UserResponse.model_validate(current_user)
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/refresh", response_model=Token)
|
|
||||||
async def refresh_token(
|
|
||||||
refresh_token: str, user_repo: UserRepository = Depends(get_user_repository)
|
|
||||||
) -> Token:
|
|
||||||
"""
|
|
||||||
刷新 Access Token
|
|
||||||
|
|
||||||
使用 Refresh Token 获取新的 Access Token
|
|
||||||
"""
|
|
||||||
from jose import jwt, JWTError
|
|
||||||
|
|
||||||
credentials_exception = HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Could not validate refresh token",
|
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
payload = jwt.decode(
|
|
||||||
refresh_token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]
|
|
||||||
)
|
|
||||||
username: str = payload.get("sub")
|
|
||||||
token_type: str = payload.get("type")
|
|
||||||
|
|
||||||
if username is None or token_type != "refresh":
|
|
||||||
raise credentials_exception
|
|
||||||
|
|
||||||
except JWTError:
|
|
||||||
raise credentials_exception
|
|
||||||
|
|
||||||
# 验证用户仍然存在且激活
|
|
||||||
user = await user_repo.get_user_by_username(username)
|
|
||||||
if not user or not user.is_active:
|
|
||||||
raise credentials_exception
|
|
||||||
|
|
||||||
# 生成新的 Access Token
|
|
||||||
new_access_token = create_access_token(subject=user.username)
|
|
||||||
|
|
||||||
return Token(
|
|
||||||
access_token=new_access_token,
|
|
||||||
refresh_token=refresh_token, # 保持原 refresh token
|
|
||||||
token_type="bearer",
|
|
||||||
expires_in=settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
|
|
||||||
)
|
|
||||||
@@ -10,7 +10,7 @@ from app.services.tjnetwork import (
|
|||||||
get_network_node_coords,
|
get_network_node_coords,
|
||||||
get_node_coord,
|
get_node_coord,
|
||||||
)
|
)
|
||||||
from app.auth.dependencies import get_current_user as verify_token
|
from app.auth.metadata_dependencies import get_current_metadata_user
|
||||||
from app.infra.cache.redis_client import redis_client, encode_datetime, decode_datetime
|
from app.infra.cache.redis_client import redis_client, encode_datetime, decode_datetime
|
||||||
import msgpack
|
import msgpack
|
||||||
|
|
||||||
@@ -64,7 +64,7 @@ async def fastapi_get_network_in_extent(
|
|||||||
|
|
||||||
@router.get(
|
@router.get(
|
||||||
"/getnetworkgeometries/",
|
"/getnetworkgeometries/",
|
||||||
dependencies=[Depends(verify_token)],
|
dependencies=[Depends(get_current_metadata_user)],
|
||||||
summary="获取完整网络几何信息",
|
summary="获取完整网络几何信息",
|
||||||
description="获取整个水网的所有节点、管线和SCADA点的几何信息(需要身份验证)"
|
description="获取整个水网的所有节点、管线和SCADA点的几何信息(需要身份验证)"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,215 +0,0 @@
|
|||||||
"""
|
|
||||||
用户管理 API 接口
|
|
||||||
|
|
||||||
演示权限控制的使用
|
|
||||||
"""
|
|
||||||
|
|
||||||
from typing import List
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, status, Path, Query
|
|
||||||
from app.domain.schemas.user import UserResponse, UserUpdate, UserCreate
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from app.domain.schemas.user import UserInDB
|
|
||||||
from app.infra.db.metadb.repositories.user_repository import UserRepository
|
|
||||||
from app.auth.dependencies import get_user_repository, get_current_active_user
|
|
||||||
from app.auth.permissions import get_current_admin, require_role, check_resource_owner
|
|
||||||
|
|
||||||
router = APIRouter()
|
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
|
||||||
"/",
|
|
||||||
summary="列出所有用户",
|
|
||||||
description="获取用户列表(仅管理员)",
|
|
||||||
response_model=List[UserResponse],
|
|
||||||
)
|
|
||||||
async def list_users(
|
|
||||||
skip: int = Query(0, ge=0, description="跳过的用户数"),
|
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大用户数"),
|
|
||||||
current_user: UserInDB = Depends(require_role(UserRole.ADMIN)),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> List[UserResponse]:
|
|
||||||
"""
|
|
||||||
获取用户列表
|
|
||||||
|
|
||||||
获取系统中所有的用户信息(需要管理员权限)
|
|
||||||
"""
|
|
||||||
users = await user_repo.get_all_users(skip=skip, limit=limit)
|
|
||||||
return [UserResponse.model_validate(user) for user in users]
|
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
|
||||||
"/{user_id}",
|
|
||||||
summary="获取用户详情",
|
|
||||||
description="获取指定用户的详细信息",
|
|
||||||
response_model=UserResponse,
|
|
||||||
)
|
|
||||||
async def get_user(
|
|
||||||
user_id: int = Path(..., gt=0, description="用户ID"),
|
|
||||||
current_user: UserInDB = Depends(get_current_active_user),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> UserResponse:
|
|
||||||
"""
|
|
||||||
获取用户详情
|
|
||||||
|
|
||||||
管理员可查看所有用户,普通用户只能查看自己
|
|
||||||
"""
|
|
||||||
# 检查权限
|
|
||||||
if not check_resource_owner(user_id, current_user):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail="You don't have permission to view this user",
|
|
||||||
)
|
|
||||||
|
|
||||||
user = await user_repo.get_user_by_id(user_id)
|
|
||||||
if not user:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
|
|
||||||
)
|
|
||||||
|
|
||||||
return UserResponse.model_validate(user)
|
|
||||||
|
|
||||||
|
|
||||||
@router.put(
|
|
||||||
"/{user_id}",
|
|
||||||
summary="更新用户信息",
|
|
||||||
description="更新指定用户的信息",
|
|
||||||
response_model=UserResponse,
|
|
||||||
)
|
|
||||||
async def update_user(
|
|
||||||
user_id: int = Path(..., gt=0, description="用户ID"),
|
|
||||||
user_update: UserUpdate = None,
|
|
||||||
current_user: UserInDB = Depends(get_current_active_user),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> UserResponse:
|
|
||||||
"""
|
|
||||||
更新用户信息
|
|
||||||
|
|
||||||
管理员可更新所有用户,普通用户只能更新自己(且不能修改角色)
|
|
||||||
"""
|
|
||||||
# 检查用户是否存在
|
|
||||||
target_user = await user_repo.get_user_by_id(user_id)
|
|
||||||
if not target_user:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 权限检查
|
|
||||||
is_owner = current_user.id == user_id
|
|
||||||
is_admin = UserRole(current_user.role).has_permission(UserRole.ADMIN)
|
|
||||||
|
|
||||||
if not is_owner and not is_admin:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail="You don't have permission to update this user",
|
|
||||||
)
|
|
||||||
|
|
||||||
# 非管理员不能修改角色和激活状态
|
|
||||||
if not is_admin:
|
|
||||||
if user_update.role is not None:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail="Only admins can change user roles",
|
|
||||||
)
|
|
||||||
if user_update.is_active is not None:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail="Only admins can change user active status",
|
|
||||||
)
|
|
||||||
|
|
||||||
# 更新用户
|
|
||||||
updated_user = await user_repo.update_user(user_id, user_update)
|
|
||||||
if not updated_user:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
||||||
detail="Failed to update user",
|
|
||||||
)
|
|
||||||
|
|
||||||
return UserResponse.model_validate(updated_user)
|
|
||||||
|
|
||||||
|
|
||||||
@router.delete("/{user_id}", summary="删除用户", description="删除指定用户(仅管理员)")
|
|
||||||
async def delete_user(
|
|
||||||
user_id: int = Path(..., gt=0, description="用户ID"),
|
|
||||||
current_user: UserInDB = Depends(get_current_admin),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> dict:
|
|
||||||
"""
|
|
||||||
删除用户
|
|
||||||
|
|
||||||
删除指定用户(需要管理员权限,不能删除自己)
|
|
||||||
"""
|
|
||||||
# 不能删除自己
|
|
||||||
if current_user.id == user_id:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
|
||||||
detail="You cannot delete your own account",
|
|
||||||
)
|
|
||||||
|
|
||||||
success = await user_repo.delete_user(user_id)
|
|
||||||
if not success:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
|
|
||||||
)
|
|
||||||
|
|
||||||
return {"message": "User deleted successfully"}
|
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
|
||||||
"/{user_id}/activate",
|
|
||||||
summary="激活用户",
|
|
||||||
description="激活指定用户账户(仅管理员)",
|
|
||||||
response_model=UserResponse,
|
|
||||||
)
|
|
||||||
async def activate_user(
|
|
||||||
user_id: int = Path(..., gt=0, description="用户ID"),
|
|
||||||
current_user: UserInDB = Depends(get_current_admin),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> UserResponse:
|
|
||||||
"""
|
|
||||||
激活用户
|
|
||||||
|
|
||||||
激活指定用户的账户(需要管理员权限)
|
|
||||||
"""
|
|
||||||
user_update = UserUpdate(is_active=True)
|
|
||||||
updated_user = await user_repo.update_user(user_id, user_update)
|
|
||||||
|
|
||||||
if not updated_user:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
|
|
||||||
)
|
|
||||||
|
|
||||||
return UserResponse.model_validate(updated_user)
|
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
|
||||||
"/{user_id}/deactivate",
|
|
||||||
summary="停用用户",
|
|
||||||
description="停用指定用户账户(仅管理员)",
|
|
||||||
response_model=UserResponse,
|
|
||||||
)
|
|
||||||
async def deactivate_user(
|
|
||||||
user_id: int = Path(..., gt=0, description="用户ID"),
|
|
||||||
current_user: UserInDB = Depends(get_current_admin),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> UserResponse:
|
|
||||||
"""
|
|
||||||
停用用户
|
|
||||||
|
|
||||||
停用指定用户的账户(需要管理员权限,不能停用自己)
|
|
||||||
"""
|
|
||||||
# 不能停用自己
|
|
||||||
if current_user.id == user_id:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
|
||||||
detail="You cannot deactivate your own account",
|
|
||||||
)
|
|
||||||
|
|
||||||
user_update = UserUpdate(is_active=False)
|
|
||||||
updated_user = await user_repo.update_user(user_id, user_update)
|
|
||||||
|
|
||||||
if not updated_user:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
|
|
||||||
)
|
|
||||||
|
|
||||||
return UserResponse.model_validate(updated_user)
|
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
from app.api.v1.endpoints import (
|
from app.api.v1.endpoints import (
|
||||||
auth,
|
agent_auth,
|
||||||
project,
|
project,
|
||||||
simulation,
|
simulation,
|
||||||
scada,
|
scada,
|
||||||
@@ -15,7 +15,6 @@ from app.api.v1.endpoints import (
|
|||||||
leakage,
|
leakage,
|
||||||
burst_detection,
|
burst_detection,
|
||||||
burst_location,
|
burst_location,
|
||||||
user_management, # 新增:用户管理
|
|
||||||
audit, # 新增:审计日志
|
audit, # 新增:审计日志
|
||||||
meta,
|
meta,
|
||||||
web_search,
|
web_search,
|
||||||
@@ -54,10 +53,7 @@ from app.api.v1.endpoints.timeseries import (
|
|||||||
api_router = APIRouter()
|
api_router = APIRouter()
|
||||||
|
|
||||||
# Core Services
|
# Core Services
|
||||||
api_router.include_router(auth.router, prefix="/auth", tags=["Auth"])
|
api_router.include_router(agent_auth.router, tags=["Agent Auth"])
|
||||||
api_router.include_router(
|
|
||||||
user_management.router, prefix="/users", tags=["User Management"]
|
|
||||||
) # 新增
|
|
||||||
api_router.include_router(audit.router, prefix="/audit", tags=["Audit Logs"]) # 新增
|
api_router.include_router(audit.router, prefix="/audit", tags=["Audit Logs"]) # 新增
|
||||||
api_router.include_router(meta.router, tags=["Metadata"])
|
api_router.include_router(meta.router, tags=["Metadata"])
|
||||||
api_router.include_router(project.router, tags=["Project"])
|
api_router.include_router(project.router, tags=["Project"])
|
||||||
|
|||||||
@@ -1,100 +0,0 @@
|
|||||||
from typing import Annotated, Optional
|
|
||||||
from fastapi import Depends, HTTPException, status, Request
|
|
||||||
from fastapi.security import OAuth2PasswordBearer
|
|
||||||
from jose import jwt, JWTError
|
|
||||||
from app.core.config import settings
|
|
||||||
from app.domain.schemas.user import UserInDB, TokenPayload
|
|
||||||
from app.infra.db.metadb.repositories.user_repository import UserRepository
|
|
||||||
from app.infra.db.postgresql.database import Database
|
|
||||||
|
|
||||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl=f"{settings.API_V1_STR}/auth/login")
|
|
||||||
|
|
||||||
|
|
||||||
# 数据库依赖
|
|
||||||
async def get_db(request: Request) -> Database:
|
|
||||||
"""
|
|
||||||
获取数据库实例
|
|
||||||
|
|
||||||
从 FastAPI app.state 中获取在启动时初始化的数据库连接
|
|
||||||
"""
|
|
||||||
if not hasattr(request.app.state, "db"):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
|
||||||
detail="Database not initialized",
|
|
||||||
)
|
|
||||||
return request.app.state.db
|
|
||||||
|
|
||||||
|
|
||||||
async def get_user_repository(db: Database = Depends(get_db)) -> UserRepository:
|
|
||||||
"""获取用户仓储实例"""
|
|
||||||
return UserRepository(db)
|
|
||||||
|
|
||||||
|
|
||||||
async def get_current_user(
|
|
||||||
token: str = Depends(oauth2_scheme),
|
|
||||||
user_repo: UserRepository = Depends(get_user_repository),
|
|
||||||
) -> UserInDB:
|
|
||||||
"""
|
|
||||||
获取当前登录用户
|
|
||||||
|
|
||||||
从 JWT Token 中解析用户信息,并从数据库验证
|
|
||||||
"""
|
|
||||||
credentials_exception = HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Could not validate credentials",
|
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
payload = jwt.decode(
|
|
||||||
token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]
|
|
||||||
)
|
|
||||||
username: str = payload.get("sub")
|
|
||||||
token_type: str = payload.get("type", "access")
|
|
||||||
|
|
||||||
if username is None:
|
|
||||||
raise credentials_exception
|
|
||||||
|
|
||||||
if token_type != "access":
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Invalid token type. Access token required.",
|
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
|
||||||
)
|
|
||||||
|
|
||||||
except JWTError:
|
|
||||||
raise credentials_exception
|
|
||||||
|
|
||||||
# 从数据库获取用户
|
|
||||||
user = await user_repo.get_user_by_username(username)
|
|
||||||
if user is None:
|
|
||||||
raise credentials_exception
|
|
||||||
|
|
||||||
return user
|
|
||||||
|
|
||||||
|
|
||||||
async def get_current_active_user(
|
|
||||||
current_user: UserInDB = Depends(get_current_user),
|
|
||||||
) -> UserInDB:
|
|
||||||
"""
|
|
||||||
获取当前活跃用户(必须是激活状态)
|
|
||||||
"""
|
|
||||||
if not current_user.is_active:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user"
|
|
||||||
)
|
|
||||||
return current_user
|
|
||||||
|
|
||||||
|
|
||||||
async def get_current_superuser(
|
|
||||||
current_user: UserInDB = Depends(get_current_user),
|
|
||||||
) -> UserInDB:
|
|
||||||
"""
|
|
||||||
获取当前超级管理员用户
|
|
||||||
"""
|
|
||||||
if not current_user.is_superuser:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail="Not enough privileges. Superuser access required.",
|
|
||||||
)
|
|
||||||
return current_user
|
|
||||||
@@ -8,35 +8,41 @@ from jose import JWTError, jwt
|
|||||||
from app.core.config import settings
|
from app.core.config import settings
|
||||||
|
|
||||||
oauth2_optional = OAuth2PasswordBearer(
|
oauth2_optional = OAuth2PasswordBearer(
|
||||||
tokenUrl=f"{settings.API_V1_STR}/auth/login", auto_error=False
|
tokenUrl="keycloak", auto_error=False
|
||||||
)
|
)
|
||||||
|
|
||||||
# logger = logging.getLogger(__name__)
|
# logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
async def get_current_keycloak_sub(
|
def _decode_keycloak_token(token: str) -> dict:
|
||||||
|
if not settings.KEYCLOAK_PUBLIC_KEY:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||||
|
detail="Keycloak public key is not configured",
|
||||||
|
)
|
||||||
|
|
||||||
|
key = settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n")
|
||||||
|
|
||||||
|
return jwt.decode(
|
||||||
|
token,
|
||||||
|
key,
|
||||||
|
algorithms=[settings.KEYCLOAK_ALGORITHM],
|
||||||
|
audience=settings.KEYCLOAK_AUDIENCE or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_current_keycloak_payload(
|
||||||
token: str | None = Depends(oauth2_optional),
|
token: str | None = Depends(oauth2_optional),
|
||||||
) -> UUID:
|
) -> dict:
|
||||||
if not token:
|
if not token:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
detail="Not authenticated",
|
detail="Not authenticated",
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
headers={"WWW-Authenticate": "Bearer"},
|
||||||
)
|
)
|
||||||
if settings.KEYCLOAK_PUBLIC_KEY:
|
|
||||||
key = settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n")
|
|
||||||
algorithms = [settings.KEYCLOAK_ALGORITHM]
|
|
||||||
else:
|
|
||||||
key = settings.SECRET_KEY
|
|
||||||
algorithms = [settings.ALGORITHM]
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
payload = jwt.decode(
|
return _decode_keycloak_token(token)
|
||||||
token,
|
|
||||||
key,
|
|
||||||
algorithms=algorithms,
|
|
||||||
audience=settings.KEYCLOAK_AUDIENCE or None,
|
|
||||||
)
|
|
||||||
except JWTError as exc:
|
except JWTError as exc:
|
||||||
# logger.warning("Keycloak token validation failed: %s", exc)
|
# logger.warning("Keycloak token validation failed: %s", exc)
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -45,6 +51,10 @@ async def get_current_keycloak_sub(
|
|||||||
headers={"WWW-Authenticate": "Bearer"},
|
headers={"WWW-Authenticate": "Bearer"},
|
||||||
) from exc
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
|
async def get_current_keycloak_sub(
|
||||||
|
payload: dict = Depends(get_current_keycloak_payload),
|
||||||
|
) -> UUID:
|
||||||
sub = payload.get("sub")
|
sub = payload.get("sub")
|
||||||
if not sub:
|
if not sub:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -64,35 +74,8 @@ async def get_current_keycloak_sub(
|
|||||||
|
|
||||||
|
|
||||||
async def get_current_keycloak_username(
|
async def get_current_keycloak_username(
|
||||||
token: str | None = Depends(oauth2_optional),
|
payload: dict = Depends(get_current_keycloak_payload),
|
||||||
) -> str:
|
) -> str:
|
||||||
if not token:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Not authenticated",
|
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
|
||||||
)
|
|
||||||
if settings.KEYCLOAK_PUBLIC_KEY:
|
|
||||||
key = settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n")
|
|
||||||
algorithms = [settings.KEYCLOAK_ALGORITHM]
|
|
||||||
else:
|
|
||||||
key = settings.SECRET_KEY
|
|
||||||
algorithms = [settings.ALGORITHM]
|
|
||||||
|
|
||||||
try:
|
|
||||||
payload = jwt.decode(
|
|
||||||
token,
|
|
||||||
key,
|
|
||||||
algorithms=algorithms,
|
|
||||||
audience=settings.KEYCLOAK_AUDIENCE or None,
|
|
||||||
)
|
|
||||||
except JWTError as exc:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
||||||
detail="Invalid token",
|
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
|
||||||
) from exc
|
|
||||||
|
|
||||||
username = payload.get("preferred_username") or payload.get("username")
|
username = payload.get("preferred_username") or payload.get("username")
|
||||||
if not username:
|
if not username:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
|
|||||||
@@ -1,106 +0,0 @@
|
|||||||
"""
|
|
||||||
权限控制依赖项和装饰器
|
|
||||||
|
|
||||||
基于角色的访问控制(RBAC)
|
|
||||||
"""
|
|
||||||
from typing import Callable
|
|
||||||
from fastapi import Depends, HTTPException, status
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from app.domain.schemas.user import UserInDB
|
|
||||||
from app.auth.dependencies import get_current_active_user
|
|
||||||
|
|
||||||
def require_role(required_role: UserRole):
|
|
||||||
"""
|
|
||||||
要求特定角色或更高权限
|
|
||||||
|
|
||||||
用法:
|
|
||||||
@router.get("/admin-only")
|
|
||||||
async def admin_endpoint(user: UserInDB = Depends(require_role(UserRole.ADMIN))):
|
|
||||||
...
|
|
||||||
|
|
||||||
Args:
|
|
||||||
required_role: 需要的最低角色
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
依赖函数
|
|
||||||
"""
|
|
||||||
async def role_checker(
|
|
||||||
current_user: UserInDB = Depends(get_current_active_user)
|
|
||||||
) -> UserInDB:
|
|
||||||
user_role = UserRole(current_user.role)
|
|
||||||
|
|
||||||
if not user_role.has_permission(required_role):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail=f"Insufficient permissions. Required role: {required_role.value}, "
|
|
||||||
f"Your role: {user_role.value}"
|
|
||||||
)
|
|
||||||
|
|
||||||
return current_user
|
|
||||||
|
|
||||||
return role_checker
|
|
||||||
|
|
||||||
# 预定义的权限检查依赖
|
|
||||||
require_admin = require_role(UserRole.ADMIN)
|
|
||||||
require_operator = require_role(UserRole.OPERATOR)
|
|
||||||
require_user = require_role(UserRole.USER)
|
|
||||||
|
|
||||||
def get_current_admin(
|
|
||||||
current_user: UserInDB = Depends(require_admin)
|
|
||||||
) -> UserInDB:
|
|
||||||
"""
|
|
||||||
获取当前管理员用户
|
|
||||||
|
|
||||||
等同于 Depends(require_role(UserRole.ADMIN))
|
|
||||||
"""
|
|
||||||
return current_user
|
|
||||||
|
|
||||||
def get_current_operator(
|
|
||||||
current_user: UserInDB = Depends(require_operator)
|
|
||||||
) -> UserInDB:
|
|
||||||
"""
|
|
||||||
获取当前操作员用户(或更高权限)
|
|
||||||
|
|
||||||
等同于 Depends(require_role(UserRole.OPERATOR))
|
|
||||||
"""
|
|
||||||
return current_user
|
|
||||||
|
|
||||||
def check_resource_owner(user_id: int, current_user: UserInDB) -> bool:
|
|
||||||
"""
|
|
||||||
检查是否是资源拥有者或管理员
|
|
||||||
|
|
||||||
Args:
|
|
||||||
user_id: 资源拥有者ID
|
|
||||||
current_user: 当前用户
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
是否有权限
|
|
||||||
"""
|
|
||||||
# 管理员可以访问所有资源
|
|
||||||
if UserRole(current_user.role).has_permission(UserRole.ADMIN):
|
|
||||||
return True
|
|
||||||
|
|
||||||
# 检查是否是资源拥有者
|
|
||||||
return current_user.id == user_id
|
|
||||||
|
|
||||||
def require_owner_or_admin(user_id: int):
|
|
||||||
"""
|
|
||||||
要求是资源拥有者或管理员
|
|
||||||
|
|
||||||
Args:
|
|
||||||
user_id: 资源拥有者ID
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
依赖函数
|
|
||||||
"""
|
|
||||||
async def owner_or_admin_checker(
|
|
||||||
current_user: UserInDB = Depends(get_current_active_user)
|
|
||||||
) -> UserInDB:
|
|
||||||
if not check_resource_owner(user_id, current_user):
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_403_FORBIDDEN,
|
|
||||||
detail="You don't have permission to access this resource"
|
|
||||||
)
|
|
||||||
return current_user
|
|
||||||
|
|
||||||
return owner_or_admin_checker
|
|
||||||
+1
-9
@@ -11,14 +11,6 @@ class Settings(BaseSettings):
|
|||||||
|
|
||||||
NETWORK_NAME: str = "default_network"
|
NETWORK_NAME: str = "default_network"
|
||||||
|
|
||||||
# JWT 配置
|
|
||||||
SECRET_KEY: str = (
|
|
||||||
"your-secret-key-here-change-in-production-use-openssl-rand-hex-32"
|
|
||||||
)
|
|
||||||
ALGORITHM: str = "HS256"
|
|
||||||
ACCESS_TOKEN_EXPIRE_MINUTES: int = 30
|
|
||||||
REFRESH_TOKEN_EXPIRE_DAYS: int = 7
|
|
||||||
|
|
||||||
# 数据加密密钥 (使用 Fernet)
|
# 数据加密密钥 (使用 Fernet)
|
||||||
ENCRYPTION_KEY: str = "" # 必须从环境变量设置
|
ENCRYPTION_KEY: str = "" # 必须从环境变量设置
|
||||||
DATABASE_ENCRYPTION_KEY: str = "" # project_databases.dsn_encrypted 专用
|
DATABASE_ENCRYPTION_KEY: str = "" # project_databases.dsn_encrypted 专用
|
||||||
@@ -59,7 +51,7 @@ class Settings(BaseSettings):
|
|||||||
PROJECT_TS_POOL_MIN_SIZE: int = 1
|
PROJECT_TS_POOL_MIN_SIZE: int = 1
|
||||||
PROJECT_TS_POOL_MAX_SIZE: int = 10
|
PROJECT_TS_POOL_MAX_SIZE: int = 10
|
||||||
|
|
||||||
# Keycloak JWT (optional override)
|
# Keycloak access token verification
|
||||||
KEYCLOAK_PUBLIC_KEY: str = ""
|
KEYCLOAK_PUBLIC_KEY: str = ""
|
||||||
KEYCLOAK_ALGORITHM: str = "RS256"
|
KEYCLOAK_ALGORITHM: str = "RS256"
|
||||||
KEYCLOAK_AUDIENCE: str = ""
|
KEYCLOAK_AUDIENCE: str = ""
|
||||||
|
|||||||
@@ -1,95 +0,0 @@
|
|||||||
from datetime import datetime, timedelta, timezone
|
|
||||||
from typing import Optional, Union, Any
|
|
||||||
|
|
||||||
from jose import jwt
|
|
||||||
from passlib.context import CryptContext
|
|
||||||
from app.core.config import settings
|
|
||||||
|
|
||||||
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
|
||||||
|
|
||||||
|
|
||||||
def _utc_now() -> datetime:
|
|
||||||
return datetime.now(timezone.utc)
|
|
||||||
|
|
||||||
|
|
||||||
def create_access_token(
|
|
||||||
subject: Union[str, Any], expires_delta: Optional[timedelta] = None
|
|
||||||
) -> str:
|
|
||||||
"""
|
|
||||||
创建 JWT Access Token
|
|
||||||
|
|
||||||
Args:
|
|
||||||
subject: 用户标识(通常是用户名或用户ID)
|
|
||||||
expires_delta: 过期时间增量
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
JWT token 字符串
|
|
||||||
"""
|
|
||||||
if expires_delta:
|
|
||||||
expire = _utc_now() + expires_delta
|
|
||||||
else:
|
|
||||||
expire = _utc_now() + timedelta(
|
|
||||||
minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES
|
|
||||||
)
|
|
||||||
|
|
||||||
to_encode = {
|
|
||||||
"exp": expire,
|
|
||||||
"sub": str(subject),
|
|
||||||
"type": "access",
|
|
||||||
"iat": _utc_now(),
|
|
||||||
}
|
|
||||||
encoded_jwt = jwt.encode(
|
|
||||||
to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM
|
|
||||||
)
|
|
||||||
return encoded_jwt
|
|
||||||
|
|
||||||
|
|
||||||
def create_refresh_token(subject: Union[str, Any]) -> str:
|
|
||||||
"""
|
|
||||||
创建 JWT Refresh Token(长期有效)
|
|
||||||
|
|
||||||
Args:
|
|
||||||
subject: 用户标识
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
JWT refresh token 字符串
|
|
||||||
"""
|
|
||||||
expire = _utc_now() + timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS)
|
|
||||||
|
|
||||||
to_encode = {
|
|
||||||
"exp": expire,
|
|
||||||
"sub": str(subject),
|
|
||||||
"type": "refresh",
|
|
||||||
"iat": _utc_now(),
|
|
||||||
}
|
|
||||||
encoded_jwt = jwt.encode(
|
|
||||||
to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM
|
|
||||||
)
|
|
||||||
return encoded_jwt
|
|
||||||
|
|
||||||
|
|
||||||
def verify_password(plain_password: str, hashed_password: str) -> bool:
|
|
||||||
"""
|
|
||||||
验证密码
|
|
||||||
|
|
||||||
Args:
|
|
||||||
plain_password: 明文密码
|
|
||||||
hashed_password: 密码哈希
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
是否匹配
|
|
||||||
"""
|
|
||||||
return pwd_context.verify(plain_password, hashed_password)
|
|
||||||
|
|
||||||
|
|
||||||
def get_password_hash(password: str) -> str:
|
|
||||||
"""
|
|
||||||
生成密码哈希
|
|
||||||
|
|
||||||
Args:
|
|
||||||
password: 明文密码
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
bcrypt 哈希字符串
|
|
||||||
"""
|
|
||||||
return pwd_context.hash(password)
|
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
from enum import Enum
|
|
||||||
|
|
||||||
class UserRole(str, Enum):
|
|
||||||
"""用户角色枚举"""
|
|
||||||
ADMIN = "ADMIN" # 管理员 - 完全权限
|
|
||||||
OPERATOR = "OPERATOR" # 操作员 - 可修改数据
|
|
||||||
USER = "USER" # 普通用户 - 读写权限
|
|
||||||
VIEWER = "VIEWER" # 观察者 - 仅查询权限
|
|
||||||
|
|
||||||
def __str__(self):
|
|
||||||
return self.value
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get_hierarchy(cls) -> dict:
|
|
||||||
"""
|
|
||||||
获取角色层级(数字越大权限越高)
|
|
||||||
"""
|
|
||||||
return {
|
|
||||||
cls.VIEWER: 1,
|
|
||||||
cls.USER: 2,
|
|
||||||
cls.OPERATOR: 3,
|
|
||||||
cls.ADMIN: 4,
|
|
||||||
}
|
|
||||||
|
|
||||||
def has_permission(self, required_role: 'UserRole') -> bool:
|
|
||||||
"""
|
|
||||||
检查当前角色是否有足够权限
|
|
||||||
|
|
||||||
Args:
|
|
||||||
required_role: 需要的最低角色
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if has permission
|
|
||||||
"""
|
|
||||||
hierarchy = self.get_hierarchy()
|
|
||||||
return hierarchy[self] >= hierarchy[required_role]
|
|
||||||
@@ -1,68 +0,0 @@
|
|||||||
from datetime import datetime
|
|
||||||
from typing import Optional
|
|
||||||
from pydantic import BaseModel, EmailStr, Field, ConfigDict
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
|
|
||||||
# ============================================
|
|
||||||
# Request Schemas (输入)
|
|
||||||
# ============================================
|
|
||||||
|
|
||||||
class UserCreate(BaseModel):
|
|
||||||
"""用户注册"""
|
|
||||||
username: str = Field(..., min_length=3, max_length=50,
|
|
||||||
description="用户名,3-50个字符")
|
|
||||||
email: EmailStr = Field(..., description="邮箱地址")
|
|
||||||
password: str = Field(..., min_length=6, max_length=100,
|
|
||||||
description="密码,至少6个字符")
|
|
||||||
role: UserRole = Field(default=UserRole.USER, description="用户角色")
|
|
||||||
|
|
||||||
class UserLogin(BaseModel):
|
|
||||||
"""用户登录"""
|
|
||||||
username: str = Field(..., description="用户名或邮箱")
|
|
||||||
password: str = Field(..., description="密码")
|
|
||||||
|
|
||||||
class UserUpdate(BaseModel):
|
|
||||||
"""用户信息更新"""
|
|
||||||
email: Optional[EmailStr] = None
|
|
||||||
password: Optional[str] = Field(None, min_length=6, max_length=100)
|
|
||||||
role: Optional[UserRole] = None
|
|
||||||
is_active: Optional[bool] = None
|
|
||||||
|
|
||||||
# ============================================
|
|
||||||
# Response Schemas (输出)
|
|
||||||
# ============================================
|
|
||||||
|
|
||||||
class UserResponse(BaseModel):
|
|
||||||
"""用户信息响应(不含密码)"""
|
|
||||||
id: int
|
|
||||||
username: str
|
|
||||||
email: str
|
|
||||||
role: UserRole
|
|
||||||
is_active: bool
|
|
||||||
is_superuser: bool
|
|
||||||
created_at: datetime
|
|
||||||
updated_at: datetime
|
|
||||||
|
|
||||||
model_config = ConfigDict(from_attributes=True)
|
|
||||||
|
|
||||||
class UserInDB(UserResponse):
|
|
||||||
"""数据库中的用户(含密码哈希)"""
|
|
||||||
hashed_password: str
|
|
||||||
|
|
||||||
# ============================================
|
|
||||||
# Token Schemas
|
|
||||||
# ============================================
|
|
||||||
|
|
||||||
class Token(BaseModel):
|
|
||||||
"""JWT Token 响应"""
|
|
||||||
access_token: str
|
|
||||||
refresh_token: Optional[str] = None
|
|
||||||
token_type: str = "bearer"
|
|
||||||
expires_in: int = Field(..., description="过期时间(秒)")
|
|
||||||
|
|
||||||
class TokenPayload(BaseModel):
|
|
||||||
"""JWT Token Payload"""
|
|
||||||
sub: str = Field(..., description="用户ID或用户名")
|
|
||||||
exp: Optional[int] = None
|
|
||||||
iat: Optional[int] = None
|
|
||||||
type: str = Field(default="access", description="token类型: access 或 refresh")
|
|
||||||
@@ -33,8 +33,6 @@ class AuditMiddleware(BaseHTTPMiddleware):
|
|||||||
|
|
||||||
# 需要审计的路径前缀
|
# 需要审计的路径前缀
|
||||||
AUDIT_PATHS = [
|
AUDIT_PATHS = [
|
||||||
# "/api/v1/auth/",
|
|
||||||
# "/api/v1/users/",
|
|
||||||
# "/api/v1/projects/",
|
# "/api/v1/projects/",
|
||||||
# "/api/v1/networks/",
|
# "/api/v1/networks/",
|
||||||
]
|
]
|
||||||
@@ -193,20 +191,14 @@ class AuditMiddleware(BaseHTTPMiddleware):
|
|||||||
return None
|
return None
|
||||||
sub = None
|
sub = None
|
||||||
try:
|
try:
|
||||||
key = (
|
if not settings.KEYCLOAK_PUBLIC_KEY:
|
||||||
settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n")
|
return None
|
||||||
if settings.KEYCLOAK_PUBLIC_KEY
|
|
||||||
else settings.SECRET_KEY
|
key = settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n")
|
||||||
)
|
|
||||||
algorithms = (
|
|
||||||
[settings.KEYCLOAK_ALGORITHM]
|
|
||||||
if settings.KEYCLOAK_PUBLIC_KEY
|
|
||||||
else [settings.ALGORITHM]
|
|
||||||
)
|
|
||||||
payload = jwt.decode(
|
payload = jwt.decode(
|
||||||
token,
|
token,
|
||||||
key,
|
key,
|
||||||
algorithms=algorithms,
|
algorithms=[settings.KEYCLOAK_ALGORITHM],
|
||||||
audience=settings.KEYCLOAK_AUDIENCE or None,
|
audience=settings.KEYCLOAK_AUDIENCE or None,
|
||||||
)
|
)
|
||||||
sub = payload.get("sub")
|
sub = payload.get("sub")
|
||||||
@@ -221,7 +213,7 @@ class AuditMiddleware(BaseHTTPMiddleware):
|
|||||||
keycloak_id = UUID(sub)
|
keycloak_id = UUID(sub)
|
||||||
user = await repo.get_user_by_keycloak_id(keycloak_id)
|
user = await repo.get_user_by_keycloak_id(keycloak_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
user = await repo.get_user_by_username(sub)
|
return None
|
||||||
if user and user.is_active:
|
if user and user.is_active:
|
||||||
return user.id
|
return user.id
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -1,235 +0,0 @@
|
|||||||
from typing import Optional, List
|
|
||||||
from datetime import datetime
|
|
||||||
from app.infra.db.postgresql.database import Database
|
|
||||||
from app.domain.schemas.user import UserCreate, UserUpdate, UserInDB
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from app.core.security import get_password_hash
|
|
||||||
import logging
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
class UserRepository:
|
|
||||||
"""用户数据访问层"""
|
|
||||||
|
|
||||||
def __init__(self, db: Database):
|
|
||||||
self.db = db
|
|
||||||
|
|
||||||
async def create_user(self, user: UserCreate) -> Optional[UserInDB]:
|
|
||||||
"""
|
|
||||||
创建新用户
|
|
||||||
|
|
||||||
Args:
|
|
||||||
user: 用户创建数据
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
创建的用户对象
|
|
||||||
"""
|
|
||||||
hashed_password = get_password_hash(user.password)
|
|
||||||
|
|
||||||
query = """
|
|
||||||
INSERT INTO users (username, email, hashed_password, role, is_active, is_superuser)
|
|
||||||
VALUES (%(username)s, %(email)s, %(hashed_password)s, %(role)s, TRUE, FALSE)
|
|
||||||
RETURNING id, username, email, hashed_password, role, is_active, is_superuser,
|
|
||||||
created_at, updated_at
|
|
||||||
"""
|
|
||||||
|
|
||||||
try:
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, {
|
|
||||||
'username': user.username,
|
|
||||||
'email': user.email,
|
|
||||||
'hashed_password': hashed_password,
|
|
||||||
'role': user.role.value
|
|
||||||
})
|
|
||||||
row = await cur.fetchone()
|
|
||||||
if row:
|
|
||||||
return UserInDB(**row)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error creating user: {e}")
|
|
||||||
raise
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def get_user_by_id(self, user_id: int) -> Optional[UserInDB]:
|
|
||||||
"""根据ID获取用户"""
|
|
||||||
query = """
|
|
||||||
SELECT id, username, email, hashed_password, role, is_active, is_superuser,
|
|
||||||
created_at, updated_at
|
|
||||||
FROM users
|
|
||||||
WHERE id = %(user_id)s
|
|
||||||
"""
|
|
||||||
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, {'user_id': user_id})
|
|
||||||
row = await cur.fetchone()
|
|
||||||
if row:
|
|
||||||
return UserInDB(**row)
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def get_user_by_username(self, username: str) -> Optional[UserInDB]:
|
|
||||||
"""根据用户名获取用户"""
|
|
||||||
query = """
|
|
||||||
SELECT id, username, email, hashed_password, role, is_active, is_superuser,
|
|
||||||
created_at, updated_at
|
|
||||||
FROM users
|
|
||||||
WHERE username = %(username)s
|
|
||||||
"""
|
|
||||||
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, {'username': username})
|
|
||||||
row = await cur.fetchone()
|
|
||||||
if row:
|
|
||||||
return UserInDB(**row)
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def get_user_by_email(self, email: str) -> Optional[UserInDB]:
|
|
||||||
"""根据邮箱获取用户"""
|
|
||||||
query = """
|
|
||||||
SELECT id, username, email, hashed_password, role, is_active, is_superuser,
|
|
||||||
created_at, updated_at
|
|
||||||
FROM users
|
|
||||||
WHERE email = %(email)s
|
|
||||||
"""
|
|
||||||
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, {'email': email})
|
|
||||||
row = await cur.fetchone()
|
|
||||||
if row:
|
|
||||||
return UserInDB(**row)
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def get_all_users(self, skip: int = 0, limit: int = 100) -> List[UserInDB]:
|
|
||||||
"""获取所有用户(分页)"""
|
|
||||||
query = """
|
|
||||||
SELECT id, username, email, hashed_password, role, is_active, is_superuser,
|
|
||||||
created_at, updated_at
|
|
||||||
FROM users
|
|
||||||
ORDER BY created_at DESC
|
|
||||||
LIMIT %(limit)s OFFSET %(skip)s
|
|
||||||
"""
|
|
||||||
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, {'skip': skip, 'limit': limit})
|
|
||||||
rows = await cur.fetchall()
|
|
||||||
return [UserInDB(**row) for row in rows]
|
|
||||||
|
|
||||||
async def update_user(self, user_id: int, user_update: UserUpdate) -> Optional[UserInDB]:
|
|
||||||
"""
|
|
||||||
更新用户信息
|
|
||||||
|
|
||||||
Args:
|
|
||||||
user_id: 用户ID
|
|
||||||
user_update: 更新数据
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
更新后的用户对象
|
|
||||||
"""
|
|
||||||
# 构建动态更新语句
|
|
||||||
update_fields = []
|
|
||||||
params = {'user_id': user_id}
|
|
||||||
|
|
||||||
if user_update.email is not None:
|
|
||||||
update_fields.append("email = %(email)s")
|
|
||||||
params['email'] = user_update.email
|
|
||||||
|
|
||||||
if user_update.password is not None:
|
|
||||||
update_fields.append("hashed_password = %(hashed_password)s")
|
|
||||||
params['hashed_password'] = get_password_hash(user_update.password)
|
|
||||||
|
|
||||||
if user_update.role is not None:
|
|
||||||
update_fields.append("role = %(role)s")
|
|
||||||
params['role'] = user_update.role.value
|
|
||||||
|
|
||||||
if user_update.is_active is not None:
|
|
||||||
update_fields.append("is_active = %(is_active)s")
|
|
||||||
params['is_active'] = user_update.is_active
|
|
||||||
|
|
||||||
if not update_fields:
|
|
||||||
return await self.get_user_by_id(user_id)
|
|
||||||
|
|
||||||
query = f"""
|
|
||||||
UPDATE users
|
|
||||||
SET {', '.join(update_fields)}, updated_at = CURRENT_TIMESTAMP
|
|
||||||
WHERE id = %(user_id)s
|
|
||||||
RETURNING id, username, email, hashed_password, role, is_active, is_superuser,
|
|
||||||
created_at, updated_at
|
|
||||||
"""
|
|
||||||
|
|
||||||
try:
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, params)
|
|
||||||
row = await cur.fetchone()
|
|
||||||
if row:
|
|
||||||
return UserInDB(**row)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error updating user {user_id}: {e}")
|
|
||||||
raise
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def delete_user(self, user_id: int) -> bool:
|
|
||||||
"""
|
|
||||||
删除用户
|
|
||||||
|
|
||||||
Args:
|
|
||||||
user_id: 用户ID
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
是否成功删除
|
|
||||||
"""
|
|
||||||
query = "DELETE FROM users WHERE id = %(user_id)s"
|
|
||||||
|
|
||||||
try:
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, {'user_id': user_id})
|
|
||||||
return cur.rowcount > 0
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error deleting user {user_id}: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def user_exists(self, username: str = None, email: str = None) -> bool:
|
|
||||||
"""
|
|
||||||
检查用户是否存在
|
|
||||||
|
|
||||||
Args:
|
|
||||||
username: 用户名
|
|
||||||
email: 邮箱
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
是否存在
|
|
||||||
"""
|
|
||||||
conditions = []
|
|
||||||
params = {}
|
|
||||||
|
|
||||||
if username:
|
|
||||||
conditions.append("username = %(username)s")
|
|
||||||
params['username'] = username
|
|
||||||
|
|
||||||
if email:
|
|
||||||
conditions.append("email = %(email)s")
|
|
||||||
params['email'] = email
|
|
||||||
|
|
||||||
if not conditions:
|
|
||||||
return False
|
|
||||||
|
|
||||||
query = f"""
|
|
||||||
SELECT EXISTS(
|
|
||||||
SELECT 1 FROM users WHERE {' OR '.join(conditions)}
|
|
||||||
)
|
|
||||||
"""
|
|
||||||
|
|
||||||
async with self.db.get_connection() as conn:
|
|
||||||
async with conn.cursor() as cur:
|
|
||||||
await cur.execute(query, params)
|
|
||||||
result = await cur.fetchone()
|
|
||||||
return result['exists'] if result else False
|
|
||||||
@@ -32,7 +32,6 @@ def test_load_auth_context_supports_aliases(monkeypatch):
|
|||||||
monkeypatch.setenv("TJWATER_SERVER", "http://server")
|
monkeypatch.setenv("TJWATER_SERVER", "http://server")
|
||||||
monkeypatch.setenv("TJWATER_ACCESS_TOKEN", "abc")
|
monkeypatch.setenv("TJWATER_ACCESS_TOKEN", "abc")
|
||||||
monkeypatch.setenv("TJWATER_PROJECT_ID", "p1")
|
monkeypatch.setenv("TJWATER_PROJECT_ID", "p1")
|
||||||
monkeypatch.setenv("TJWATER_USER_ID", "u1")
|
|
||||||
monkeypatch.setenv("TJWATER_USERNAME", "tester")
|
monkeypatch.setenv("TJWATER_USERNAME", "tester")
|
||||||
monkeypatch.setenv("TJWATER_NETWORK", "net1")
|
monkeypatch.setenv("TJWATER_NETWORK", "net1")
|
||||||
|
|
||||||
@@ -41,7 +40,6 @@ def test_load_auth_context_supports_aliases(monkeypatch):
|
|||||||
assert auth.server == "http://server"
|
assert auth.server == "http://server"
|
||||||
assert auth.access_token == "abc"
|
assert auth.access_token == "abc"
|
||||||
assert auth.project_id == "p1"
|
assert auth.project_id == "p1"
|
||||||
assert auth.user_id == "u1"
|
|
||||||
assert auth.username == "tester"
|
assert auth.username == "tester"
|
||||||
assert auth.network == "net1"
|
assert auth.network == "net1"
|
||||||
|
|
||||||
@@ -50,7 +48,6 @@ def test_build_runtime_context_uses_default_server(monkeypatch):
|
|||||||
monkeypatch.delenv("TJWATER_SERVER", raising=False)
|
monkeypatch.delenv("TJWATER_SERVER", raising=False)
|
||||||
monkeypatch.delenv("TJWATER_ACCESS_TOKEN", raising=False)
|
monkeypatch.delenv("TJWATER_ACCESS_TOKEN", raising=False)
|
||||||
monkeypatch.delenv("TJWATER_PROJECT_ID", raising=False)
|
monkeypatch.delenv("TJWATER_PROJECT_ID", raising=False)
|
||||||
monkeypatch.delenv("TJWATER_USER_ID", raising=False)
|
|
||||||
monkeypatch.delenv("TJWATER_USERNAME", raising=False)
|
monkeypatch.delenv("TJWATER_USERNAME", raising=False)
|
||||||
monkeypatch.delenv("TJWATER_NETWORK", raising=False)
|
monkeypatch.delenv("TJWATER_NETWORK", raising=False)
|
||||||
monkeypatch.delenv("TJWATER_EXTRA_HEADERS", raising=False)
|
monkeypatch.delenv("TJWATER_EXTRA_HEADERS", raising=False)
|
||||||
|
|||||||
@@ -46,7 +46,6 @@ class AuthContext:
|
|||||||
server: str | None = None
|
server: str | None = None
|
||||||
access_token: str | None = None
|
access_token: str | None = None
|
||||||
project_id: str | None = None
|
project_id: str | None = None
|
||||||
user_id: str | None = None
|
|
||||||
username: str | None = None
|
username: str | None = None
|
||||||
network: str | None = None
|
network: str | None = None
|
||||||
headers: dict[str, str] = field(default_factory=dict)
|
headers: dict[str, str] = field(default_factory=dict)
|
||||||
@@ -98,7 +97,6 @@ def load_auth_context(auth_stdin: bool = False) -> AuthContext:
|
|||||||
"server": os.getenv("TJWATER_SERVER"),
|
"server": os.getenv("TJWATER_SERVER"),
|
||||||
"access_token": os.getenv("TJWATER_ACCESS_TOKEN"),
|
"access_token": os.getenv("TJWATER_ACCESS_TOKEN"),
|
||||||
"project_id": os.getenv("TJWATER_PROJECT_ID"),
|
"project_id": os.getenv("TJWATER_PROJECT_ID"),
|
||||||
"user_id": os.getenv("TJWATER_USER_ID"),
|
|
||||||
"username": os.getenv("TJWATER_USERNAME"),
|
"username": os.getenv("TJWATER_USERNAME"),
|
||||||
"network": os.getenv("TJWATER_NETWORK"),
|
"network": os.getenv("TJWATER_NETWORK"),
|
||||||
"headers": json.loads(extra_headers) if extra_headers else {},
|
"headers": json.loads(extra_headers) if extra_headers else {},
|
||||||
@@ -117,7 +115,6 @@ def load_auth_context(auth_stdin: bool = False) -> AuthContext:
|
|||||||
server=_pick(raw, "server", "base_url"),
|
server=_pick(raw, "server", "base_url"),
|
||||||
access_token=_pick(raw, "access_token", "token", "accessToken"),
|
access_token=_pick(raw, "access_token", "token", "accessToken"),
|
||||||
project_id=_pick(raw, "project_id", "projectId", "x_project_id"),
|
project_id=_pick(raw, "project_id", "projectId", "x_project_id"),
|
||||||
user_id=_pick(raw, "user_id", "userId", "x_user_id"),
|
|
||||||
username=_pick(raw, "username", "preferred_username"),
|
username=_pick(raw, "username", "preferred_username"),
|
||||||
network=_pick(raw, "network", "project_code", "projectCode", "project"),
|
network=_pick(raw, "network", "project_code", "projectCode", "project"),
|
||||||
headers={str(key): str(value) for key, value in headers.items()},
|
headers={str(key): str(value) for key, value in headers.items()},
|
||||||
@@ -350,8 +347,6 @@ def build_headers(
|
|||||||
headers["X-Project-Id"] = require_project_id(ctx)
|
headers["X-Project-Id"] = require_project_id(ctx)
|
||||||
elif ctx.auth.project_id:
|
elif ctx.auth.project_id:
|
||||||
headers["X-Project-Id"] = ctx.auth.project_id
|
headers["X-Project-Id"] = ctx.auth.project_id
|
||||||
if ctx.auth.user_id:
|
|
||||||
headers["X-User-Id"] = ctx.auth.user_id
|
|
||||||
return headers
|
return headers
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -306,10 +306,9 @@ app/api/v1/endpoints/snapshots.py
|
|||||||
app/api/v1/endpoints/cache.py
|
app/api/v1/endpoints/cache.py
|
||||||
app/api/v1/endpoints/audit.py
|
app/api/v1/endpoints/audit.py
|
||||||
app/api/v1/endpoints/users.py
|
app/api/v1/endpoints/users.py
|
||||||
app/api/v1/endpoints/user_management.py
|
|
||||||
```
|
```
|
||||||
|
|
||||||
这些接口不纳入首批 Agent CLI。原因是它们更偏运维、审计、用户管理或状态回滚,不属于 Agent 面向水务业务分析的核心调用范围。
|
这些接口不纳入首批 Agent CLI。原因是它们更偏运维、审计或状态回滚,不属于 Agent 面向水务业务分析的核心调用范围。
|
||||||
|
|
||||||
暂不暴露:
|
暂不暴露:
|
||||||
|
|
||||||
@@ -339,10 +338,6 @@ GET /audit/logs/count
|
|||||||
GET /getuserschema/
|
GET /getuserschema/
|
||||||
GET /getuser/
|
GET /getuser/
|
||||||
GET /getallusers/
|
GET /getallusers/
|
||||||
PUT /users/{user_id}
|
|
||||||
DELETE /users/{user_id}
|
|
||||||
POST /users/{user_id}/activate
|
|
||||||
POST /users/{user_id}/deactivate
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## Help
|
## Help
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
from types import SimpleNamespace
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from fastapi import HTTPException, status
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from app.api.v1.endpoints import agent_auth as agent_auth_endpoint
|
||||||
|
from app.auth.keycloak_dependencies import get_current_keycloak_payload
|
||||||
|
from app.auth.metadata_dependencies import get_current_metadata_user
|
||||||
|
from app.auth.project_dependencies import ProjectContext, get_project_context
|
||||||
|
from tests.conftest import build_test_app
|
||||||
|
|
||||||
|
|
||||||
|
def _build_client(*, project_context=None, current_user=None) -> TestClient:
|
||||||
|
app = build_test_app(agent_auth_endpoint.router, "/api/v1")
|
||||||
|
if project_context is not None:
|
||||||
|
app.dependency_overrides[get_project_context] = lambda: project_context
|
||||||
|
if current_user is not None:
|
||||||
|
app.dependency_overrides[get_current_metadata_user] = lambda: current_user
|
||||||
|
app.dependency_overrides[get_current_keycloak_payload] = lambda: {"exp": 1781183400}
|
||||||
|
return TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_auth_context_returns_metadata_user_and_project_context():
|
||||||
|
user_id = uuid4()
|
||||||
|
keycloak_sub = uuid4()
|
||||||
|
project_id = uuid4()
|
||||||
|
client = _build_client(
|
||||||
|
project_context=ProjectContext(
|
||||||
|
project_id=project_id,
|
||||||
|
user_id=user_id,
|
||||||
|
project_role="editor",
|
||||||
|
),
|
||||||
|
current_user=SimpleNamespace(
|
||||||
|
id=user_id,
|
||||||
|
keycloak_id=keycloak_sub,
|
||||||
|
username="alice",
|
||||||
|
role="user",
|
||||||
|
is_superuser=False,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
response = client.get("/api/v1/agent/auth/context")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {
|
||||||
|
"user_id": str(user_id),
|
||||||
|
"keycloak_sub": str(keycloak_sub),
|
||||||
|
"username": "alice",
|
||||||
|
"role": "user",
|
||||||
|
"is_superuser": False,
|
||||||
|
"project_id": str(project_id),
|
||||||
|
"project_role": "editor",
|
||||||
|
"token_expires_at": "2026-06-11T13:10:00+00:00",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_auth_context_propagates_project_auth_failures():
|
||||||
|
def reject_project():
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
detail="No access to project",
|
||||||
|
)
|
||||||
|
|
||||||
|
app = build_test_app(agent_auth_endpoint.router, "/api/v1")
|
||||||
|
app.dependency_overrides[get_project_context] = reject_project
|
||||||
|
app.dependency_overrides[get_current_metadata_user] = lambda: SimpleNamespace(
|
||||||
|
id=uuid4(),
|
||||||
|
keycloak_id=uuid4(),
|
||||||
|
username="alice",
|
||||||
|
role="user",
|
||||||
|
is_superuser=False,
|
||||||
|
)
|
||||||
|
app.dependency_overrides[get_current_keycloak_payload] = lambda: {"exp": 1781183400}
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
response = client.get("/api/v1/agent/auth/context")
|
||||||
|
|
||||||
|
assert response.status_code == 403
|
||||||
|
assert response.json()["detail"] == "No access to project"
|
||||||
@@ -2,7 +2,7 @@
|
|||||||
"""
|
"""
|
||||||
测试新增 API 集成
|
测试新增 API 集成
|
||||||
|
|
||||||
验证新的认证、用户管理和审计日志接口是否正确集成
|
验证 Keycloak/metadata 认证和审计日志接口是否正确集成
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import sys
|
import sys
|
||||||
@@ -17,16 +17,15 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.
|
|||||||
"module_name, desc",
|
"module_name, desc",
|
||||||
[
|
[
|
||||||
("app.core.encryption", "加密模块"),
|
("app.core.encryption", "加密模块"),
|
||||||
("app.core.security", "安全模块"),
|
|
||||||
("app.core.audit", "审计模块"),
|
("app.core.audit", "审计模块"),
|
||||||
("app.domain.models.role", "角色模型"),
|
|
||||||
("app.domain.schemas.user", "用户Schema"),
|
|
||||||
("app.domain.schemas.audit", "审计Schema"),
|
("app.domain.schemas.audit", "审计Schema"),
|
||||||
("app.auth.permissions", "权限控制"),
|
("app.auth.keycloak_dependencies", "Keycloak Token 校验"),
|
||||||
("app.api.v1.endpoints.auth", "认证接口"),
|
("app.auth.metadata_dependencies", "Metadata 用户解析"),
|
||||||
("app.api.v1.endpoints.user_management", "用户管理接口"),
|
("app.auth.project_dependencies", "项目权限控制"),
|
||||||
|
("app.api.v1.endpoints.agent_auth", "Agent 认证上下文接口"),
|
||||||
|
("app.api.v1.endpoints.meta", "Metadata 接口"),
|
||||||
("app.api.v1.endpoints.audit", "审计日志接口"),
|
("app.api.v1.endpoints.audit", "审计日志接口"),
|
||||||
("app.infra.db.metadb.repositories.user_repository", "用户仓储"),
|
("app.infra.db.metadb.repositories.metadata_repository", "Metadata 仓储"),
|
||||||
("app.infra.db.metadb.repositories.audit_repository", "审计仓储"),
|
("app.infra.db.metadb.repositories.audit_repository", "审计仓储"),
|
||||||
("app.infra.audit.middleware", "审计中间件"),
|
("app.infra.audit.middleware", "审计中间件"),
|
||||||
],
|
],
|
||||||
@@ -49,8 +48,8 @@ def test_router_configuration():
|
|||||||
routes = [r.path for r in api_router.routes if hasattr(r, "path")]
|
routes = [r.path for r in api_router.routes if hasattr(r, "path")]
|
||||||
|
|
||||||
# 验证基础路径是否存在
|
# 验证基础路径是否存在
|
||||||
assert any("/auth" in r for r in routes), "缺少认证相关路由 (/auth)"
|
assert any("/agent/auth/context" in r for r in routes), "缺少 Agent 认证上下文路由"
|
||||||
assert any("/users" in r for r in routes), "缺少用户管理路由 (/users)"
|
assert any("/meta" in r for r in routes), "缺少 Metadata 路由"
|
||||||
assert any("/audit" in r for r in routes), "缺少审计日志路由 (/audit)"
|
assert any("/audit" in r for r in routes), "缺少审计日志路由 (/audit)"
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -1,139 +0,0 @@
|
|||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import AsyncMock
|
|
||||||
|
|
||||||
from fastapi.testclient import TestClient
|
|
||||||
|
|
||||||
from app.api.v1.endpoints import auth as auth_endpoint
|
|
||||||
from app.auth.dependencies import get_current_active_user, get_user_repository
|
|
||||||
from app.core.security import create_access_token, create_refresh_token, get_password_hash
|
|
||||||
from tests.conftest import build_test_app, make_user
|
|
||||||
|
|
||||||
|
|
||||||
def _build_client(repo, current_user=None) -> TestClient:
|
|
||||||
app = build_test_app(auth_endpoint.router, "/api/v1/auth")
|
|
||||||
app.dependency_overrides[get_user_repository] = lambda: repo
|
|
||||||
if current_user is not None:
|
|
||||||
app.dependency_overrides[get_current_active_user] = lambda: current_user
|
|
||||||
return TestClient(app)
|
|
||||||
|
|
||||||
|
|
||||||
def test_register_success():
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
user_exists=AsyncMock(side_effect=[False, False]),
|
|
||||||
create_user=AsyncMock(return_value=make_user()),
|
|
||||||
)
|
|
||||||
client = _build_client(repo)
|
|
||||||
|
|
||||||
response = client.post(
|
|
||||||
"/api/v1/auth/register",
|
|
||||||
json={
|
|
||||||
"username": "tester",
|
|
||||||
"email": "tester@example.com",
|
|
||||||
"password": "secret123",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 201
|
|
||||||
assert response.json()["username"] == "tester"
|
|
||||||
|
|
||||||
|
|
||||||
def test_register_rejects_duplicate_username():
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
user_exists=AsyncMock(side_effect=[True]),
|
|
||||||
create_user=AsyncMock(),
|
|
||||||
)
|
|
||||||
client = _build_client(repo)
|
|
||||||
|
|
||||||
response = client.post(
|
|
||||||
"/api/v1/auth/register",
|
|
||||||
json={
|
|
||||||
"username": "tester",
|
|
||||||
"email": "tester@example.com",
|
|
||||||
"password": "secret123",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 400
|
|
||||||
assert response.json()["detail"] == "Username already registered"
|
|
||||||
repo.create_user.assert_not_awaited()
|
|
||||||
|
|
||||||
|
|
||||||
def test_login_supports_email_lookup():
|
|
||||||
hashed_password = get_password_hash("secret123")
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
get_user_by_username=AsyncMock(return_value=None),
|
|
||||||
get_user_by_email=AsyncMock(
|
|
||||||
return_value=make_user(
|
|
||||||
email="tester@example.com",
|
|
||||||
hashed_password=hashed_password,
|
|
||||||
)
|
|
||||||
),
|
|
||||||
)
|
|
||||||
client = _build_client(repo)
|
|
||||||
|
|
||||||
response = client.post(
|
|
||||||
"/api/v1/auth/login",
|
|
||||||
data={"username": "tester@example.com", "password": "secret123"},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert response.json()["token_type"] == "bearer"
|
|
||||||
repo.get_user_by_email.assert_awaited_once_with("tester@example.com")
|
|
||||||
|
|
||||||
|
|
||||||
def test_login_simple_uses_query_params():
|
|
||||||
hashed_password = get_password_hash("secret123")
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
get_user_by_username=AsyncMock(
|
|
||||||
return_value=make_user(hashed_password=hashed_password)
|
|
||||||
),
|
|
||||||
get_user_by_email=AsyncMock(),
|
|
||||||
)
|
|
||||||
client = _build_client(repo)
|
|
||||||
|
|
||||||
response = client.post(
|
|
||||||
"/api/v1/auth/login/simple",
|
|
||||||
params={"username": "tester", "password": "secret123"},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert response.json()["token_type"] == "bearer"
|
|
||||||
|
|
||||||
|
|
||||||
def test_me_returns_current_user_info():
|
|
||||||
client = _build_client(SimpleNamespace(), current_user=make_user(username="alice"))
|
|
||||||
|
|
||||||
response = client.get("/api/v1/auth/me")
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert response.json()["username"] == "alice"
|
|
||||||
|
|
||||||
|
|
||||||
def test_refresh_rejects_access_token():
|
|
||||||
repo = SimpleNamespace(get_user_by_username=AsyncMock())
|
|
||||||
client = _build_client(repo)
|
|
||||||
|
|
||||||
response = client.post(
|
|
||||||
"/api/v1/auth/refresh",
|
|
||||||
params={"refresh_token": create_access_token("tester")},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 401
|
|
||||||
|
|
||||||
|
|
||||||
def test_refresh_success_returns_new_access_token():
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
get_user_by_username=AsyncMock(return_value=make_user()),
|
|
||||||
)
|
|
||||||
client = _build_client(repo)
|
|
||||||
refresh_token = create_refresh_token("tester")
|
|
||||||
|
|
||||||
response = client.post(
|
|
||||||
"/api/v1/auth/refresh",
|
|
||||||
params={"refresh_token": refresh_token},
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
payload = response.json()
|
|
||||||
assert payload["refresh_token"] == refresh_token
|
|
||||||
assert payload["token_type"] == "bearer"
|
|
||||||
@@ -1,95 +0,0 @@
|
|||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import AsyncMock
|
|
||||||
|
|
||||||
from fastapi.testclient import TestClient
|
|
||||||
|
|
||||||
from app.api.v1.endpoints import user_management as user_management_endpoint
|
|
||||||
from app.auth.dependencies import get_current_active_user, get_user_repository
|
|
||||||
from app.auth.permissions import get_current_admin
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from tests.conftest import build_test_app, make_user
|
|
||||||
|
|
||||||
|
|
||||||
def _build_client(repo, *, current_user=None, admin_user=None) -> TestClient:
|
|
||||||
app = build_test_app(user_management_endpoint.router, "/users")
|
|
||||||
app.dependency_overrides[get_user_repository] = lambda: repo
|
|
||||||
if current_user is not None:
|
|
||||||
app.dependency_overrides[get_current_active_user] = lambda: current_user
|
|
||||||
if admin_user is not None:
|
|
||||||
app.dependency_overrides[get_current_admin] = lambda: admin_user
|
|
||||||
return TestClient(app)
|
|
||||||
|
|
||||||
|
|
||||||
def test_list_users_requires_admin_role():
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
get_all_users=AsyncMock(
|
|
||||||
return_value=[
|
|
||||||
make_user(id=1, username="admin", role=UserRole.ADMIN),
|
|
||||||
make_user(id=2, username="user2"),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
client = _build_client(
|
|
||||||
repo,
|
|
||||||
current_user=make_user(id=1, role=UserRole.ADMIN),
|
|
||||||
)
|
|
||||||
|
|
||||||
response = client.get("/users/", params={"skip": 5, "limit": 2})
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert len(response.json()) == 2
|
|
||||||
repo.get_all_users.assert_awaited_once_with(skip=5, limit=2)
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_user_rejects_non_owner_non_admin():
|
|
||||||
repo = SimpleNamespace(get_user_by_id=AsyncMock())
|
|
||||||
client = _build_client(repo, current_user=make_user(id=2, role=UserRole.USER))
|
|
||||||
|
|
||||||
response = client.get("/users/3")
|
|
||||||
|
|
||||||
assert response.status_code == 403
|
|
||||||
assert response.json()["detail"] == "You don't have permission to view this user"
|
|
||||||
repo.get_user_by_id.assert_not_awaited()
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_user_blocks_role_change_for_non_admin():
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
get_user_by_id=AsyncMock(return_value=make_user(id=1)),
|
|
||||||
update_user=AsyncMock(),
|
|
||||||
)
|
|
||||||
client = _build_client(repo, current_user=make_user(id=1, role=UserRole.USER))
|
|
||||||
|
|
||||||
response = client.put("/users/1", json={"role": "ADMIN"})
|
|
||||||
|
|
||||||
assert response.status_code == 403
|
|
||||||
assert response.json()["detail"] == "Only admins can change user roles"
|
|
||||||
repo.update_user.assert_not_awaited()
|
|
||||||
|
|
||||||
|
|
||||||
def test_delete_user_blocks_self_delete_for_admin():
|
|
||||||
admin_user = make_user(id=1, role=UserRole.ADMIN, is_superuser=True)
|
|
||||||
repo = SimpleNamespace(delete_user=AsyncMock())
|
|
||||||
client = _build_client(repo, admin_user=admin_user)
|
|
||||||
|
|
||||||
response = client.delete("/users/1")
|
|
||||||
|
|
||||||
assert response.status_code == 400
|
|
||||||
assert response.json()["detail"] == "You cannot delete your own account"
|
|
||||||
repo.delete_user.assert_not_awaited()
|
|
||||||
|
|
||||||
|
|
||||||
def test_activate_user_updates_active_flag():
|
|
||||||
repo = SimpleNamespace(
|
|
||||||
update_user=AsyncMock(return_value=make_user(id=2, is_active=True)),
|
|
||||||
)
|
|
||||||
client = _build_client(
|
|
||||||
repo,
|
|
||||||
admin_user=make_user(id=1, role=UserRole.ADMIN, is_superuser=True),
|
|
||||||
)
|
|
||||||
|
|
||||||
response = client.post("/users/2/activate")
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert response.json()["is_active"] is True
|
|
||||||
user_update = repo.update_user.await_args.args[1]
|
|
||||||
assert user_update.is_active is True
|
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
from jose import jwt
|
|
||||||
|
|
||||||
from app.core.config import settings
|
|
||||||
from app.core.security import (
|
|
||||||
create_access_token,
|
|
||||||
create_refresh_token,
|
|
||||||
get_password_hash,
|
|
||||||
verify_password,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_password_hash_roundtrip():
|
|
||||||
hashed = get_password_hash("secret123")
|
|
||||||
assert hashed != "secret123"
|
|
||||||
assert verify_password("secret123", hashed) is True
|
|
||||||
assert verify_password("wrong", hashed) is False
|
|
||||||
|
|
||||||
|
|
||||||
def test_create_access_token_sets_access_type():
|
|
||||||
token = create_access_token("alice")
|
|
||||||
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
|
||||||
|
|
||||||
assert payload["sub"] == "alice"
|
|
||||||
assert payload["type"] == "access"
|
|
||||||
assert "exp" in payload
|
|
||||||
assert "iat" in payload
|
|
||||||
|
|
||||||
|
|
||||||
def test_create_refresh_token_sets_refresh_type():
|
|
||||||
token = create_refresh_token("alice")
|
|
||||||
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
|
||||||
|
|
||||||
assert payload["sub"] == "alice"
|
|
||||||
assert payload["type"] == "refresh"
|
|
||||||
assert "exp" in payload
|
|
||||||
assert "iat" in payload
|
|
||||||
@@ -159,25 +159,6 @@ class FakeAsyncSession:
|
|||||||
self.refreshed.append(obj)
|
self.refreshed.append(obj)
|
||||||
|
|
||||||
|
|
||||||
def make_user(**overrides):
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from app.domain.schemas.user import UserInDB
|
|
||||||
|
|
||||||
data = {
|
|
||||||
"id": 1,
|
|
||||||
"username": "tester",
|
|
||||||
"email": "tester@example.com",
|
|
||||||
"hashed_password": "hashed-password",
|
|
||||||
"role": UserRole.USER,
|
|
||||||
"is_active": True,
|
|
||||||
"is_superuser": False,
|
|
||||||
"created_at": datetime(2025, 1, 1, tzinfo=timezone.utc),
|
|
||||||
"updated_at": datetime(2025, 1, 1, tzinfo=timezone.utc),
|
|
||||||
}
|
|
||||||
data.update(overrides)
|
|
||||||
return UserInDB(**data)
|
|
||||||
|
|
||||||
|
|
||||||
def make_audit_log(**overrides):
|
def make_audit_log(**overrides):
|
||||||
data = {
|
data = {
|
||||||
"id": uuid4(),
|
"id": uuid4(),
|
||||||
|
|||||||
@@ -22,14 +22,14 @@ def test_create_log_adds_commits_and_refreshes(monkeypatch):
|
|||||||
|
|
||||||
result = asyncio.run(
|
result = asyncio.run(
|
||||||
repo.create_log(
|
repo.create_log(
|
||||||
action="LOGIN",
|
action="CREATE_PROJECT",
|
||||||
request_method="POST",
|
request_method="POST",
|
||||||
request_path="/auth/login",
|
request_path="/api/v1/projects",
|
||||||
response_status=200,
|
response_status=200,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.action == "LOGIN"
|
assert result.action == "CREATE_PROJECT"
|
||||||
assert result.request_method == "POST"
|
assert result.request_method == "POST"
|
||||||
assert session.commit_count == 1
|
assert session.commit_count == 1
|
||||||
assert len(session.added) == 1
|
assert len(session.added) == 1
|
||||||
|
|||||||
@@ -1,97 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import AsyncMock
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
from fastapi import HTTPException
|
|
||||||
|
|
||||||
from app.auth import dependencies
|
|
||||||
from app.core.security import create_access_token, create_refresh_token
|
|
||||||
from tests.conftest import make_user
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_db_returns_app_state_db():
|
|
||||||
request = SimpleNamespace(app=SimpleNamespace(state=SimpleNamespace(db="db-instance")))
|
|
||||||
|
|
||||||
result = asyncio.run(dependencies.get_db(request))
|
|
||||||
|
|
||||||
assert result == "db-instance"
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_db_raises_when_database_missing():
|
|
||||||
request = SimpleNamespace(app=SimpleNamespace(state=SimpleNamespace()))
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(dependencies.get_db(request))
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 503
|
|
||||||
assert exc_info.value.detail == "Database not initialized"
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_current_user_accepts_valid_access_token():
|
|
||||||
repo = SimpleNamespace(get_user_by_username=AsyncMock(return_value=make_user()))
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
dependencies.get_current_user(
|
|
||||||
token=create_access_token("tester"),
|
|
||||||
user_repo=repo,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result.username == "tester"
|
|
||||||
repo.get_user_by_username.assert_awaited_once_with("tester")
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_current_user_rejects_refresh_token():
|
|
||||||
repo = SimpleNamespace(get_user_by_username=AsyncMock())
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(
|
|
||||||
dependencies.get_current_user(
|
|
||||||
token=create_refresh_token("tester"),
|
|
||||||
user_repo=repo,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 401
|
|
||||||
assert exc_info.value.detail == "Invalid token type. Access token required."
|
|
||||||
repo.get_user_by_username.assert_not_awaited()
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_current_user_rejects_missing_user():
|
|
||||||
repo = SimpleNamespace(get_user_by_username=AsyncMock(return_value=None))
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(
|
|
||||||
dependencies.get_current_user(
|
|
||||||
token=create_access_token("ghost"),
|
|
||||||
user_repo=repo,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 401
|
|
||||||
assert exc_info.value.detail == "Could not validate credentials"
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_current_active_user_rejects_inactive_user():
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(
|
|
||||||
dependencies.get_current_active_user(
|
|
||||||
current_user=make_user(is_active=False),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 403
|
|
||||||
assert exc_info.value.detail == "Inactive user"
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_current_superuser_rejects_non_superuser():
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(
|
|
||||||
dependencies.get_current_superuser(
|
|
||||||
current_user=make_user(is_superuser=False),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 403
|
|
||||||
assert exc_info.value.detail == "Not enough privileges. Superuser access required."
|
|
||||||
@@ -1,56 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
import pytest
|
|
||||||
from fastapi import HTTPException
|
|
||||||
|
|
||||||
from app.auth import permissions
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from tests.conftest import make_user
|
|
||||||
|
|
||||||
|
|
||||||
def test_require_role_allows_higher_privilege_user():
|
|
||||||
checker = permissions.require_role(UserRole.OPERATOR)
|
|
||||||
|
|
||||||
result = asyncio.run(checker(current_user=make_user(role=UserRole.ADMIN)))
|
|
||||||
|
|
||||||
assert result.role == UserRole.ADMIN
|
|
||||||
|
|
||||||
|
|
||||||
def test_require_role_rejects_insufficient_role():
|
|
||||||
checker = permissions.require_role(UserRole.ADMIN)
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(checker(current_user=make_user(role=UserRole.USER)))
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 403
|
|
||||||
assert "Required role: ADMIN" in exc_info.value.detail
|
|
||||||
|
|
||||||
|
|
||||||
def test_check_resource_owner_allows_admin():
|
|
||||||
assert permissions.check_resource_owner(
|
|
||||||
99,
|
|
||||||
make_user(id=1, role=UserRole.ADMIN),
|
|
||||||
) is True
|
|
||||||
|
|
||||||
|
|
||||||
def test_check_resource_owner_allows_owner():
|
|
||||||
assert permissions.check_resource_owner(
|
|
||||||
7,
|
|
||||||
make_user(id=7, role=UserRole.USER),
|
|
||||||
) is True
|
|
||||||
|
|
||||||
|
|
||||||
def test_check_resource_owner_rejects_other_user():
|
|
||||||
assert permissions.check_resource_owner(
|
|
||||||
7,
|
|
||||||
make_user(id=8, role=UserRole.USER),
|
|
||||||
) is False
|
|
||||||
|
|
||||||
|
|
||||||
def test_require_owner_or_admin_rejects_other_user():
|
|
||||||
checker = permissions.require_owner_or_admin(7)
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
asyncio.run(checker(current_user=make_user(id=8, role=UserRole.USER)))
|
|
||||||
|
|
||||||
assert exc_info.value.status_code == 403
|
|
||||||
assert exc_info.value.detail == "You don't have permission to access this resource"
|
|
||||||
@@ -1,124 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
from unittest.mock import AsyncMock
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
from app.domain.models.role import UserRole
|
|
||||||
from app.domain.schemas.user import UserCreate, UserUpdate
|
|
||||||
from app.infra.db.metadb.repositories.user_repository import UserRepository
|
|
||||||
from tests.conftest import FakeCursor, FakeDB
|
|
||||||
|
|
||||||
|
|
||||||
def _user_row(**overrides):
|
|
||||||
base = {
|
|
||||||
"id": 1,
|
|
||||||
"username": "tester",
|
|
||||||
"email": "tester@example.com",
|
|
||||||
"hashed_password": "hashed-password",
|
|
||||||
"role": "USER",
|
|
||||||
"is_active": True,
|
|
||||||
"is_superuser": False,
|
|
||||||
"created_at": "2025-01-01T00:00:00+00:00",
|
|
||||||
"updated_at": "2025-01-01T00:00:00+00:00",
|
|
||||||
}
|
|
||||||
base.update(overrides)
|
|
||||||
return base
|
|
||||||
|
|
||||||
|
|
||||||
def test_create_user_hashes_password_and_returns_model(monkeypatch):
|
|
||||||
cursor = FakeCursor(fetchone_results=[_user_row()])
|
|
||||||
repo = UserRepository(FakeDB(cursor))
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"app.infra.db.metadb.repositories.user_repository.get_password_hash",
|
|
||||||
lambda password: f"hashed::{password}",
|
|
||||||
)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
repo.create_user(
|
|
||||||
UserCreate(
|
|
||||||
username="tester",
|
|
||||||
email="tester@example.com",
|
|
||||||
password="secret123",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result is not None
|
|
||||||
assert result.username == "tester"
|
|
||||||
assert cursor.executed[0][1]["hashed_password"] == "hashed::secret123"
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_user_without_fields_returns_existing_user(monkeypatch):
|
|
||||||
repo = UserRepository(FakeDB(FakeCursor()))
|
|
||||||
existing_user = AsyncMock(return_value="existing")
|
|
||||||
monkeypatch.setattr(repo, "get_user_by_id", existing_user)
|
|
||||||
|
|
||||||
result = asyncio.run(repo.update_user(1, UserUpdate()))
|
|
||||||
|
|
||||||
assert result == "existing"
|
|
||||||
existing_user.assert_awaited_once_with(1)
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_user_builds_dynamic_query(monkeypatch):
|
|
||||||
cursor = FakeCursor(fetchone_results=[_user_row(role="ADMIN", email="new@example.com")])
|
|
||||||
repo = UserRepository(FakeDB(cursor))
|
|
||||||
monkeypatch.setattr(
|
|
||||||
"app.infra.db.metadb.repositories.user_repository.get_password_hash",
|
|
||||||
lambda password: f"hashed::{password}",
|
|
||||||
)
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
repo.update_user(
|
|
||||||
1,
|
|
||||||
UserUpdate(
|
|
||||||
email="new@example.com",
|
|
||||||
password="new-secret",
|
|
||||||
role=UserRole.ADMIN,
|
|
||||||
is_active=False,
|
|
||||||
),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result is not None
|
|
||||||
query, params = cursor.executed[0]
|
|
||||||
assert "email = %(email)s" in query
|
|
||||||
assert "hashed_password = %(hashed_password)s" in query
|
|
||||||
assert "role = %(role)s" in query
|
|
||||||
assert "is_active = %(is_active)s" in query
|
|
||||||
assert params["hashed_password"] == "hashed::new-secret"
|
|
||||||
assert params["role"] == "ADMIN"
|
|
||||||
assert params["is_active"] is False
|
|
||||||
|
|
||||||
|
|
||||||
def test_delete_user_returns_false_when_execute_raises():
|
|
||||||
cursor = FakeCursor()
|
|
||||||
cursor.execute = AsyncMock(side_effect=RuntimeError("boom"))
|
|
||||||
repo = UserRepository(FakeDB(cursor))
|
|
||||||
|
|
||||||
result = asyncio.run(repo.delete_user(1))
|
|
||||||
|
|
||||||
assert result is False
|
|
||||||
|
|
||||||
|
|
||||||
def test_user_exists_short_circuits_without_filters():
|
|
||||||
cursor = FakeCursor()
|
|
||||||
repo = UserRepository(FakeDB(cursor))
|
|
||||||
|
|
||||||
result = asyncio.run(repo.user_exists())
|
|
||||||
|
|
||||||
assert result is False
|
|
||||||
assert cursor.executed == []
|
|
||||||
|
|
||||||
|
|
||||||
def test_user_exists_checks_username_or_email():
|
|
||||||
cursor = FakeCursor(fetchone_results=[{"exists": True}])
|
|
||||||
repo = UserRepository(FakeDB(cursor))
|
|
||||||
|
|
||||||
result = asyncio.run(
|
|
||||||
repo.user_exists(username="tester", email="tester@example.com")
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result is True
|
|
||||||
query, params = cursor.executed[0]
|
|
||||||
assert "username = %(username)s OR email = %(email)s" in query
|
|
||||||
assert params == {"username": "tester", "email": "tester@example.com"}
|
|
||||||
Reference in New Issue
Block a user