diff --git a/core/api.py b/core/api.py index 77fd2ac..d47a476 100644 --- a/core/api.py +++ b/core/api.py @@ -1,16 +1,95 @@ -from ninja import NinjaAPI, Schema -from typing import List -from datetime import date from decimal import Decimal +from datetime import datetime, timedelta, date +from typing import List +import jwt +from django.conf import settings +from django.contrib.auth import authenticate +from django.contrib.auth.models import User +from django.shortcuts import get_object_or_404 +from ninja import NinjaAPI, Schema +from ninja.errors import HttpError +from ninja.security import HttpBearer from core.models import Transaction, Category -api = NinjaAPI() +api = NinjaAPI(title="JAST API", version="1.0.0") + + +class JWTAuth(HttpBearer): + def authenticate(self, request, token): + try: + payload = jwt.decode(token, settings.SECRET_KEY, algorithms=["HS256"]) + if payload.get("type") != "access": + return None + user = User.objects.get(id=payload["user_id"]) + return user + except (jwt.PyJWTError, User.DoesNotExist): + return None + + +auth = JWTAuth() + + +def generate_access_token(user: User) -> str: + payload = { + "user_id": user.id, + "type": "access", + "exp": datetime.utcnow() + timedelta(minutes=30), + "iat": datetime.utcnow(), + } + return jwt.encode(payload, settings.SECRET_KEY, algorithm="HS256") + + +def generate_refresh_token(user: User) -> str: + payload = { + "user_id": user.id, + "type": "refresh", + "exp": datetime.utcnow() + timedelta(days=14), + "iat": datetime.utcnow(), + } + return jwt.encode(payload, settings.SECRET_KEY, algorithm="HS256") + + +# Schemas +class RegisterIn(Schema): + username: str + password: str + email: str | None = None + first_name: str | None = None + last_name: str | None = None + + +class LoginIn(Schema): + username: str + password: str + + +class RefreshIn(Schema): + refresh_token: str + + +class RefreshOut(Schema): + access_token: str + + +class UserOut(Schema): + id: int + username: str + email: str | None = None + first_name: str | None = None + last_name: str | None = None + + +class TokenOut(Schema): + access_token: str + refresh_token: str + user: UserOut + -# --- Schemas (Data validation for Flutter) --- class CategorySchema(Schema): id: int name: str + class TransactionOut(Schema): id: int title: str @@ -18,27 +97,111 @@ class TransactionOut(Schema): date: date category: CategorySchema | None = None + class TransactionIn(Schema): title: str amount: Decimal date: date category_id: int | None = None -# --- Endpoints --- -@api.get("/transactions", response=List[TransactionOut]) -def list_transactions(request): - return Transaction.objects.select_related('category').all().order_by('-date') -@api.post("/transactions", response=TransactionOut) +class SuccessResponse(Schema): + success: bool + + +# Endpoints +@api.post("/auth/register", response=TokenOut) +def register(request, payload: RegisterIn): + if User.objects.filter(username=payload.username).exists(): + raise HttpError(400, "Username already taken") + + user = User.objects.create_user( + username=payload.username, + password=payload.password, + email=payload.email or "", + first_name=payload.first_name or "", + last_name=payload.last_name or "", + ) + access_token = generate_access_token(user) + refresh_token = generate_refresh_token(user) + return { + "access_token": access_token, + "refresh_token": refresh_token, + "user": user, + } + + +@api.post("/auth/login", response=TokenOut) +def login_view(request, payload: LoginIn): + user = authenticate(username=payload.username, password=payload.password) + if user is None: + raise HttpError(401, "Invalid username or password") + + access_token = generate_access_token(user) + refresh_token = generate_refresh_token(user) + return { + "access_token": access_token, + "refresh_token": refresh_token, + "user": user, + } + + +@api.post("/auth/refresh", response=RefreshOut) +def refresh_token_view(request, payload: RefreshIn): + try: + data = jwt.decode(payload.refresh_token, settings.SECRET_KEY, algorithms=["HS256"]) + if data.get("type") != "refresh": + raise HttpError(401, "Invalid token type") + user = User.objects.get(id=data["user_id"]) + new_access_token = generate_access_token(user) + return {"access_token": new_access_token} + except (jwt.PyJWTError, User.DoesNotExist): + raise HttpError(401, "Invalid or expired refresh token") + + +@api.get("/auth/me", response=UserOut, auth=auth) +def get_current_user(request): + return request.auth + + +@api.get("/transactions", response=List[TransactionOut], auth=auth) +def list_transactions(request): + return ( + Transaction.objects.filter(user=request.auth) + .select_related("category") + .order_by("-date") + ) + + +@api.post("/transactions", response=TransactionOut, auth=auth) def create_transaction(request, payload: TransactionIn): - transaction = Transaction.objects.create( + return Transaction.objects.create( + user=request.auth, title=payload.title, amount=payload.amount, date=payload.date, - category_id=payload.category_id + category_id=payload.category_id, ) + + +@api.put("/transactions/{transaction_id}", response=TransactionOut, auth=auth) +def update_transaction(request, transaction_id: int, payload: TransactionIn): + transaction = get_object_or_404(Transaction, id=transaction_id, user=request.auth) + transaction.title = payload.title + transaction.amount = payload.amount + transaction.date = payload.date + transaction.category_id = payload.category_id + transaction.save() return transaction -@api.get("/categories", response=List[CategorySchema]) + +@api.delete("/transactions/{transaction_id}", response=SuccessResponse, auth=auth) +def delete_transaction(request, transaction_id: int): + transaction = get_object_or_404(Transaction, id=transaction_id, user=request.auth) + transaction.delete() + return {"success": True} + + +@api.get("/categories", response=List[CategorySchema], auth=auth) def list_categories(request): return Category.objects.all() diff --git a/core/settings.py b/core/settings.py index 71e4eb3..521d8c7 100644 --- a/core/settings.py +++ b/core/settings.py @@ -140,3 +140,7 @@ STATIC_URL = 'static/' STATICFILES_DIRS = [ BASE_DIR / 'static', ] +SESSION_ENGINE = "django.contrib.sessions.backends.db" +SESSION_COOKIE_AGE = 1209600 # Persist session for 2 weeks (in seconds) +SESSION_SAVE_EVERY_REQUEST = True +SESSION_COOKIE_HTTPONLY = True diff --git a/pyproject.toml b/pyproject.toml index d42b2ac..d859a96 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,5 +8,6 @@ dependencies = [ "django>=6.1.1", "django-ninja>=1.7.1", "gunicorn>=26.2.0", + "pyjwt>=2.14.0", "whitenoise>=6.12.0", ] diff --git a/uv.lock b/uv.lock index dd1adf4..b33ab74 100644 --- a/uv.lock +++ b/uv.lock @@ -64,6 +64,7 @@ dependencies = [ { name = "django" }, { name = "django-ninja" }, { name = "gunicorn" }, + { name = "pyjwt" }, { name = "whitenoise" }, ] @@ -72,6 +73,7 @@ requires-dist = [ { name = "django", specifier = ">=6.1.1" }, { name = "django-ninja", specifier = ">=1.7.1" }, { name = "gunicorn", specifier = ">=26.2.0" }, + { name = "pyjwt", specifier = ">=2.14.0" }, { name = "whitenoise", specifier = ">=6.12.0" }, ] @@ -131,6 +133,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/63/19/67830dda664e6bdf9285ee2e40f355d0d7d6b92aa0c42e8d217bb8d33d36/pydantic_core-2.46.5-cp314-cp314t-win_arm64.whl", hash = "sha256:acf8a67ba51f4ca9ddbd0e6b3000a65ac51ab734661778b3e7ba64d99a710f2f", size = 1989276, upload-time = "2026-08-28T10:00:16.984Z" }, ] +[[package]] +name = "pyjwt" +version = "2.14.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/af/c3/8a3b59c25070cc61dc517fbdfa5dc0904670c96f605cc69759dc09166b99/pyjwt-2.14.0.tar.gz", hash = "sha256:77283c83fb56ecf566a886c757a714bc83668e38156de2cce8263302f42e0b86", size = 113177, upload-time = "2026-09-11T13:11:54.638Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9c/97/672cb32ce0dfea44b740cb7b4f97038463b9cf7c0ead1aacf595572851d6/pyjwt-2.14.0-py3-none-any.whl", hash = "sha256:ad0cef71c756a56e74863c2919cf0985f72decbcfcb550ee2f422e7c62b5eedc", size = 32896, upload-time = "2026-09-11T13:11:53.409Z" }, +] + [[package]] name = "sqlparse" version = "0.6.0"