diff --git a/apps/gooyal_oauth2/rest_framework.py b/apps/gooyal_oauth2/rest_framework.py index 3a08f0e..a8a30a9 100644 --- a/apps/gooyal_oauth2/rest_framework.py +++ b/apps/gooyal_oauth2/rest_framework.py @@ -1,8 +1,9 @@ import logging +from django.core.exceptions import ImproperlyConfigured from oauth2_provider.contrib.rest_framework import TokenMatchesOASRequirements, OAuth2Authentication from rest_framework.permissions import ( - IsAuthenticated + IsAuthenticated, BasePermission ) logger = logging.getLogger("oauth2_provider") @@ -25,3 +26,66 @@ class IsAuthenticatedOrTokenMatchesOASRequirements(TokenMatchesOASRequirements): result = (is_authenticated and not oauth2authenticated) or token_has_scope.has_permission(request, view) logger.debug(f'authentication result: {result}') return result + + +class TokenMatchesViewSetActions(BasePermission): + """ + :attr:action_required_scopes: dict keyed by view set action name with value: iterable action scope lists + + This fulfills the [Open API Specification (OAS; formerly Swagger)](https://www.openapis.org/) + list of alternative Security Requirements Objects for oauth2 or openIdConnect: + When a list of Security Requirement Objects is defined on the Open API object or Operation Object, + only one of Security Requirement Objects in the list needs to be satisfied to authorize the request. + [1](https://github.com/OAI/OpenAPI-Specification/blob/master/versions/3.0.0.md#securityRequirementObject) + + For each method, a list of lists of allowed scopes is tried in order and the first to match succeeds. + + @example + required_action_scopes = { + 'list': [['read']], + 'create': [['create1','scope2'], ['alt-scope3'], ['alt-scope4','alt-scope5']], + } + + TODO: DRY: subclass TokenHasScope and iterate over values of required_scope? + """ + + def has_permission(self, request, view): + token = request.auth + + if not token: + return False + + if hasattr(token, "scope"): # OAuth 2 + required_action_scopes = self.get_required_action_scopes(request, view) + + # m = request.method.upper() + a = view.action + if a in required_action_scopes: + logger.debug( + "Required scopes alternatives to access resource: {0}".format( + required_action_scopes[a] + ) + ) + for alt in required_action_scopes[a]: + if token.is_valid(alt): + return True + return False + else: + logger.warning("no scope action defined for action {0}".format(a)) + return False + + assert False, ( + "TokenMatchesViewSetActions requires the" + "`oauth2_provider.rest_framework.OAuth2Authentication` authentication " + "class to be used." + ) + + def get_required_action_scopes(self, request, view): + try: + return getattr(view, "required_action_scopes") + except AttributeError: + raise ImproperlyConfigured( + "TokenMatchesViewSetActions requires the view to" + " define the required_action_scopes attribute" + ) + diff --git a/apps/promotions/urls_application.py b/apps/promotions/urls_application.py new file mode 100644 index 0000000..4e6f095 --- /dev/null +++ b/apps/promotions/urls_application.py @@ -0,0 +1,60 @@ +from rest_framework import routers +from rest_framework.routers import DefaultRouter, DynamicRoute, Route +from . import views_application +from django.urls import NoReverseMatch, path, re_path + + +app_name = 'promotions-application' + + +class ApplicationRouter(DefaultRouter): + + routes = [ + # List route. + Route( + url=r'^{prefix}/{trailing_slash}$', + mapping={ + 'get': 'list', + 'post': 'create' + }, + name='{basename}-list', + detail=False, + initkwargs={'suffix': 'List'} + ), + # Dynamically generated list routes. Generated using + # @action(detail=False) decorator on methods of the viewset. + DynamicRoute( + url=r'^{prefix}//{url_path}{trailing_slash}$', + name='{basename}-{url_name}', + detail=False, + initkwargs={} + ), + # # Detail route. + Route( + url=r'^{prefix}//{lookup}{trailing_slash}$', + mapping={ + 'get': 'retrieve', + 'put': 'update', + 'patch': 'partial_update', + 'delete': 'destroy' + }, + name='{basename}-detail', + detail=True, + initkwargs={'suffix': 'Instance'} + ), + # # Dynamically generated detail routes. Generated using + # # @action(detail=True) decorator on methods of the viewset. + DynamicRoute( + url=r'^{prefix}//{lookup}/{url_path}{trailing_slash}$', + name='{basename}-{url_name}', + detail=True, + initkwargs={} + ), + ] + +router = DefaultRouter() + +router.register('plan', views_application.ApplicationPlanViewSet, basename='plan') +router.register('event', views_application.ApplicationEventViewSet, basename='event') + +urlpatterns = router.urls diff --git a/apps/promotions/urls.py b/apps/promotions/urls_user.py similarity index 79% rename from apps/promotions/urls.py rename to apps/promotions/urls_user.py index 4e1f04b..10ef410 100644 --- a/apps/promotions/urls.py +++ b/apps/promotions/urls_user.py @@ -2,17 +2,17 @@ from rest_framework import routers from rest_framework.routers import DefaultRouter, DynamicRoute, Route from utils.router import ProfileRouter -from . import views +from . import views_user from django.urls import NoReverseMatch, path, re_path, include -from .views import ApplicationPromoteUserApiView, ApplicationEventSubmitAPIView, ApplicationEventRetrieveAPIView, ApplicationPromotionListApiView +from .views_user import ApplicationPromoteUserApiView, ApplicationEventSubmitAPIView, ApplicationPromotionListApiView app_name = 'promotions' router = DefaultRouter() # router.register(r'application', views.ApplicationPromotionViewSet, basename='application_promotions') # router.register(r'events', views.ApplicationViewSet, basename='application-events') -router.register(r'plans', views.UserPlanViewSet, basename='user-plans') +router.register(r'plans', views_user.UserPlanViewSet, basename='user-plans') # urlpatterns = diff --git a/apps/promotions/views.py b/apps/promotions/views_application.py similarity index 74% rename from apps/promotions/views.py rename to apps/promotions/views_application.py index 709fbd2..2d9dead 100644 --- a/apps/promotions/views.py +++ b/apps/promotions/views_application.py @@ -4,9 +4,10 @@ from rest_framework.decorators import action from rest_framework import exceptions from rest_framework.generics import CreateAPIView, get_object_or_404, RetrieveAPIView, ListAPIView from rest_framework.response import Response +from rest_framework.settings import api_settings from rest_framework.viewsets import GenericViewSet -from apps.gooyal_oauth2.rest_framework import IsAuthenticatedOrTokenMatchesOASRequirements +from apps.gooyal_oauth2.rest_framework import IsAuthenticatedOrTokenMatchesOASRequirements, TokenMatchesViewSetActions from apps.gooyal_oauth2.utils import get_application from utils.exceptions import UnprocessableEntity from .models import Plan, Promotion, EventSaver @@ -15,63 +16,71 @@ from .tasks import analyze_event_task from ..users.models import User -# class ApplicationPlanViewSet(mixins.ListModelMixin, -# mixins.CreateModelMixin, -# GenericViewSet): -# queryset = Plan.objects.all() -# serializer_class = PlanSerializer -# -# permission_classes = [IsAuthenticatedOrTokenMatchesOASRequirements] -# required_alternate_scopes = { -# "POST": [["promotions.application.user-plans:submit"]], -# "GET": [["promotions.application.user-promotions:list-retrieve"]], -# } -# -# def get_queryset(self): -# user_uuid = self.kwargs.get('user_uuid') -# user = User.objects.get(uuid=user_uuid) -# application = get_application(self.request) -# return Plan.objects.filter(application=application) -# - -class ApplicationPromotionViewSet( +class ApplicationPlanViewSet( mixins.RetrieveModelMixin, mixins.ListModelMixin, - mixins.CreateModelMixin, + # mixins.CreateModelMixin, GenericViewSet ): serializer_class = PromotionSerializer - permission_classes = [IsAuthenticatedOrTokenMatchesOASRequirements] - required_alternate_scopes = { - "POST": [["promotions.application.user-promotions:promote"]], - "GET": [["promotions.application.user-promotions:list-retrieve"]], + permission_classes = [TokenMatchesViewSetActions] + required_action_scopes = { + "promote": [["promotions.application.user-plan:promote"]], + "retrieve": [["promotions.application.user-plan:list-retrieve"]], + "list": [["promotions.application.user-plan:list-retrieve"]], } def get_queryset(self): - user_uuid = self.kwargs.get('user_uuid') - user = User.objects.get(uuid=user_uuid) application = get_application(self.request) - return Promotion.objects.filter(application=application, user=user) + return Promotion.objects.filter(application=application) - def perform_create(self, serializer): - user_uuid = self.kwargs.get('user_uuid') - user = User.objects.get(uuid=user_uuid) + # def perform_create(self, serializer): + # user = self.get_user() + # application = get_application(self.request) + # + # plan = serializer.validated_data['plan'] + # + # promotion_args = dict(user=user, application=application, **serializer.validated_data) + # promotion_amount = plan.calculate_promotion(**promotion_args) + # if promotion_amount: + # promotion = serializer.save(promotion_amount=promotion_amount, user=user, application=application) + # promotion.promote(**promotion_args) + + @action(detail=True, methods=['POST'], serializer_class=PromoteSerializer) + def promote(self, request, pk=None): + plan = self.get_object() application = get_application(self.request) - plan = serializer.validated_data['plan'] - promotion_args = dict(user=user, application=application, **serializer.validated_data) - promotion_amount = plan.calculate_promotion(**promotion_args) - if promotion_amount: - promotion = serializer.save(promotion_amount=promotion_amount, user=user, application=application) - promotion.promote(**promotion_args) + serializer = self.get_serializer(data=request.data) + serializer.is_valid(raise_exception=True) + data = serializer.validated_data - @action(detail=False, methods=['post'], serializer_class=EventSerializer) - def submit_event(self, request): - pass + event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() + + # TODO + # try: + user_uuid = self.kwargs.get('user_uuid') + event = event_saver.save_event(user=user_uuid, application=application, **serializer.validated_data) + + # except Exception as e: + # raise exceptions.ValidationError(str(e)) + + result_list = plan.process_event(event) + data["promotions"] = result_list + + headers = self.get_success_headers(serializer.data) + return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) + + def get_success_headers(self, data): + try: + return {'Location': str(data[api_settings.URL_FIELD_NAME])} + except (TypeError, KeyError): + return {} -class ApplicationViewSet( + +class ApplicationEventViewSet( mixins.CreateModelMixin, GenericViewSet ): @@ -100,6 +109,7 @@ class ApplicationViewSet( serializer.instance = event +# NOT class UserPlanViewSet(mixins.RetrieveModelMixin, # mixins.ListModelMixin, # mixins.CreateModelMixin, @@ -134,28 +144,6 @@ class UserPlanViewSet(mixins.RetrieveModelMixin, # serializer = self.get_serializer(plans, many=True) # return self.get_paginated_response(serializer.data) - @action(detail=True, methods=['POST'], serializer_class=PromoteSerializer) - def promote(self, request, pk=None): - plan = self.get_object() - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - data = serializer.validated_data - application = get_application(self.request) - - event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() - - user = self.request.user.pk or serializer.validated_data['user'] - # try: - event = event_saver.save_event(user=user, application=application, **serializer.validated_data) - # except Exception as e: - # raise exceptions.ValidationError(str(e)) - - result_list = plan.process_event(event) - data["promotions"] = result_list - - headers = self.get_success_headers(serializer.data) - return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) - class ApplicationPromoteUserApiView(CreateAPIView): model = Promotion @@ -192,6 +180,7 @@ class ApplicationPromoteUserApiView(CreateAPIView): headers = self.get_success_headers(serializer.data) return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) + class ApplicationPromotionListApiView(ListAPIView): model = Promotion serializer_class = PromotionSerializer @@ -212,7 +201,6 @@ class ApplicationPromotionListApiView(ListAPIView): qs = Promotion.objects.filter(user_uuid=user.uuid, plan=plan) return qs - # def create(self, request, *args, **kwargs): # plan = self.get_plan() # serializer = self.get_serializer(data=request.data) @@ -232,7 +220,6 @@ class ApplicationPromotionListApiView(ListAPIView): # return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) - class ApplicationEventSubmitAPIView(CreateAPIView): model = Promotion serializer_class = PromoteSerializer @@ -258,20 +245,20 @@ class ApplicationEventSubmitAPIView(CreateAPIView): headers = self.get_success_headers(serializer.data) return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) -class ApplicationEventRetrieveAPIView(RetrieveAPIView): - model = Promotion - serializer_class = PromoteSerializer - - def create(self, request, *args, **kwargs): - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - data = serializer.validated_data - application = get_application(self.request) - event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() - event = event_saver.save_event(user=None, application=application, **serializer.validated_data) - - analyze_event_task.delay(event.uuid) - - headers = self.get_success_headers(serializer.data) - return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) - +# class ApplicationEventRetrieveAPIView(RetrieveAPIView): +# model = Promotion +# serializer_class = PromoteSerializer +# +# def create(self, request, *args, **kwargs): +# serializer = self.get_serializer(data=request.data) +# serializer.is_valid(raise_exception=True) +# data = serializer.validated_data +# application = get_application(self.request) +# event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() +# event = event_saver.save_event(user=None, application=application, **serializer.validated_data) +# +# analyze_event_task.delay(event.uuid) +# +# headers = self.get_success_headers(serializer.data) +# return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) +# diff --git a/apps/promotions/views_user.py b/apps/promotions/views_user.py new file mode 100644 index 0000000..843377f --- /dev/null +++ b/apps/promotions/views_user.py @@ -0,0 +1,192 @@ +from drf_spectacular.utils import extend_schema +from rest_framework import mixins, status +from rest_framework.decorators import action +from rest_framework import exceptions +from rest_framework.generics import CreateAPIView, get_object_or_404, RetrieveAPIView, ListAPIView +from rest_framework.response import Response +from rest_framework.viewsets import GenericViewSet + +from apps.gooyal_oauth2.rest_framework import IsAuthenticatedOrTokenMatchesOASRequirements +from apps.gooyal_oauth2.utils import get_application +from utils.exceptions import UnprocessableEntity +from .models import Plan, Promotion, EventSaver +from .serializers import PlanSerializer, PromotionSerializer, EventSerializer, PromoteSerializer, UserPlanSerializer +from .tasks import analyze_event_task +from ..users.models import User + + +class UserPlanViewSet(mixins.RetrieveModelMixin, + # mixins.ListModelMixin, + # mixins.CreateModelMixin, + GenericViewSet): + serializer_class = UserPlanSerializer + permission_classes = [IsAuthenticatedOrTokenMatchesOASRequirements] + # required_alternate_scopes = { + # "GET": [["promotions.user.plans:list-retrieve"]], + # "POST": [["promotions.user.plans:list-retrieve"]], + # } + required_alternate_scopes = { + "GET": [[]], + "POST": [[]], + } + + def get_queryset(self): + user = self.request.user + application = get_application(self.request) + + queryset = Plan.objects.all() + return queryset + + # @extend_schema(responses=PlanSerializer(many=True)) + # @action(detail=False, methods=['POST'], serializer_class=PlanSerializer) + # def available(self, request): + # serializer = self.get_serializer(data=self.request.data) + # serializer.is_valid(raise_exception=True) + # print(serializer.validated_data) + # plans = self.paginate_queryset(self.filter_queryset(self.get_queryset())) + # for plan in plans: + # plan.calculate_promotion(**serializer.validated_data) + # serializer = self.get_serializer(plans, many=True) + # return self.get_paginated_response(serializer.data) + + @action(detail=True, methods=['POST'], serializer_class=PromoteSerializer) + def promote(self, request, pk=None): + plan = self.get_object() + serializer = self.get_serializer(data=request.data) + serializer.is_valid(raise_exception=True) + data = serializer.validated_data + application = get_application(self.request) + + event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() + + user = self.request.user.pk or serializer.validated_data['user'] + # try: + event = event_saver.save_event(user=user, application=application, **serializer.validated_data) + # except Exception as e: + # raise exceptions.ValidationError(str(e)) + + result_list = plan.process_event(event) + data["promotions"] = result_list + + headers = self.get_success_headers(serializer.data) + return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) + + +class ApplicationPromoteUserApiView(CreateAPIView): + model = Promotion + serializer_class = PromoteSerializer + + permission_classes = [IsAuthenticatedOrTokenMatchesOASRequirements] + required_alternate_scopes = { + "POST": [[]], + } + + def get_plan(self): + plan_uuid = self.kwargs.get('plan') + plan = get_object_or_404(Plan, uuid=plan_uuid) + return plan + + def create(self, request, *args, **kwargs): + plan = self.get_plan() + serializer = self.get_serializer(data=request.data) + serializer.is_valid(raise_exception=True) + data = serializer.validated_data + application = get_application(self.request) + event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() + user = self.request.user.pk or serializer.validated_data['user'] + # try: + event = event_saver.save_event(user=user, application=application, **serializer.validated_data) + # except Exception as e: + # raise exceptions.ValidationError(str(e)) + + result_list = list(plan.process_event(event)) + # data["promotions"] = list(result_list) + + serializer = self.get_serializer(instance=plan) + + headers = self.get_success_headers(serializer.data) + return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) + +class ApplicationPromotionListApiView(ListAPIView): + model = Promotion + serializer_class = PromotionSerializer + + permission_classes = [IsAuthenticatedOrTokenMatchesOASRequirements] + required_alternate_scopes = { + "GET": [[]], + } + + def get_plan(self): + plan_uuid = self.kwargs.get('plan') + plan = get_object_or_404(Plan, uuid=plan_uuid) + return plan + + def get_queryset(self): + user = self.request.user + plan = self.get_plan() + qs = Promotion.objects.filter(user_uuid=user.uuid, plan=plan) + return qs + + + # def create(self, request, *args, **kwargs): + # plan = self.get_plan() + # serializer = self.get_serializer(data=request.data) + # serializer.is_valid(raise_exception=True) + # data = serializer.validated_data + # application = get_application(self.request) + # event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() + # # try: + # event = event_saver.save_event(user=None, application=application, **serializer.validated_data) + # # except Exception as e: + # # raise exceptions.ValidationError(str(e)) + # + # result_list = plan.process_event(event) + # data["promotions"] = result_list + # + # headers = self.get_success_headers(serializer.data) + # return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) + + + +class ApplicationEventSubmitAPIView(CreateAPIView): + model = Promotion + serializer_class = PromoteSerializer + permission_classes = [IsAuthenticatedOrTokenMatchesOASRequirements] + required_alternate_scopes = { + "POST": [[]], + } + + def create(self, request, *args, **kwargs): + serializer = self.get_serializer(data=request.data) + serializer.is_valid(raise_exception=True) + data = serializer.validated_data + application = get_application(self.request) + event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() + user = self.request.user.pk or serializer.validated_data['user'] + try: + event = event_saver.save_event(user=user, application=application, **serializer.validated_data) + except Exception as e: + raise UnprocessableEntity(str(e)) + + analyze_event_task.delay(event.uuid) + + headers = self.get_success_headers(serializer.data) + return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) + +# class ApplicationEventRetrieveAPIView(RetrieveAPIView): +# model = Promotion +# serializer_class = PromoteSerializer +# +# def create(self, request, *args, **kwargs): +# serializer = self.get_serializer(data=request.data) +# serializer.is_valid(raise_exception=True) +# data = serializer.validated_data +# application = get_application(self.request) +# event_saver: EventSaver = EventSaver.objects.filter(event_label=serializer.validated_data['label']).first() +# event = event_saver.save_event(user=None, application=application, **serializer.validated_data) +# +# analyze_event_task.delay(event.uuid) +# +# headers = self.get_success_headers(serializer.data) +# return Response(serializer.data, status=status.HTTP_201_CREATED, headers=headers) +# diff --git a/main/urls.py b/main/urls.py index 0198a17..1431bb6 100644 --- a/main/urls.py +++ b/main/urls.py @@ -20,7 +20,6 @@ from django.urls import path, include from django.contrib import admin from drf_spectacular.views import SpectacularAPIView, SpectacularRedocView, SpectacularSwaggerView - urlpatterns = [ path('swagger/', SpectacularAPIView.as_view(), name='schema'), path('swagger/swagger-ui/', SpectacularSwaggerView.as_view(url_name='schema'), name='swagger-ui'), @@ -29,7 +28,8 @@ urlpatterns = [ path('admin/', admin.site.urls), path('', include('utils.urls')), path('oauth2/', include('oauth2_provider.urls', namespace='oauth2_provider')), - path('promotions/', include('apps.promotions.urls', namespace='promotions')), + path('promotions/', include('apps.promotions.urls_user', namespace='promotions')), + path('api/v2/promotions/application//', include('apps.promotions.urls_application', namespace='promotions-application')), ] urlpatterns += static(settings.MEDIA_URL, document_root=settings.MEDIA_ROOT)