diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 00000000..104dcc4c --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,30 @@ +name: CI + +on: + push: + branches: ["main", "dev"] + pull_request: + branches: ["main", "dev"] + +jobs: + lint-and-typecheck: + name: Lint & Type Check + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6 + - name: Install uv + uses: astral-sh/setup-uv@v7 + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version-file: "pyproject.toml" + - name: Install dependencies + run: uv sync --locked --all-extras --dev + - name: Lint + run: uv run flake8 + + - name: Type Check + run: uv run mypy . + + # - name: Run tests + # run: uv run pytest --maxfail=1 --disable-warnings \ No newline at end of file diff --git a/app/container.py b/app/container.py index 9b864210..d2f19f2c 100644 --- a/app/container.py +++ b/app/container.py @@ -1,12 +1,48 @@ -from fastapi import Depends import sqlalchemy.ext.asyncio +from fastapi import Depends from app.infra.database import get_db -from db.generated import user as user_queries,session as session_queries,devices as device_queries +from app.infra.redis import RedisClient +from db.generated import user as user_queries +from db.generated import session as session_queries +from db.generated import devices as device_queries +from app.service.users import AuthService +from app.service.session import SessionService +from app.service.device import DeviceService + + + +class Container: + + def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): + # infrastructure + self.redis = RedisClient.get_instance() + + # queriers + self.user_querier = user_queries.AsyncQuerier(conn) + self.session_querier = session_queries.AsyncQuerier(conn) + self.device_querier = device_queries.AsyncQuerier(conn) + + # services + self.session_service = SessionService() + self.session_service.init( + session=self.session_querier, + redis=self.redis, + ) + + self.device_service = DeviceService() + self.device_service.init( + device_querier=self.device_querier, + ) + + self.auth_service = AuthService( + user_querier=self.user_querier, + device_querier=self.device_querier, + session_querier=self.session_querier, + ) + -async def init_repo(conn: sqlalchemy.ext.asyncio.AsyncConnection = Depends(get_db)): - user_querier:user_queries.AsyncQuerier = user_queries.AsyncQuerier(conn) - session_querier:session_queries.AsyncQuerier = session_queries.AsyncQuerier(conn) - device_querier :device_queries.AsyncQuerier = device_queries.AsyncQuerier(conn) - return user_querier,session_querier,device_querier - \ No newline at end of file +async def get_container( + conn: sqlalchemy.ext.asyncio.AsyncConnection = Depends(get_db), +) -> Container: + return Container(conn) \ No newline at end of file diff --git a/app/core/exceptions.py b/app/core/exceptions.py index 8b37caaf..0006c207 100644 --- a/app/core/exceptions.py +++ b/app/core/exceptions.py @@ -20,6 +20,10 @@ def forbidden(detail: str = "Forbidden") -> HTTPException: @staticmethod def bad_request(detail: str = "Bad request") -> HTTPException: return HTTPException(status_code=400, detail=detail) + + @staticmethod + def payement_required(detail:str = "payement required")->HTTPException: + return HTTPException(status_code=402,detail=detail) @staticmethod def internal_error(detail: str = "Internal server error") -> HTTPException: diff --git a/app/core/securite.py b/app/core/securite.py index ad669185..610e1b4a 100644 --- a/app/core/securite.py +++ b/app/core/securite.py @@ -3,8 +3,8 @@ import jwt from passlib.context import CryptContext import pyotp -from core.config import settings -from core.exceptions import AppException +from app.core.config import settings +from app.core.exceptions import AppException pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") diff --git a/app/deps/auth.py b/app/deps/auth.py index b03a7dfe..21aac93b 100644 --- a/app/deps/auth.py +++ b/app/deps/auth.py @@ -1,57 +1,51 @@ -from typing import Any +from typing import Annotated import uuid -from fastapi import Depends +from fastapi import Depends, HTTPException from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer - -from app.core import constant -from app.core.exceptions import AppException +from pydantic import BaseModel +from app.container import get_container, Container from app.core.securite import decode_access_mobile_token -from app.infra.redis import RedisClient -from app.service.users import AuthService -from db.generated import user as user_queries -from db.generated import session as session_queries - security = HTTPBearer() - +class MobileUserSchema(BaseModel): + user_id: uuid.UUID + email: str + session_id: uuid.UUID + + async def get_current_mobile_user( - credentials: HTTPAuthorizationCredentials = Depends(security), -) -> dict[str, Any]: + credentials: Annotated[HTTPAuthorizationCredentials, Depends(security)], + container: Annotated[Container, Depends(get_container)], +) -> MobileUserSchema: + """ + Dependency to get the current logged-in mobile user. + Returns a strict Pydantic model. + """ token = credentials.credentials payload = decode_access_mobile_token(token) - session_id = payload.get("session_id") + session_id_str = payload.get("session_id") - if not session_id: - raise AppException.unauthorized("Invalid token") + if not session_id_str: + raise HTTPException(status_code=401, detail="Invalid token") - session_querier = session_queries.AsyncQuerier(conn) - session = await session_querier.get_session_by_id(id=uuid.UUID(session_id)) + session_id = uuid.UUID(session_id_str) + # Validate session via SessionService + session = await container.session_service.session_querier.get_session_by_id(id=session_id) if not session: - raise AppException.unauthorized("Session not found") - - if session.expires_at.replace(tzinfo=None) < payload.get("exp", 0): - raise AppException.unauthorized("Session expired") - - session_key = constant.RedisKey.UserSessionByUser.value.format( - user_id=session.user_id - ) - redis_session = await redis.get(session_key) - - if not redis_session or redis_session != session_id: - raise AppException.forbidden("Invalid session. Please login again.") - - await redis.expire(session_key, AuthService.REDIS_SESSION_TTL) + raise HTTPException(status_code=401, detail="Session not found") - user_querier = user_queries.AsyncQuerier(conn) - user = await user_querier.get_user_by_id(id=session.user_id) + exp_ts = payload.get("exp") + if exp_ts and session.expires_at.timestamp() < exp_ts: + raise HTTPException(status_code=401, detail="Session expired") + user = await container.auth_service.user_querier.get_user_by_id(id=session.user_id) if not user: - raise AppException.unauthorized("User not found") + raise HTTPException(status_code=401, detail="User not found") - return { - "user_id": str(user.id), - "email": user.email, - "session_id": session_id, - } + return MobileUserSchema( + user_id=user.id, + email=user.email, + session_id=session.id, + ) \ No newline at end of file diff --git a/app/infra/nats.py b/app/infra/nats.py index a5fe8729..0e6f93e9 100644 --- a/app/infra/nats.py +++ b/app/infra/nats.py @@ -47,11 +47,11 @@ async def publish(subject: NatsSubjects, message: bytes) -> None: async def subscribe(subject: NatsSubjects, callback: Callable[[Any], Any]) -> None: if NatsClient._nc is None: await NatsClient.connect() - + assert NatsClient._nc is not None async def _wrapper(msg:Msg): await callback(msg.data) - await NatsClient._nc.subscribe(subject.value, cb=_wrapper)#TODO:fix it here + await NatsClient._nc.subscribe(subject.value, cb=_wrapper)# type: ignore @staticmethod @@ -76,7 +76,7 @@ async def _wrapper(msg:Msg): await msg.ack() if NatsClient._js is None : print("no client ") - await NatsClient._js.subscribe( + await NatsClient._js.subscribe( # type: ignore subject=subject.value, stream=stream_name, durable=durable_name, diff --git a/app/infra/redis.py b/app/infra/redis.py index cf2d5965..7562f880 100644 --- a/app/infra/redis.py +++ b/app/infra/redis.py @@ -17,8 +17,8 @@ def __init__(self, host: str, port: int, password: str): f"redis://{host}:{port}", password=password, decode_responses=True ) - async def set(self, key: RedisKey | str, value: str, expire: int | None = None): - await self.client.set(key, value, ex=expire) + async def set(self, key: RedisKey | str, value: str, expire: int | None = None,nx:bool=False): + await self.client.set(key, value, ex=expire,nx=nx) async def get(self, key: RedisKey | str) -> str | None: return await self.client.get(key) @@ -32,5 +32,11 @@ async def exists(self, key: RedisKey | str) -> bool: async def expire(self, key: RedisKey | str, seconds: int): await self.client.expire(key, seconds) + @classmethod + def get_instance(cls) -> "RedisClient": + if cls._instance is None: + raise RuntimeError("RedisClient not initialized") + return cls._instance + async def close(self): await self.client.close() diff --git a/app/main.py b/app/main.py index d346f7c8..23ab8371 100644 --- a/app/main.py +++ b/app/main.py @@ -1,6 +1,10 @@ +import logging +import time from contextlib import asynccontextmanager - -from fastapi import FastAPI +import asyncio +from fastapi import FastAPI, Request, Response +from fastapi.middleware.cors import CORSMiddleware +from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint from app.core.config import settings from app.infra.minio import init_minio_client @@ -9,37 +13,104 @@ from app.router.mobile.auth import router as mobile_auth_router -app = FastAPI() +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s | %(levelname)s | %(name)s | %(message)s", +) +logger = logging.getLogger("api") -@app.get("/") -def read_root(): - return {"Hello": "World"} -@app.get("/health") -def health_check(): - return {"status": "healthy"} +class RequestLoggingMiddleware(BaseHTTPMiddleware): + async def dispatch( + self, + request: Request, + call_next: RequestResponseEndpoint, + ) -> Response: + + start_time = time.time() + + response = await call_next(request) + process_time = time.time() - start_time -app.include_router(mobile_auth_router, prefix="/mobile") + logger.info( + "%s %s status=%s time=%.3fs", + request.method, + request.url.path, + response.status_code, + process_time, + ) + return response + +MAX_RETRIES = 5 +RETRY_DELAY = 2 # seconds @asynccontextmanager async def lifespan(app: FastAPI): - await init_minio_client( - minio_host=settings.MINIO_HOST, - minio_port=settings.MINIO_API_PORT, - minio_root_user=settings.MINIO_ROOT_USER, - minio_root_password=settings.MINIO_ROOT_PASSWORD, - ) + + for attempt in range(1, MAX_RETRIES + 1): + try: + await init_minio_client( + minio_host=settings.MINIO_HOST, + minio_port=settings.MINIO_API_PORT, + minio_root_user=settings.MINIO_ROOT_USER, + minio_root_password=settings.MINIO_ROOT_PASSWORD, + ) + break + except Exception as e: + print(f"[MINIO] Attempt {attempt} failed: {e}") + if attempt == MAX_RETRIES: + raise RuntimeError("Cannot connect to MinIO after multiple attempts") from e + await asyncio.sleep(RETRY_DELAY) + RedisClient( host=settings.REDIS_HOST, port=settings.REDIS_PORT, password=settings.REDIS_PASSWORD, ) + await NatsClient.connect() + yield - await RedisClient.close(RedisClient._instance.client)#todo:fix for self + + await RedisClient.get_instance().close() await NatsClient.close() + + + +app = FastAPI( + title="multAI API", + description="Mobile and Web API for multAI", + version="1.0.0", + lifespan=lifespan, +) + + + +app.add_middleware(RequestLoggingMiddleware) + +app.add_middleware( + CORSMiddleware, + allow_origins=["*"], + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + + + +@app.get("/") +def read_root(): + return {"Hello": "World"} + + +@app.get("/health") +def health_check(): + return {"status": "healthy"} + + +app.include_router(mobile_auth_router, prefix="/mobile") \ No newline at end of file diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index b5252af4..54866bac 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -1,17 +1,14 @@ +from typing import Optional + from fastapi import APIRouter, Depends -import sqlalchemy.ext.asyncio +from uuid import UUID -from app.deps.auth import get_current_mobile_user -from app.infra.database import get_db -from app.infra.redis import RedisClient -from app.schema.auth.mobile.auth import ( - MobileAuthRequest, - MobileAuthResponse, - RefreshTokenRequest, - LogoutRequest, -) -from app.service.users import AuthService +from app.container import get_container, Container +from app.core.exceptions import AppException +from app.deps.auth import MobileUserSchema, get_current_mobile_user +from app.schema.request.mobile.auth import MobileAuthRequest, RefreshTokenRequest +from app.schema.response.mobile.auth import MeResponse, DeviceSchema, MobileAuthResponse, SessionSchema, UserSchema router = APIRouter(prefix="/auth", tags=["mobile-auth"]) @@ -19,31 +16,86 @@ @router.post("/register-login", response_model=MobileAuthResponse) async def mobile_register_login( req: MobileAuthRequest, - conn: sqlalchemy.ext.asyncio.AsyncConnection = Depends(get_db), - redis: RedisClient = Depends(), + container: Container = Depends(get_container), ): - return await AuthService.mobile_register_login(conn, redis, req) + + return await container.auth_service.mobile_register_login(container.redis, req) @router.post("/refresh", response_model=MobileAuthResponse) async def refresh_token( req: RefreshTokenRequest, - conn: sqlalchemy.ext.asyncio.AsyncConnection = Depends(get_db), - redis: RedisClient = Depends(), + container: Container = Depends(get_container), ): - return await AuthService.refresh_token(conn, redis, req.refresh_token) + + return await container.auth_service.refresh_token(container.redis, req.refresh_token) @router.post("/logout") async def logout( - req: LogoutRequest, - redis: RedisClient = Depends(), + container: Container = Depends(get_container), + User:MobileUserSchema = Depends(get_current_mobile_user) ): - return await AuthService.logout(redis, req.user_id, req.session_id) + return await container.auth_service.logout( + container.redis, + str(User.user_id), + str(User.session_id), + ) + + +@router.post("/revoke-device") +async def revoke_device( + device_id: UUID, + container: Container = Depends(get_container), + current_user:MobileUserSchema = Depends(get_current_mobile_user), +): -@router.get("/me") + await container.device_service.revoke_device( + device_id=device_id, + user_id=current_user.user_id, + ) + return {"message": "Device revoked successfully"} + + +@router.get("/me", response_model=MeResponse) async def get_me( - user: dict = Depends(get_current_mobile_user), + current_user:MobileUserSchema = Depends(get_current_mobile_user), + container: Container = Depends(get_container), ): - return user + + user = await container.auth_service.user_querier.get_user_by_id(id=current_user.user_id) + if user is None : + raise AppException.not_found("user not found") + + devices, _ = await container.device_service.get_all_devices(current_user.user_id) + device_list = [ + DeviceSchema( + id=d.id, + device_name=d.device_name or "uknown ", + device_type=d.device_type or "uknown ", + totp_secret=d.totp_secret, + ) + for d in devices + ] + + session_schema: Optional[SessionSchema] = None + sessions_objs = await container.session_service.session_querier.get_session_by_id( + id=current_user.session_id + ) + + if sessions_objs: + session_schema = SessionSchema( + session_id=sessions_objs.id, + device_id=sessions_objs.device_id, + last_active=sessions_objs.last_active, + expires_at=sessions_objs.expires_at, + ) + + + + return MeResponse( + user=UserSchema(id=user.id, email=user.email), + devices=device_list, + sessions=session_schema, + ) \ No newline at end of file diff --git a/app/schema/auth/mobile/__init__.py b/app/schema/auth/mobile/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/app/schema/auth/mobile/auth.py b/app/schema/request/mobile/auth.py similarity index 51% rename from app/schema/auth/mobile/auth.py rename to app/schema/request/mobile/auth.py index 14889657..c52520b5 100644 --- a/app/schema/auth/mobile/auth.py +++ b/app/schema/request/mobile/auth.py @@ -4,22 +4,15 @@ class MobileAuthRequest(BaseModel): email: EmailStr password: str - device_id: str device_name: str device_type: str + device_id: str -class RefreshTokenRequest(BaseModel): - refresh_token: str -class MobileAuthResponse(BaseModel): - access_token: str + +class RefreshTokenRequest(BaseModel): refresh_token: str - session_id: str - expires_in: int -class LogoutRequest(BaseModel): - session_id: str - user_id: str diff --git a/app/schema/response/mobile/auth.py b/app/schema/response/mobile/auth.py new file mode 100644 index 00000000..89c9ec51 --- /dev/null +++ b/app/schema/response/mobile/auth.py @@ -0,0 +1,33 @@ +from typing import List, Optional +from pydantic import BaseModel +import uuid +from datetime import datetime + +class DeviceSchema(BaseModel): + id: uuid.UUID + device_name: str + device_type: str + totp_secret: str | None + +class SessionSchema(BaseModel): + session_id: uuid.UUID + device_id: uuid.UUID + last_active: datetime + expires_at: datetime + +class UserSchema(BaseModel): + id: uuid.UUID + email: str + +class MeResponse(BaseModel): + user: UserSchema + devices: List[DeviceSchema] + sessions: Optional[SessionSchema] + + + +class MobileAuthResponse(BaseModel): + access_token: str + refresh_token: str + session_id: str + expires_in: int diff --git a/app/service/device.py b/app/service/device.py index e828a675..ad1b557e 100644 --- a/app/service/device.py +++ b/app/service/device.py @@ -1,9 +1,7 @@ -from typing import AsyncIterator - from db.generated import devices as device_queries from app.core.securite import create_totp_secret import uuid -from app.core.exceptions import DBException,AppException +from app.core.exceptions import DBException,AppException, DBExceptionImpl from db.generated.models import UserDevice @@ -15,9 +13,9 @@ def init(self, device_querier: device_queries.AsyncQuerier): @staticmethod async def create_device(user_id: uuid.UUID,device_name: str,device_type: str)->UserDevice|None: try : - if await DeviceService.device_querier.count_user_devices(user_id=user_id) >= 3: + DeviceCount = await DeviceService.count_devices(user_id=user_id) + if DeviceCount >=3: raise AppException.bad_request("You can only have 3 devices") - return await DeviceService.device_querier.create_device( user_id=user_id, device_name=device_name, @@ -39,24 +37,34 @@ async def revoke_device(device_id:uuid.UUID,user_id :uuid.UUID): DBException.handle(e) @staticmethod - async def get_all_devices(user_id:uuid.UUID)->tuple(AsyncIterator[UserDevice],int) - devices= DeviceService.device_querier.list_user_devices(user_id=user_id) - count = await DeviceService.device_querier.count_user_devices(user_id=user_id) - return devices,count + async def get_all_devices(user_id: uuid.UUID) -> tuple[list[UserDevice], int]: + devices: list[UserDevice] = [] + + async for device in DeviceService.device_querier.list_user_devices(user_id=user_id): + devices.append(device) + + count = await DeviceService.count_devices(user_id=user_id) + + return devices, count @staticmethod async def get_device_by_id(device_id:uuid.UUID,user_id:uuid.UUID)->UserDevice|None: try : - return await DeviceService.device_querier.get_device_by_id(id=device_id,user_id=user_id) + device = await DeviceService.device_querier.get_device__by_id(id=device_id) + if device is None : + raise AppException.not_found("device not found ") except Exception as e : - DBException.handle(e) + raise DBExceptionImpl.handle(e) @staticmethod async def count_devices(user_id:uuid.UUID)->int: try : - return await DeviceService.device_querier.count_user_devices(user_id=user_id) + count = await DeviceService.device_querier.count__user__devices(user_id=user_id) + if count is None : + raise AppException.internal_error("db failed to count ") + return count except Exception as e : - DBException.handle(e) + raise DBExceptionImpl.handle(e) diff --git a/app/service/session.py b/app/service/session.py index e0039e42..c500594d 100644 --- a/app/service/session.py +++ b/app/service/session.py @@ -1,16 +1,151 @@ +from pydantic import BaseModel +from app.core.exceptions import AppException, DBExceptionImpl from db.generated import session as session_queries -from app.core.exceptions import DBException import uuid -from db.generated.models import Session as Session_querier +from db.generated.models import UserSession +from datetime import datetime,timedelta,timezone +from app.infra.redis import RedisClient +from app.core.constant import RedisKey +from db.generated.session import UpsertSessionRow + +class SessionRedis(BaseModel): + session_id:uuid.UUID + user_id:uuid.UUID + device_id:uuid.UUID + last_active:datetime + expires_at:datetime + class SessionService : session_querier : session_queries.AsyncQuerier + redis : RedisClient - def init(self,session:session_queries.AsyncQuerier): + def init(self,session:session_queries.AsyncQuerier,redis:RedisClient): self.session_querier = session + self.redis = redis + + @staticmethod + async def create_session(user_id:uuid.UUID,device_id:uuid.UUID)->UpsertSessionRow: + try : + session = await SessionService.session_querier.upsert_session( + user_id=user_id, + device_id=device_id, + expires_at=datetime.now(timezone.utc) + timedelta(days=7), + ) + if session is None : + raise AppException.internal_error("session creation failed ") + + result = await SessionService.redis.set( + key=RedisKey.UserSessionByUser.format(user_id=user_id), + value=SessionRedis( + session_id=session.id, + user_id=session.user_id, + device_id=session.device_id, + last_active=session.last_active, + expires_at=session.expires_at, + ).model_dump_json(), + expire=60*60*5, + nx=True + ) + if not result: + AppException.forbidden("You already logged in in another device") + return session + except Exception as e : + raise DBExceptionImpl.handle(e) + + + + @staticmethod + async def get_session_by_id(session_id:uuid.UUID)->UserSession: + try : + session = await SessionService.session_querier.get_session_by_id(id=session_id) + if session is None : + raise AppException.not_found("session Not found ") + return session + except Exception as e : + raise DBExceptionImpl.handle(e) + + + @staticmethod + async def check_session( + session_id: uuid.UUID, + user_id: uuid.UUID, + device_id: uuid.UUID + ) -> bool: + try: + session_in_redis = await SessionService.redis.get( + RedisKey.UserSessionByUser.format(user_id=user_id) + ) + + if session_in_redis is None: + return False + + session_info = SessionRedis.model_validate_json(session_in_redis) + + if session_info: + if session_info.device_id != device_id and session_info.session_id != session_id: + raise AppException.forbidden("You already logged in on another device") + + await SessionService.redis.set( + key=RedisKey.UserSessionByUser.format(user_id=user_id), + value=SessionRedis( + session_id=session_info.session_id, + user_id=session_info.user_id, + device_id=session_info.device_id, + last_active=session_info.last_active, + expires_at=session_info.expires_at, + ).model_dump_json(), + expire=60 * 60 * 5, + nx=False, + ) + + return True + + session = await SessionService.session_querier.get_session_by_id(id=session_id) + + if session is None: + raise AppException.forbidden("Session not found") + + await SessionService.redis.set( + key=RedisKey.UserSessionByUser.format(user_id=user_id), + value=SessionRedis( + session_id=session.id, + user_id=session.user_id, + device_id=session.device_id, + last_active=session.last_active, + expires_at=session.expires_at, + ).model_dump_json(), + expire=60 * 60 * 5, + nx=True, + ) + + return True + + except Exception as e: + raise DBExceptionImpl.handle(e) + + + @staticmethod + async def delete_session(session_id:uuid.UUID,user_id:uuid.UUID,device_id : uuid.UUID): + try : + await SessionService.session_querier.delete_session_by_device(user_id=user_id,device_id=device_id) + except Exception as e : + raise DBExceptionImpl.handle(e) + + + @staticmethod + async def delete_expired_sessions(): + try : + await SessionService.session_querier.delete_expired_sessions() + except Exception as e : + raise DBExceptionImpl.handle(e) @staticmethod - async def create_session(user_id:uuid.UUID,device_id:uuid.UUID)->Session: + async def count_user_sessions(user_id:uuid.UUID)->int: try : - return await SessionService.session_querier. + count = await SessionService.session_querier.count_user_sessions(user_id=user_id) + if count is None : + raise AppException.internal_error("failed to count ") + else : + return count except Exception as e : - DBException.handle(e) \ No newline at end of file + raise DBExceptionImpl.handle(e) diff --git a/app/service/users.py b/app/service/users.py index af7cca7c..a920e9e2 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -1,7 +1,5 @@ from datetime import datetime, timedelta, timezone -from typing import Any import uuid -import sqlalchemy.ext.asyncio from app.core import constant from app.core.exceptions import AppException @@ -14,10 +12,9 @@ Get_expiry_time, ) from app.infra.redis import RedisClient -from app.schema.auth.mobile.auth import ( - MobileAuthRequest, - MobileAuthResponse, -) + +from app.schema.request.mobile.auth import MobileAuthRequest +from app.schema.response.mobile.auth import MobileAuthResponse from db.generated import user as user_queries from db.generated import devices as device_queries from db.generated import session as session_queries @@ -40,9 +37,8 @@ def __init__( self.device_querier = device_querier self.session_querier = session_querier - @staticmethod async def mobile_register_login( - conn: sqlalchemy.ext.asyncio.AsyncConnection, + self, redis: RedisClient, req: MobileAuthRequest, ) -> MobileAuthResponse: @@ -54,13 +50,13 @@ async def mobile_register_login( user = existing_user else: hashed = hash_password(req.password) - user = await user_querier.create_user( + user = await self.user_querier.create_user( email=req.email, hashed_password=hashed ) if not user: raise AppException.internal_error("Failed to create user") - user_id = user.id + user_id: uuid.UUID = user.id session_key = constant.RedisKey.UserSessionByUser.value.format(user_id=user_id) if await redis.exists(session_key): @@ -103,13 +99,12 @@ async def mobile_register_login( return MobileAuthResponse( access_token=access_token, refresh_token=refresh_token, - session_id=str(session.id), + session_id=str(), expires_in=expiry, ) - @staticmethod async def refresh_token( - conn: sqlalchemy.ext.asyncio.AsyncConnection, + self, redis: RedisClient, refresh_token: str, ) -> MobileAuthResponse: @@ -119,8 +114,7 @@ async def refresh_token( if not session_id: raise AppException.unauthorized("Invalid refresh token") - session_querier = session_queries.AsyncQuerier(conn) - session = await session_querier.get_session_by_id(id=uuid.UUID(session_id)) + session = await self.session_querier.get_session_by_id(id=uuid.UUID(session_id)) if not session: raise AppException.unauthorized("Session not found") @@ -149,8 +143,8 @@ async def refresh_token( expires_in=expiry, ) - @staticmethod async def logout( + self, redis: RedisClient, user_id: str, session_id: str, @@ -159,14 +153,12 @@ async def logout( await redis.delete(session_key) return {"message": "Logged out successfully"} - @staticmethod async def validate_session( - conn: sqlalchemy.ext.asyncio.AsyncConnection, + self, redis: RedisClient, session_id: str, ) -> bool: - session_querier = session_queries.AsyncQuerier(conn) - session = await session_querier.get_session_by_id(id=uuid.UUID(session_id)) + session = await self.session_querier.get_session_by_id(id=uuid.UUID(session_id)) if not session: return False diff --git a/db/generated/devices.py b/db/generated/devices.py index 6763065d..a92c73a2 100644 --- a/db/generated/devices.py +++ b/db/generated/devices.py @@ -8,7 +8,14 @@ import sqlalchemy import sqlalchemy.ext.asyncio -from generated import models +from . import models + + +COUNT__USER__DEVICES = """-- name: count__user__devices \\:one +SELECT COUNT(*) +FROM user_devices +WHERE user_id = :p1 +""" CREATE_DEVICE = """-- name: create_device \\:one @@ -33,6 +40,12 @@ """ +GET_DEVICE__BY_ID = """-- name: get_device__by_id \\:one +SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at from user_devices +WHERE id =:p1 +""" + + LIST_USER_DEVICES = """-- name: list_user_devices \\:many SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at FROM user_devices @@ -59,6 +72,12 @@ class AsyncQuerier: def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): self._conn = conn + async def count__user__devices(self, *, user_id: uuid.UUID) -> Optional[int]: + row = (await self._conn.execute(sqlalchemy.text(COUNT__USER__DEVICES), {"p1": user_id})).first() + if row is None: + return None + return row[0] + async def create_device(self, *, user_id: uuid.UUID, device_name: Optional[str], device_type: Optional[str], totp_secret: Optional[str]) -> Optional[models.UserDevice]: row = (await self._conn.execute(sqlalchemy.text(CREATE_DEVICE), { "p1": user_id, @@ -82,6 +101,21 @@ async def create_device(self, *, user_id: uuid.UUID, device_name: Optional[str], async def enable_device2_fa(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(ENABLE_DEVICE2_FA), {"p1": id, "p2": user_id}) + async def get_device__by_id(self, *, id: uuid.UUID) -> Optional[models.UserDevice]: + row = (await self._conn.execute(sqlalchemy.text(GET_DEVICE__BY_ID), {"p1": id})).first() + if row is None: + return None + return models.UserDevice( + id=row[0], + user_id=row[1], + device_name=row[2], + device_type=row[3], + totp_secret=row[4], + is_2fa_enabled=row[5], + last_active=row[6], + created_at=row[7], + ) + async def list_user_devices(self, *, user_id: uuid.UUID) -> AsyncIterator[models.UserDevice]: result = await self._conn.stream(sqlalchemy.text(LIST_USER_DEVICES), {"p1": user_id}) async for row in result: diff --git a/db/generated/session.py b/db/generated/session.py index 8ebc29f3..f45707d7 100644 --- a/db/generated/session.py +++ b/db/generated/session.py @@ -2,6 +2,7 @@ # versions: # sqlc v1.30.0 # source: session.sql +import dataclasses import datetime from typing import Optional import uuid @@ -9,7 +10,7 @@ import sqlalchemy import sqlalchemy.ext.asyncio -from generated import models +from . import models COUNT_USER_SESSIONS = """-- name: count_user_sessions \\:one @@ -69,10 +70,26 @@ DO UPDATE SET last_active = NOW(), expires_at = EXCLUDED.expires_at -RETURNING id, user_id, device_id, created_at, last_active, expires_at +RETURNING + id, + user_id, + device_id, + last_active, + expires_at, + created_at """ +@dataclasses.dataclass() +class UpsertSessionRow: + id: uuid.UUID + user_id: uuid.UUID + device_id: uuid.UUID + last_active: datetime.datetime + expires_at: datetime.datetime + created_at: datetime.datetime + + class AsyncQuerier: def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): self._conn = conn @@ -121,15 +138,15 @@ async def get_session_by_id(self, *, id: uuid.UUID) -> Optional[models.UserSessi async def update_session_activity(self, *, id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(UPDATE_SESSION_ACTIVITY), {"p1": id}) - async def upsert_session(self, *, user_id: uuid.UUID, device_id: uuid.UUID, expires_at: datetime.datetime) -> Optional[models.UserSession]: + async def upsert_session(self, *, user_id: uuid.UUID, device_id: uuid.UUID, expires_at: datetime.datetime) -> Optional[UpsertSessionRow]: row = (await self._conn.execute(sqlalchemy.text(UPSERT_SESSION), {"p1": user_id, "p2": device_id, "p3": expires_at})).first() if row is None: return None - return models.UserSession( + return UpsertSessionRow( id=row[0], user_id=row[1], device_id=row[2], - created_at=row[3], - last_active=row[4], - expires_at=row[5], + last_active=row[3], + expires_at=row[4], + created_at=row[5], ) diff --git a/db/generated/stuff_user.py b/db/generated/stuff_user.py index e5b08e23..437f73aa 100644 --- a/db/generated/stuff_user.py +++ b/db/generated/stuff_user.py @@ -8,7 +8,7 @@ import sqlalchemy import sqlalchemy.ext.asyncio -from generated import models +from . import models CREATE_ADMIN = """-- name: create_admin \\:one diff --git a/db/generated/user.py b/db/generated/user.py index b946fa39..02258d86 100644 --- a/db/generated/user.py +++ b/db/generated/user.py @@ -8,7 +8,7 @@ import sqlalchemy import sqlalchemy.ext.asyncio -from generated import models +from . import models CREATE_USER = """-- name: create_user \\:one diff --git a/db/queries/devices.sql b/db/queries/devices.sql index f37dac74..525a041e 100644 --- a/db/queries/devices.sql +++ b/db/queries/devices.sql @@ -31,4 +31,13 @@ UPDATE user_devices SET is_2fa_enabled = TRUE WHERE id = $1 AND user_id = $2 -AND is_2fa_enabled = FALSE; \ No newline at end of file +AND is_2fa_enabled = FALSE; + +-- name: Get_device_By_id :one +SELECT * from user_devices +WHERE id =$1; + +-- name: Count_User_Devices :one +SELECT COUNT(*) +FROM user_devices +WHERE user_id = $1; \ No newline at end of file diff --git a/db/queries/session.sql b/db/queries/session.sql index 51f42798..2a5b859c 100644 --- a/db/queries/session.sql +++ b/db/queries/session.sql @@ -10,8 +10,13 @@ ON CONFLICT (user_id, device_id) DO UPDATE SET last_active = NOW(), expires_at = EXCLUDED.expires_at -RETURNING *; - +RETURNING + id, + user_id, + device_id, + last_active, + expires_at, + created_at; -- name: GetSessionByDevice :one SELECT * @@ -43,4 +48,4 @@ DELETE FROM user_sessions WHERE expires_at < NOW(); -- name: CountUserSessions :one -SELECT COUNT(*) FROM user_sessions WHERE user_id = $1; \ No newline at end of file +SELECT COUNT(*) FROM user_sessions WHERE user_id = $1; diff --git a/docker-compose.yml b/docker-compose.yml index 25fe7db0..584386ce 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -13,7 +13,7 @@ services: POSTGRES_PASSWORD: ${POSTGRES_PASSWORD} POSTGRES_DB: ${POSTGRES_DB} ports: - - "${POSTGRES_PORT}:5432" + - "5432:5432" volumes: - postgres_data:/var/lib/postgresql/data networks: @@ -24,6 +24,7 @@ services: container_name: multi_nats command: > -js + -a 0.0.0.0 -m ${NATS_MONITOR_PORT} ports: - "${NATS_PORT}:4222" @@ -57,6 +58,7 @@ services: - "${PGADMIN_PORT}:80" networks: - multi_network + redis: image: redis:7-alpine container_name: multi_redis @@ -66,11 +68,13 @@ services: networks: - multi_network volumes: - - redis_data:/data + - redis_data:/data volumes: postgres_data: minio_data: + redis_data: + networks: multi_network: diff --git a/app/schema/__init__.py b/migrations/sql/down/create_photo_table.sql similarity index 100% rename from app/schema/__init__.py rename to migrations/sql/down/create_photo_table.sql diff --git a/migrations/sql/up/create_staff_user.sql b/migrations/sql/up/create_staff_user.sql index ba05ff77..311889c5 100644 --- a/migrations/sql/up/create_staff_user.sql +++ b/migrations/sql/up/create_staff_user.sql @@ -1,4 +1,4 @@ -CREATE TYPE staff_role AS ENUM ('admin', 'multi'); +CREATE TYPE staff_role AS ENUM ('admin','multi_team_lead', 'multi'); CREATE TABLE IF NOT EXISTS staff_users ( id UUID PRIMARY KEY DEFAULT uuid_generate_v4(), diff --git a/pyproject.toml b/pyproject.toml index dea7a500..faac2c9f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,8 +6,10 @@ readme = "README.md" requires-python = ">=3.12" dependencies = [ "alembic>=1.18.4", + "asyncpg>=0.31.0", "cryptography>=46.0.5", "fastapi[standard]>=0.135.1", + "greenlet>=3.3.2", "miniopy-async>=1.23.4", "nats-py>=2.14.0", "passlib[bcrypt]>=1.7.4", @@ -19,3 +21,22 @@ dependencies = [ "redis>=7.2.1", "setuptools>=82.0.0", ] + +[tool.ruff] +src = ["app", "migrations", "db"] +line-length = 88 +select = ["E", "F", "W", "C90"] +ignore = ["E501"] +exclude = ["__pycache__", ".venv", ".git", "build", "dist", ".ruff_cache"] + +[tool.mypy] +python_version = "3.12" +check_untyped_defs = true +disallow_untyped_defs = true +disallow_incomplete_defs = true +ignore_missing_imports = false +strict_optional = true +warn_unused_ignores = true +warn_redundant_casts = true +follow_imports = "silent" +files = ["app", "db", "migrations"]