diff --git a/app/api/v1/endpoints/agent_auth.py b/app/api/v1/endpoints/agent_auth.py new file mode 100644 index 0000000..c3f91d6 --- /dev/null +++ b/app/api/v1/endpoints/agent_auth.py @@ -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, + ) diff --git a/app/api/v1/endpoints/auth.py b/app/api/v1/endpoints/auth.py deleted file mode 100644 index 819a094..0000000 --- a/app/api/v1/endpoints/auth.py +++ /dev/null @@ -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, - ) diff --git a/app/api/v1/endpoints/network/geometry.py b/app/api/v1/endpoints/network/geometry.py index 8a99743..5adb575 100644 --- a/app/api/v1/endpoints/network/geometry.py +++ b/app/api/v1/endpoints/network/geometry.py @@ -10,7 +10,7 @@ from app.services.tjnetwork import ( get_network_node_coords, 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 import msgpack @@ -64,7 +64,7 @@ async def fastapi_get_network_in_extent( @router.get( "/getnetworkgeometries/", - dependencies=[Depends(verify_token)], + dependencies=[Depends(get_current_metadata_user)], summary="获取完整网络几何信息", description="获取整个水网的所有节点、管线和SCADA点的几何信息(需要身份验证)" ) diff --git a/app/api/v1/endpoints/user_management.py b/app/api/v1/endpoints/user_management.py deleted file mode 100644 index 72e40f0..0000000 --- a/app/api/v1/endpoints/user_management.py +++ /dev/null @@ -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) diff --git a/app/api/v1/router.py b/app/api/v1/router.py index 52c5431..3e4e7d3 100644 --- a/app/api/v1/router.py +++ b/app/api/v1/router.py @@ -1,6 +1,6 @@ from fastapi import APIRouter from app.api.v1.endpoints import ( - auth, + agent_auth, project, simulation, scada, @@ -15,7 +15,6 @@ from app.api.v1.endpoints import ( leakage, burst_detection, burst_location, - user_management, # 新增:用户管理 audit, # 新增:审计日志 meta, web_search, @@ -54,10 +53,7 @@ from app.api.v1.endpoints.timeseries import ( api_router = APIRouter() # Core Services -api_router.include_router(auth.router, prefix="/auth", tags=["Auth"]) -api_router.include_router( - user_management.router, prefix="/users", tags=["User Management"] -) # 新增 +api_router.include_router(agent_auth.router, tags=["Agent Auth"]) api_router.include_router(audit.router, prefix="/audit", tags=["Audit Logs"]) # 新增 api_router.include_router(meta.router, tags=["Metadata"]) api_router.include_router(project.router, tags=["Project"]) diff --git a/app/auth/dependencies.py b/app/auth/dependencies.py deleted file mode 100644 index 3524f0a..0000000 --- a/app/auth/dependencies.py +++ /dev/null @@ -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 diff --git a/app/auth/keycloak_dependencies.py b/app/auth/keycloak_dependencies.py index 6b34936..ac43799 100644 --- a/app/auth/keycloak_dependencies.py +++ b/app/auth/keycloak_dependencies.py @@ -8,35 +8,41 @@ from jose import JWTError, jwt from app.core.config import settings oauth2_optional = OAuth2PasswordBearer( - tokenUrl=f"{settings.API_V1_STR}/auth/login", auto_error=False + tokenUrl="keycloak", auto_error=False ) # 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), -) -> UUID: +) -> dict: 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, - ) + return _decode_keycloak_token(token) except JWTError as exc: # logger.warning("Keycloak token validation failed: %s", exc) raise HTTPException( @@ -45,6 +51,10 @@ async def get_current_keycloak_sub( headers={"WWW-Authenticate": "Bearer"}, ) from exc + +async def get_current_keycloak_sub( + payload: dict = Depends(get_current_keycloak_payload), +) -> UUID: sub = payload.get("sub") if not sub: raise HTTPException( @@ -64,35 +74,8 @@ async def get_current_keycloak_sub( async def get_current_keycloak_username( - token: str | None = Depends(oauth2_optional), + payload: dict = Depends(get_current_keycloak_payload), ) -> 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") if not username: raise HTTPException( diff --git a/app/auth/permissions.py b/app/auth/permissions.py deleted file mode 100644 index 0fb8d1c..0000000 --- a/app/auth/permissions.py +++ /dev/null @@ -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 diff --git a/app/core/config.py b/app/core/config.py index 791d470..a4bac3f 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -11,14 +11,6 @@ class Settings(BaseSettings): 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) ENCRYPTION_KEY: str = "" # 必须从环境变量设置 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_MAX_SIZE: int = 10 - # Keycloak JWT (optional override) + # Keycloak access token verification KEYCLOAK_PUBLIC_KEY: str = "" KEYCLOAK_ALGORITHM: str = "RS256" KEYCLOAK_AUDIENCE: str = "" diff --git a/app/core/security.py b/app/core/security.py deleted file mode 100644 index a99e69f..0000000 --- a/app/core/security.py +++ /dev/null @@ -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) diff --git a/app/domain/models/role.py b/app/domain/models/role.py deleted file mode 100644 index 1870bf8..0000000 --- a/app/domain/models/role.py +++ /dev/null @@ -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] diff --git a/app/domain/schemas/user.py b/app/domain/schemas/user.py deleted file mode 100644 index 864035a..0000000 --- a/app/domain/schemas/user.py +++ /dev/null @@ -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") diff --git a/app/infra/audit/middleware.py b/app/infra/audit/middleware.py index d8f1e3e..d0774d4 100644 --- a/app/infra/audit/middleware.py +++ b/app/infra/audit/middleware.py @@ -33,8 +33,6 @@ class AuditMiddleware(BaseHTTPMiddleware): # 需要审计的路径前缀 AUDIT_PATHS = [ - # "/api/v1/auth/", - # "/api/v1/users/", # "/api/v1/projects/", # "/api/v1/networks/", ] @@ -193,20 +191,14 @@ class AuditMiddleware(BaseHTTPMiddleware): return None sub = None try: - key = ( - settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n") - if settings.KEYCLOAK_PUBLIC_KEY - else settings.SECRET_KEY - ) - algorithms = ( - [settings.KEYCLOAK_ALGORITHM] - if settings.KEYCLOAK_PUBLIC_KEY - else [settings.ALGORITHM] - ) + if not settings.KEYCLOAK_PUBLIC_KEY: + return None + + key = settings.KEYCLOAK_PUBLIC_KEY.replace("\\n", "\n") payload = jwt.decode( token, key, - algorithms=algorithms, + algorithms=[settings.KEYCLOAK_ALGORITHM], audience=settings.KEYCLOAK_AUDIENCE or None, ) sub = payload.get("sub") @@ -221,7 +213,7 @@ class AuditMiddleware(BaseHTTPMiddleware): keycloak_id = UUID(sub) user = await repo.get_user_by_keycloak_id(keycloak_id) except ValueError: - user = await repo.get_user_by_username(sub) + return None if user and user.is_active: return user.id return None diff --git a/app/infra/db/metadb/repositories/user_repository.py b/app/infra/db/metadb/repositories/user_repository.py deleted file mode 100644 index 4d975ec..0000000 --- a/app/infra/db/metadb/repositories/user_repository.py +++ /dev/null @@ -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 diff --git a/cli/tests/unit/test_tjwater_cli.py b/cli/tests/unit/test_tjwater_cli.py index 08a119f..8e16ec2 100644 --- a/cli/tests/unit/test_tjwater_cli.py +++ b/cli/tests/unit/test_tjwater_cli.py @@ -32,7 +32,6 @@ def test_load_auth_context_supports_aliases(monkeypatch): monkeypatch.setenv("TJWATER_SERVER", "http://server") monkeypatch.setenv("TJWATER_ACCESS_TOKEN", "abc") monkeypatch.setenv("TJWATER_PROJECT_ID", "p1") - monkeypatch.setenv("TJWATER_USER_ID", "u1") monkeypatch.setenv("TJWATER_USERNAME", "tester") monkeypatch.setenv("TJWATER_NETWORK", "net1") @@ -41,7 +40,6 @@ def test_load_auth_context_supports_aliases(monkeypatch): assert auth.server == "http://server" assert auth.access_token == "abc" assert auth.project_id == "p1" - assert auth.user_id == "u1" assert auth.username == "tester" 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_ACCESS_TOKEN", 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_NETWORK", raising=False) monkeypatch.delenv("TJWATER_EXTRA_HEADERS", raising=False) diff --git a/cli/tjwater_cli/core.py b/cli/tjwater_cli/core.py index a02c4dd..16d7bde 100644 --- a/cli/tjwater_cli/core.py +++ b/cli/tjwater_cli/core.py @@ -46,7 +46,6 @@ class AuthContext: server: str | None = None access_token: str | None = None project_id: str | None = None - user_id: str | None = None username: str | None = None network: str | None = None 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"), "access_token": os.getenv("TJWATER_ACCESS_TOKEN"), "project_id": os.getenv("TJWATER_PROJECT_ID"), - "user_id": os.getenv("TJWATER_USER_ID"), "username": os.getenv("TJWATER_USERNAME"), "network": os.getenv("TJWATER_NETWORK"), "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"), access_token=_pick(raw, "access_token", "token", "accessToken"), 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"), network=_pick(raw, "network", "project_code", "projectCode", "project"), 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) elif 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 diff --git a/cli/tjwater_cli_endpoint_scope.md b/cli/tjwater_cli_endpoint_scope.md index 1413674..a86dffa 100644 --- a/cli/tjwater_cli_endpoint_scope.md +++ b/cli/tjwater_cli_endpoint_scope.md @@ -306,10 +306,9 @@ app/api/v1/endpoints/snapshots.py app/api/v1/endpoints/cache.py app/api/v1/endpoints/audit.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 /getuser/ GET /getallusers/ -PUT /users/{user_id} -DELETE /users/{user_id} -POST /users/{user_id}/activate -POST /users/{user_id}/deactivate ``` ## Help diff --git a/tests/api/test_agent_auth_endpoints.py b/tests/api/test_agent_auth_endpoints.py new file mode 100644 index 0000000..1acda78 --- /dev/null +++ b/tests/api/test_agent_auth_endpoints.py @@ -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" diff --git a/tests/api/test_api_integration.py b/tests/api/test_api_integration.py index a5afbb9..27ec1eb 100755 --- a/tests/api/test_api_integration.py +++ b/tests/api/test_api_integration.py @@ -2,7 +2,7 @@ """ 测试新增 API 集成 -验证新的认证、用户管理和审计日志接口是否正确集成 +验证 Keycloak/metadata 认证和审计日志接口是否正确集成 """ import sys @@ -17,16 +17,15 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../. "module_name, desc", [ ("app.core.encryption", "加密模块"), - ("app.core.security", "安全模块"), ("app.core.audit", "审计模块"), - ("app.domain.models.role", "角色模型"), - ("app.domain.schemas.user", "用户Schema"), ("app.domain.schemas.audit", "审计Schema"), - ("app.auth.permissions", "权限控制"), - ("app.api.v1.endpoints.auth", "认证接口"), - ("app.api.v1.endpoints.user_management", "用户管理接口"), + ("app.auth.keycloak_dependencies", "Keycloak Token 校验"), + ("app.auth.metadata_dependencies", "Metadata 用户解析"), + ("app.auth.project_dependencies", "项目权限控制"), + ("app.api.v1.endpoints.agent_auth", "Agent 认证上下文接口"), + ("app.api.v1.endpoints.meta", "Metadata 接口"), ("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.audit.middleware", "审计中间件"), ], @@ -49,8 +48,8 @@ def test_router_configuration(): 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("/users" in r for r in routes), "缺少用户管理路由 (/users)" + assert any("/agent/auth/context" in r for r in routes), "缺少 Agent 认证上下文路由" + assert any("/meta" in r for r in routes), "缺少 Metadata 路由" assert any("/audit" in r for r in routes), "缺少审计日志路由 (/audit)" except Exception as e: diff --git a/tests/api/test_auth_endpoints.py b/tests/api/test_auth_endpoints.py deleted file mode 100644 index 144e442..0000000 --- a/tests/api/test_auth_endpoints.py +++ /dev/null @@ -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" diff --git a/tests/api/test_user_management_endpoints.py b/tests/api/test_user_management_endpoints.py deleted file mode 100644 index 8991d5a..0000000 --- a/tests/api/test_user_management_endpoints.py +++ /dev/null @@ -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 diff --git a/tests/auth/test_security.py b/tests/auth/test_security.py deleted file mode 100644 index 8953108..0000000 --- a/tests/auth/test_security.py +++ /dev/null @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py index d528272..812baa9 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -159,25 +159,6 @@ class FakeAsyncSession: 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): data = { "id": uuid4(), diff --git a/tests/unit/test_audit_repository.py b/tests/unit/test_audit_repository.py index 01a16bb..8397270 100644 --- a/tests/unit/test_audit_repository.py +++ b/tests/unit/test_audit_repository.py @@ -22,14 +22,14 @@ def test_create_log_adds_commits_and_refreshes(monkeypatch): result = asyncio.run( repo.create_log( - action="LOGIN", + action="CREATE_PROJECT", request_method="POST", - request_path="/auth/login", + request_path="/api/v1/projects", response_status=200, ) ) - assert result.action == "LOGIN" + assert result.action == "CREATE_PROJECT" assert result.request_method == "POST" assert session.commit_count == 1 assert len(session.added) == 1 diff --git a/tests/unit/test_auth_dependencies.py b/tests/unit/test_auth_dependencies.py deleted file mode 100644 index 0c556a7..0000000 --- a/tests/unit/test_auth_dependencies.py +++ /dev/null @@ -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." diff --git a/tests/unit/test_permissions.py b/tests/unit/test_permissions.py deleted file mode 100644 index 3dcd86a..0000000 --- a/tests/unit/test_permissions.py +++ /dev/null @@ -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" diff --git a/tests/unit/test_user_repository.py b/tests/unit/test_user_repository.py deleted file mode 100644 index 7d52ad8..0000000 --- a/tests/unit/test_user_repository.py +++ /dev/null @@ -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"}