Revert "Merge DRF 3.1 in to master"
This commit is contained in:
@@ -124,4 +124,4 @@ def course_grading_policy(course_key):
|
||||
final grade.
|
||||
"""
|
||||
course = _retrieve_course(course_key)
|
||||
return GradingPolicySerializer(course.raw_grader, many=True).data
|
||||
return GradingPolicySerializer(course.raw_grader).data
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
"""
|
||||
API Serializers
|
||||
"""
|
||||
from collections import defaultdict
|
||||
|
||||
from rest_framework import serializers
|
||||
|
||||
|
||||
@@ -13,58 +11,23 @@ class GradingPolicySerializer(serializers.Serializer):
|
||||
dropped = serializers.IntegerField(source='drop_count')
|
||||
weight = serializers.FloatField()
|
||||
|
||||
def to_representation(self, obj):
|
||||
"""
|
||||
Return a representation of the grading policy.
|
||||
"""
|
||||
# Backwards compatibility with the behavior of DRF v2.
|
||||
# When the grader dictionary was missing keys, DRF v2 would default to None;
|
||||
# DRF v3 unhelpfully raises an exception.
|
||||
return dict(
|
||||
super(GradingPolicySerializer, self).to_representation(
|
||||
defaultdict(lambda: None, obj)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# pylint: disable=invalid-name
|
||||
class BlockSerializer(serializers.Serializer):
|
||||
""" Serializer for course structure block. """
|
||||
id = serializers.CharField(source='usage_key')
|
||||
type = serializers.CharField(source='block_type')
|
||||
parent = serializers.CharField(required=False)
|
||||
parent = serializers.CharField(source='parent')
|
||||
display_name = serializers.CharField()
|
||||
graded = serializers.BooleanField(default=False)
|
||||
format = serializers.CharField()
|
||||
children = serializers.CharField()
|
||||
|
||||
def to_representation(self, obj):
|
||||
"""
|
||||
Return a representation of the block.
|
||||
|
||||
NOTE: this method maintains backwards compatibility with the behavior
|
||||
of Django Rest Framework v2.
|
||||
"""
|
||||
data = super(BlockSerializer, self).to_representation(obj)
|
||||
|
||||
# Backwards compatibility with the behavior of DRF v2
|
||||
# Include a NULL value for "parent" in the representation
|
||||
# (instead of excluding the key entirely)
|
||||
if obj.get("parent") is None:
|
||||
data["parent"] = None
|
||||
|
||||
# Backwards compatibility with the behavior of DRF v2
|
||||
# Leave the children list as a list instead of serializing
|
||||
# it to a string.
|
||||
data["children"] = obj["children"]
|
||||
|
||||
return data
|
||||
|
||||
|
||||
class CourseStructureSerializer(serializers.Serializer):
|
||||
""" Serializer for course structure. """
|
||||
root = serializers.CharField()
|
||||
blocks = serializers.SerializerMethodField()
|
||||
root = serializers.CharField(source='root')
|
||||
blocks = serializers.SerializerMethodField('get_blocks')
|
||||
|
||||
def get_blocks(self, structure):
|
||||
""" Serialize the individual blocks. """
|
||||
|
||||
@@ -2,33 +2,12 @@
|
||||
|
||||
from rest_framework import serializers
|
||||
|
||||
from opaque_keys.edx.keys import CourseKey
|
||||
from opaque_keys import InvalidKeyError
|
||||
from openedx.core.djangoapps.credit.models import CreditCourse
|
||||
|
||||
|
||||
class CourseKeyField(serializers.Field):
|
||||
"""
|
||||
Serializer field for a model CourseKey field.
|
||||
"""
|
||||
|
||||
def to_representation(self, data):
|
||||
"""Convert a course key to unicode. """
|
||||
return unicode(data)
|
||||
|
||||
def to_internal_value(self, data):
|
||||
"""Convert unicode to a course key. """
|
||||
try:
|
||||
return CourseKey.from_string(data)
|
||||
except InvalidKeyError as ex:
|
||||
raise serializers.ValidationError("Invalid course key: {msg}".format(msg=ex.msg))
|
||||
|
||||
|
||||
class CreditCourseSerializer(serializers.ModelSerializer):
|
||||
""" CreditCourse Serializer """
|
||||
|
||||
course_key = CourseKeyField()
|
||||
|
||||
class Meta(object): # pylint: disable=missing-docstring
|
||||
model = CreditCourse
|
||||
exclude = ('id',)
|
||||
|
||||
@@ -393,7 +393,10 @@ class CreditCourseViewSetTests(TestCase):
|
||||
|
||||
# POSTs without a CSRF token should fail.
|
||||
response = client.post(self.path, data=json.dumps(data), content_type=JSON)
|
||||
self.assertEqual(response.status_code, 403)
|
||||
|
||||
# NOTE (CCB): Ordinarily we would expect a 403; however, since the CSRF validation and session authentication
|
||||
# fail, DRF considers the request to be unauthenticated.
|
||||
self.assertEqual(response.status_code, 401)
|
||||
self.assertIn('CSRF', response.content)
|
||||
|
||||
# Retrieve a CSRF token
|
||||
|
||||
@@ -18,9 +18,8 @@ from django.views.decorators.http import require_POST, require_GET
|
||||
from opaque_keys import InvalidKeyError
|
||||
from opaque_keys.edx.keys import CourseKey
|
||||
import pytz
|
||||
from rest_framework import viewsets, mixins, permissions
|
||||
from rest_framework.authentication import SessionAuthentication
|
||||
from rest_framework_oauth.authentication import OAuth2Authentication
|
||||
from rest_framework import viewsets, mixins, permissions, authentication
|
||||
|
||||
from util.json_request import JsonResponse
|
||||
from util.date_utils import from_timestamp
|
||||
from openedx.core.djangoapps.credit import api
|
||||
@@ -378,28 +377,17 @@ class CreditCourseViewSet(mixins.CreateModelMixin, mixins.UpdateModelMixin, view
|
||||
lookup_value_regex = settings.COURSE_KEY_REGEX
|
||||
queryset = CreditCourse.objects.all()
|
||||
serializer_class = CreditCourseSerializer
|
||||
authentication_classes = (OAuth2Authentication, SessionAuthentication,)
|
||||
authentication_classes = (authentication.OAuth2Authentication, authentication.SessionAuthentication,)
|
||||
permission_classes = (permissions.IsAuthenticated, permissions.IsAdminUser)
|
||||
|
||||
# In Django Rest Framework v3, there is a default pagination
|
||||
# class that transmutes the response data into a dictionary
|
||||
# with pagination information. The original response data (a list)
|
||||
# is stored in a "results" value of the dictionary.
|
||||
# For backwards compatibility with the existing API, we disable
|
||||
# the default behavior by setting the pagination_class to None.
|
||||
pagination_class = None
|
||||
|
||||
# This CSRF exemption only applies when authenticating without SessionAuthentication.
|
||||
# SessionAuthentication will enforce CSRF protection.
|
||||
@method_decorator(csrf_exempt)
|
||||
def dispatch(self, request, *args, **kwargs):
|
||||
# Convert the course ID/key from a string to an actual CourseKey object.
|
||||
course_id = kwargs.get(self.lookup_field, None)
|
||||
|
||||
if course_id:
|
||||
kwargs[self.lookup_field] = CourseKey.from_string(course_id)
|
||||
|
||||
return super(CreditCourseViewSet, self).dispatch(request, *args, **kwargs)
|
||||
|
||||
def get_object(self):
|
||||
# Convert the serialized course key into a CourseKey instance
|
||||
# so we can look up the object.
|
||||
course_key = self.kwargs.get(self.lookup_field)
|
||||
if course_key is not None:
|
||||
self.kwargs[self.lookup_field] = CourseKey.from_string(course_key)
|
||||
|
||||
return super(CreditCourseViewSet, self).get_object()
|
||||
|
||||
@@ -30,40 +30,6 @@ TEST_UPLOAD_DT = datetime.datetime(2002, 1, 9, 15, 43, 01, tzinfo=UTC)
|
||||
TEST_UPLOAD_DT2 = datetime.datetime(2003, 1, 9, 15, 43, 01, tzinfo=UTC)
|
||||
|
||||
|
||||
class PatchedClient(APIClient):
|
||||
"""
|
||||
Patch DRF's APIClient to avoid a unicode error on file upload.
|
||||
|
||||
Famous last words: This is a *temporary* fix that we should be
|
||||
able to remove once we upgrade Django past 1.4.
|
||||
"""
|
||||
|
||||
def request(self, *args, **kwargs):
|
||||
"""Construct an API request. """
|
||||
# DRF's default test client implementation uses `six.text_type()`
|
||||
# to convert the CONTENT_TYPE to `unicode`. In Django 1.4,
|
||||
# this causes a `UnicodeDecodeError` when Django parses a multipart
|
||||
# upload.
|
||||
#
|
||||
# This is the DRF code we're working around:
|
||||
# https://github.com/tomchristie/django-rest-framework/blob/3.1.3/rest_framework/compat.py#L227
|
||||
#
|
||||
# ... and this is the Django code that raises the exception:
|
||||
#
|
||||
# https://github.com/django/django/blob/1.4.22/django/http/multipartparser.py#L435
|
||||
#
|
||||
# Django unhelpfully swallows the exception, so to the application code
|
||||
# it appears as though the user didn't send any file data.
|
||||
#
|
||||
# This appears to be an issue only with requests constructed in the test
|
||||
# suite, not with the upload code used in production.
|
||||
#
|
||||
if isinstance(kwargs.get("CONTENT_TYPE"), basestring):
|
||||
kwargs["CONTENT_TYPE"] = str(kwargs["CONTENT_TYPE"])
|
||||
|
||||
return super(PatchedClient, self).request(*args, **kwargs)
|
||||
|
||||
|
||||
class ProfileImageEndpointTestCase(UserSettingsEventTestMixin, APITestCase):
|
||||
"""
|
||||
Base class / shared infrastructure for tests of profile_image "upload" and
|
||||
@@ -145,10 +111,6 @@ class ProfileImageUploadTestCase(ProfileImageEndpointTestCase):
|
||||
"""
|
||||
_view_name = "profile_image_upload"
|
||||
|
||||
# Use the patched version of the API client to workaround a unicode issue
|
||||
# with DRF 3.1 and Django 1.4. Remove this after we upgrade Django past 1.4!
|
||||
client_class = PatchedClient
|
||||
|
||||
def check_upload_event_emitted(self, old=None, new=TEST_UPLOAD_DT):
|
||||
"""
|
||||
Make sure we emit a UserProfile event corresponding to the
|
||||
|
||||
@@ -183,7 +183,7 @@ def update_account_settings(requesting_user, update, username=None):
|
||||
serializer.save()
|
||||
|
||||
if "language_proficiencies" in update:
|
||||
new_language_proficiencies = update["language_proficiencies"]
|
||||
new_language_proficiencies = legacy_profile_serializer.data["language_proficiencies"]
|
||||
emit_setting_changed_event(
|
||||
user=existing_user,
|
||||
db_table=existing_user_profile.language_proficiencies.model._meta.db_table,
|
||||
|
||||
@@ -53,7 +53,7 @@ class UserReadOnlySerializer(serializers.Serializer):
|
||||
|
||||
super(UserReadOnlySerializer, self).__init__(*args, **kwargs)
|
||||
|
||||
def to_representation(self, user):
|
||||
def to_native(self, user):
|
||||
"""
|
||||
Overwrite to_native to handle custom logic since we are serializing two models as one here
|
||||
:param user: User object
|
||||
@@ -152,8 +152,8 @@ class AccountLegacyProfileSerializer(serializers.HyperlinkedModelSerializer, Rea
|
||||
Class that serializes the portion of UserProfile model needed for account information.
|
||||
"""
|
||||
profile_image = serializers.SerializerMethodField("_get_profile_image")
|
||||
requires_parental_consent = serializers.SerializerMethodField()
|
||||
language_proficiencies = LanguageProficiencySerializer(many=True, required=False)
|
||||
requires_parental_consent = serializers.SerializerMethodField("get_requires_parental_consent")
|
||||
language_proficiencies = LanguageProficiencySerializer(many=True, allow_add_remove=True, required=False)
|
||||
|
||||
class Meta(object): # pylint: disable=missing-docstring
|
||||
model = UserProfile
|
||||
@@ -165,21 +165,25 @@ class AccountLegacyProfileSerializer(serializers.HyperlinkedModelSerializer, Rea
|
||||
read_only_fields = ()
|
||||
explicit_read_only_fields = ("profile_image", "requires_parental_consent")
|
||||
|
||||
def validate_name(self, new_name):
|
||||
def validate_name(self, attrs, source):
|
||||
""" Enforce minimum length for name. """
|
||||
if len(new_name) < NAME_MIN_LENGTH:
|
||||
raise serializers.ValidationError(
|
||||
"The name field must be at least {} characters long.".format(NAME_MIN_LENGTH)
|
||||
)
|
||||
return new_name
|
||||
if source in attrs:
|
||||
new_name = attrs[source].strip()
|
||||
if len(new_name) < NAME_MIN_LENGTH:
|
||||
raise serializers.ValidationError(
|
||||
"The name field must be at least {} characters long.".format(NAME_MIN_LENGTH)
|
||||
)
|
||||
attrs[source] = new_name
|
||||
|
||||
def validate_language_proficiencies(self, value):
|
||||
return attrs
|
||||
|
||||
def validate_language_proficiencies(self, attrs, source):
|
||||
""" Enforce all languages are unique. """
|
||||
language_proficiencies = [language for language in value]
|
||||
unique_language_proficiencies = set(language["code"] for language in language_proficiencies)
|
||||
language_proficiencies = [language for language in attrs.get(source, [])]
|
||||
unique_language_proficiencies = set(language.code for language in language_proficiencies)
|
||||
if len(language_proficiencies) != len(unique_language_proficiencies):
|
||||
raise serializers.ValidationError("The language_proficiencies field must consist of unique languages")
|
||||
return value
|
||||
return attrs
|
||||
|
||||
def transform_gender(self, user_profile, value):
|
||||
""" Converts empty string to None, to indicate not set. Replaced by to_representation in version 3. """
|
||||
@@ -226,29 +230,3 @@ class AccountLegacyProfileSerializer(serializers.HyperlinkedModelSerializer, Rea
|
||||
call the method with a single argument, the user_profile object.
|
||||
"""
|
||||
return AccountLegacyProfileSerializer.get_profile_image(user_profile, user_profile.user)
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
"""
|
||||
Update the profile, including nested fields.
|
||||
"""
|
||||
language_proficiencies = validated_data.pop("language_proficiencies", None)
|
||||
|
||||
# Update all fields on the user profile that are writeable,
|
||||
# except for "language_proficiencies", which we'll update separately
|
||||
update_fields = set(self.get_writeable_fields()) - set(["language_proficiencies"])
|
||||
for field_name in update_fields:
|
||||
default = getattr(instance, field_name)
|
||||
field_value = validated_data.get(field_name, default)
|
||||
setattr(instance, field_name, field_value)
|
||||
|
||||
instance.save()
|
||||
|
||||
# Now update the related language proficiency
|
||||
if language_proficiencies is not None:
|
||||
instance.language_proficiencies.all().delete()
|
||||
instance.language_proficiencies.bulk_create([
|
||||
LanguageProficiency(user_profile=instance, code=language["code"])
|
||||
for language in language_proficiencies
|
||||
])
|
||||
|
||||
return instance
|
||||
|
||||
@@ -164,10 +164,7 @@ class TestAccountApi(UserSettingsEventTestMixin, TestCase):
|
||||
field_errors = context_manager.exception.field_errors
|
||||
self.assertEqual(3, len(field_errors))
|
||||
self.assertEqual("This field is not editable via this API", field_errors["username"]["developer_message"])
|
||||
self.assertIn(
|
||||
"Value \'undecided\' is not valid for field \'gender\'",
|
||||
field_errors["gender"]["developer_message"]
|
||||
)
|
||||
self.assertIn("Select a valid choice", field_errors["gender"]["developer_message"])
|
||||
self.assertIn("Valid e-mail address required.", field_errors["email"]["developer_message"])
|
||||
|
||||
@patch('django.core.mail.send_mail')
|
||||
|
||||
@@ -359,19 +359,16 @@ class TestAccountAPI(UserAPITestCase):
|
||||
self.assertEqual(404, response.status_code)
|
||||
|
||||
@ddt.data(
|
||||
("gender", "f", "not a gender", u'"not a gender" is not a valid choice.'),
|
||||
("level_of_education", "none", u"ȻħȺɍłɇs", u'"ȻħȺɍłɇs" is not a valid choice.'),
|
||||
("country", "GB", "XY", u'"XY" is not a valid choice.'),
|
||||
("year_of_birth", 2009, "not_an_int", u"A valid integer is required."),
|
||||
("name", "bob", "z" * 256, u"Ensure this field has no more than 255 characters."),
|
||||
("gender", "f", "not a gender", u"Select a valid choice. not a gender is not one of the available choices."),
|
||||
("level_of_education", "none", u"ȻħȺɍłɇs", u"Select a valid choice. ȻħȺɍłɇs is not one of the available choices."),
|
||||
("country", "GB", "XY", u"Select a valid choice. XY is not one of the available choices."),
|
||||
("year_of_birth", 2009, "not_an_int", u"Enter a whole number."),
|
||||
("name", "bob", "z" * 256, u"Ensure this value has at most 255 characters (it has 256)."),
|
||||
("name", u"ȻħȺɍłɇs", "z ", u"The name field must be at least 2 characters long."),
|
||||
("goals", "Smell the roses"),
|
||||
("mailing_address", "Sesame Street"),
|
||||
# Note that we store the raw data, so it is up to client to escape the HTML.
|
||||
(
|
||||
"bio", u"<html>Lacrosse-playing superhero 壓是進界推日不復女</html>",
|
||||
"z" * 3001, u"Ensure this field has no more than 3000 characters."
|
||||
),
|
||||
("bio", u"<html>Lacrosse-playing superhero 壓是進界推日不復女</html>", "z" * 3001, u"Ensure this value has at most 3000 characters (it has 3001)."),
|
||||
# Note that email is tested below, as it is not immediately updated.
|
||||
# Note that language_proficiencies is tested below as there are multiple error and success conditions.
|
||||
)
|
||||
@@ -571,10 +568,10 @@ class TestAccountAPI(UserAPITestCase):
|
||||
self.assertItemsEqual(response.data["language_proficiencies"], proficiencies)
|
||||
|
||||
@ddt.data(
|
||||
(u"not_a_list", {u'non_field_errors': [u'Expected a list of items but got type "unicode".']}),
|
||||
([u"not_a_JSON_object"], [{u'non_field_errors': [u'Invalid data. Expected a dictionary, but got unicode.']}]),
|
||||
(u"not_a_list", [{u'non_field_errors': [u'Expected a list of items.']}]),
|
||||
([u"not_a_JSON_object"], [{u'non_field_errors': [u'Invalid data']}]),
|
||||
([{}], [{"code": [u"This field is required."]}]),
|
||||
([{u"code": u"invalid_language_code"}], [{'code': [u'"invalid_language_code" is not a valid choice.']}]),
|
||||
([{u"code": u"invalid_language_code"}], [{'code': [u'Select a valid choice. invalid_language_code is not one of the available choices.']}]),
|
||||
([{u"code": u"kw"}, {u"code": u"el"}, {u"code": u"kw"}], [u'The language_proficiencies field must consist of unique languages']),
|
||||
)
|
||||
@ddt.unpack
|
||||
|
||||
@@ -9,10 +9,9 @@ from django.conf import settings
|
||||
from django.core.exceptions import ObjectDoesNotExist
|
||||
from django.db import IntegrityError
|
||||
from django.utils.translation import ugettext as _
|
||||
from student.models import User, UserProfile
|
||||
from django.utils.translation import ugettext_noop
|
||||
|
||||
from student.models import User, UserProfile
|
||||
from request_cache import get_request_or_stub
|
||||
from ..errors import (
|
||||
UserAPIInternalError, UserAPIRequestError, UserNotFound, UserNotAuthorized,
|
||||
PreferenceValidationError, PreferenceUpdateError
|
||||
@@ -69,17 +68,7 @@ def get_user_preferences(requesting_user, username=None):
|
||||
UserAPIInternalError: the operation failed due to an unexpected error.
|
||||
"""
|
||||
existing_user = _get_user(requesting_user, username, allow_staff=True)
|
||||
|
||||
# Django Rest Framework V3 uses the current request to version
|
||||
# hyperlinked URLS, so we need to retrieve the request and pass
|
||||
# it in the serializer's context (otherwise we get an AssertionError).
|
||||
# We're retrieving the request from the cache rather than passing it in
|
||||
# as an argument because this is an implementation detail of how we're
|
||||
# serializing data, which we want to encapsulate in the API call.
|
||||
context = {
|
||||
"request": get_request_or_stub()
|
||||
}
|
||||
user_serializer = UserSerializer(existing_user, context=context)
|
||||
user_serializer = UserSerializer(existing_user)
|
||||
return user_serializer.data["preferences"]
|
||||
|
||||
|
||||
@@ -367,7 +356,7 @@ def validate_user_preference_serializer(serializer, preference_key, preference_v
|
||||
developer_message = u"Value '{preference_value}' not valid for preference '{preference_key}': {error}".format(
|
||||
preference_key=preference_key, preference_value=preference_value, error=serializer.errors
|
||||
)
|
||||
if "key" in serializer.errors:
|
||||
if serializer.errors["key"]:
|
||||
user_message = _(u"Invalid user preference key '{preference_key}'.").format(
|
||||
preference_key=preference_key
|
||||
)
|
||||
|
||||
@@ -403,7 +403,7 @@ def get_expected_validation_developer_message(preference_key, preference_value):
|
||||
preference_key=preference_key,
|
||||
preference_value=preference_value,
|
||||
error={
|
||||
"key": [u"Ensure this field has no more than 255 characters."]
|
||||
"key": [u"Ensure this value has at most 255 characters (it has 256)."]
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ from .models import UserPreference
|
||||
|
||||
|
||||
class UserSerializer(serializers.HyperlinkedModelSerializer):
|
||||
name = serializers.SerializerMethodField()
|
||||
preferences = serializers.SerializerMethodField()
|
||||
name = serializers.SerializerMethodField("get_name")
|
||||
preferences = serializers.SerializerMethodField("get_preferences")
|
||||
|
||||
def get_name(self, user):
|
||||
profile = UserProfile.objects.get(user=user)
|
||||
@@ -32,10 +32,9 @@ class UserPreferenceSerializer(serializers.HyperlinkedModelSerializer):
|
||||
|
||||
|
||||
class RawUserPreferenceSerializer(serializers.ModelSerializer):
|
||||
"""Serializer that generates a raw representation of a user preference.
|
||||
"""
|
||||
Serializer that generates a raw representation of a user preference.
|
||||
"""
|
||||
user = serializers.PrimaryKeyRelatedField(queryset=User.objects.all())
|
||||
user = serializers.PrimaryKeyRelatedField()
|
||||
|
||||
class Meta(object): # pylint: disable=missing-docstring
|
||||
model = UserPreference
|
||||
@@ -58,11 +57,3 @@ class ReadOnlyFieldsSerializerMixin(object):
|
||||
cls.Meta.read_only_fields tuple.
|
||||
"""
|
||||
return getattr(cls.Meta, 'read_only_fields', '') + getattr(cls.Meta, 'explicit_read_only_fields', '')
|
||||
|
||||
@classmethod
|
||||
def get_writeable_fields(cls):
|
||||
"""
|
||||
Return all fields on this serializer that are writeable.
|
||||
"""
|
||||
all_fields = getattr(cls.Meta, 'fields', tuple())
|
||||
return tuple(set(all_fields) - set(cls.get_read_only_fields()))
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
""" Common Authentication Handlers used across projects. """
|
||||
from rest_framework.authentication import SessionAuthentication
|
||||
from rest_framework_oauth.authentication import OAuth2Authentication
|
||||
from rest_framework import authentication
|
||||
from rest_framework.exceptions import AuthenticationFailed
|
||||
from rest_framework_oauth.compat import oauth2_provider, provider_now
|
||||
from rest_framework.compat import oauth2_provider, provider_now
|
||||
|
||||
|
||||
class SessionAuthenticationAllowInactiveUser(SessionAuthentication):
|
||||
class SessionAuthenticationAllowInactiveUser(authentication.SessionAuthentication):
|
||||
"""Ensure that the user is logged in, but do not require the account to be active.
|
||||
|
||||
We use this in the special case that a user has created an account,
|
||||
@@ -52,7 +51,7 @@ class SessionAuthenticationAllowInactiveUser(SessionAuthentication):
|
||||
return (user, None)
|
||||
|
||||
|
||||
class OAuth2AuthenticationAllowInactiveUser(OAuth2Authentication):
|
||||
class OAuth2AuthenticationAllowInactiveUser(authentication.OAuth2Authentication):
|
||||
"""
|
||||
This is a temporary workaround while the is_active field on the user is coupled
|
||||
with whether or not the user has verified ownership of their claimed email address.
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Fields useful for edX API implementations."""
|
||||
from rest_framework.serializers import Field
|
||||
from django.core.exceptions import ValidationError
|
||||
|
||||
from rest_framework.serializers import CharField, Field
|
||||
|
||||
|
||||
class ExpandableField(Field):
|
||||
@@ -16,19 +18,25 @@ class ExpandableField(Field):
|
||||
self.expanded = kwargs.pop('expanded_serializer')
|
||||
super(ExpandableField, self).__init__(**kwargs)
|
||||
|
||||
def to_representation(self, obj):
|
||||
"""
|
||||
Return a representation of the field that is either expanded or collapsed.
|
||||
"""
|
||||
should_expand = self.field_name in self.context.get("expand", [])
|
||||
field = self.expanded if should_expand else self.collapsed
|
||||
def field_to_native(self, obj, field_name):
|
||||
"""Converts obj to a native representation, using the expanded serializer if the context requires it."""
|
||||
if 'expand' in self.context and field_name in self.context['expand']:
|
||||
self.expanded.initialize(self, field_name)
|
||||
return self.expanded.field_to_native(obj, field_name)
|
||||
else:
|
||||
self.collapsed.initialize(self, field_name)
|
||||
return self.collapsed.field_to_native(obj, field_name)
|
||||
|
||||
# Avoid double-binding the field, otherwise we'll get
|
||||
# an error about the source kwarg being redundant.
|
||||
if field.source is None:
|
||||
field.bind(self.field_name, self)
|
||||
|
||||
if should_expand:
|
||||
self.expanded.context["expand"] = set(field.context.get("expand", []))
|
||||
class NonEmptyCharField(CharField):
|
||||
"""
|
||||
A field that enforces non-emptiness even for partial updates.
|
||||
|
||||
return field.to_representation(obj)
|
||||
This is necessary because prior to version 3, DRF skips validation for empty
|
||||
values. Thus, CharField's min_length and RegexField cannot be used to
|
||||
enforce this constraint.
|
||||
"""
|
||||
def validate(self, value):
|
||||
super(NonEmptyCharField, self).validate(value)
|
||||
if not value.strip():
|
||||
raise ValidationError(self.error_messages["required"])
|
||||
|
||||
@@ -1,33 +0,0 @@
|
||||
"""
|
||||
Django Rest Framework view mixins.
|
||||
"""
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.http import Http404
|
||||
from rest_framework import status
|
||||
from rest_framework.mixins import CreateModelMixin
|
||||
from rest_framework.response import Response
|
||||
|
||||
|
||||
class PutAsCreateMixin(CreateModelMixin):
|
||||
"""
|
||||
Backwards compatibility with Django Rest Framework v2, which allowed
|
||||
creation of a new resource using PUT.
|
||||
"""
|
||||
|
||||
def update(self, request, *args, **kwargs):
|
||||
"""
|
||||
Create/update course modes for a course.
|
||||
"""
|
||||
# First, try to update the existing instance
|
||||
try:
|
||||
try:
|
||||
return super(PutAsCreateMixin, self).update(request, *args, **kwargs)
|
||||
except Http404:
|
||||
# If no instance exists yet, create it.
|
||||
# This is backwards-compatible with the behavior of DRF v2.
|
||||
return super(PutAsCreateMixin, self).create(request, *args, **kwargs)
|
||||
|
||||
# Backwards compatibility with DRF v2 behavior, which would catch model-level
|
||||
# validation errors and return a 400
|
||||
except ValidationError as err:
|
||||
return Response(err.messages, status=status.HTTP_400_BAD_REQUEST)
|
||||
@@ -3,31 +3,6 @@
|
||||
from django.http import Http404
|
||||
from django.core.paginator import Paginator, InvalidPage
|
||||
|
||||
from rest_framework.response import Response
|
||||
from rest_framework import pagination
|
||||
|
||||
|
||||
class DefaultPagination(pagination.PageNumberPagination):
|
||||
"""
|
||||
Default paginator for APIs in edx-platform.
|
||||
|
||||
This is configured in settings to be automatically used
|
||||
by any subclass of Django Rest Framework's generic API views.
|
||||
"""
|
||||
page_size_query_param = "page_size"
|
||||
|
||||
def get_paginated_response(self, data):
|
||||
"""
|
||||
Annotate the response with pagination information.
|
||||
"""
|
||||
return Response({
|
||||
'next': self.get_next_link(),
|
||||
'previous': self.get_previous_link(),
|
||||
'count': self.page.paginator.count,
|
||||
'num_pages': self.page.paginator.num_pages,
|
||||
'results': data
|
||||
})
|
||||
|
||||
|
||||
def paginate_search_results(object_class, search_results, page_size, page):
|
||||
"""
|
||||
|
||||
@@ -1,8 +1,32 @@
|
||||
"""
|
||||
Serializers to be used in APIs.
|
||||
"""
|
||||
from rest_framework import pagination, serializers
|
||||
|
||||
from rest_framework import serializers
|
||||
|
||||
class PaginationSerializer(pagination.PaginationSerializer):
|
||||
"""
|
||||
Custom PaginationSerializer for openedx.
|
||||
|
||||
Adds the following fields:
|
||||
- num_pages: total number of pages
|
||||
- current_page: the current page being returned
|
||||
- start: the index of the first page item within the overall collection
|
||||
"""
|
||||
start_page = 1 # django Paginator.page objects have 1-based indexes
|
||||
num_pages = serializers.Field(source='paginator.num_pages')
|
||||
current_page = serializers.SerializerMethodField('get_current_page')
|
||||
start = serializers.SerializerMethodField('get_start')
|
||||
sort_order = serializers.SerializerMethodField('get_sort_order')
|
||||
|
||||
def get_current_page(self, page):
|
||||
"""Get the current page"""
|
||||
return page.number
|
||||
|
||||
def get_start(self, page):
|
||||
"""Get the index of the first page item within the overall collection"""
|
||||
return (self.get_current_page(page) - self.start_page) * page.paginator.per_page
|
||||
|
||||
def get_sort_order(self, page): # pylint: disable=unused-argument
|
||||
"""Get the order by which this collection was sorted"""
|
||||
return self.context.get('sort_order')
|
||||
|
||||
|
||||
class CollapsedReferenceSerializer(serializers.HyperlinkedModelSerializer):
|
||||
@@ -30,10 +54,9 @@ class CollapsedReferenceSerializer(serializers.HyperlinkedModelSerializer):
|
||||
|
||||
super(CollapsedReferenceSerializer, self).__init__(*args, **kwargs)
|
||||
|
||||
self.fields[id_source] = serializers.CharField(read_only=True)
|
||||
self.fields[id_source] = serializers.CharField(read_only=True, source=id_source)
|
||||
self.fields['url'].view_name = view_name
|
||||
self.fields['url'].lookup_field = lookup_field
|
||||
self.fields['url'].lookup_url_kwarg = lookup_field
|
||||
|
||||
class Meta(object):
|
||||
"""Defines meta information for the ModelSerializer.
|
||||
|
||||
@@ -1,235 +1,63 @@
|
||||
"""
|
||||
Tests for OAuth2. This module is copied from django-rest-framework-oauth (tests/test_authentication.py)
|
||||
and updated to use our subclass of OAuth2Authentication.
|
||||
"""
|
||||
|
||||
from __future__ import unicode_literals
|
||||
import datetime
|
||||
|
||||
from django.conf.urls import patterns, url, include
|
||||
from django.contrib.auth.models import User
|
||||
from django.http import HttpResponse
|
||||
from django.test import TestCase
|
||||
from django.utils import unittest
|
||||
from django.utils.http import urlencode
|
||||
|
||||
from rest_framework import status
|
||||
from rest_framework.permissions import IsAuthenticated
|
||||
from rest_framework_oauth import permissions
|
||||
from rest_framework_oauth.compat import oauth2_provider, oauth2_provider_scope
|
||||
from rest_framework.test import APIRequestFactory, APIClient
|
||||
from rest_framework.views import APIView
|
||||
"""Tests for util.authentication module."""
|
||||
|
||||
from mock import patch
|
||||
from django.conf import settings
|
||||
from rest_framework import permissions
|
||||
from rest_framework.compat import patterns, url
|
||||
from rest_framework.tests import test_authentication
|
||||
from provider import scope, constants
|
||||
from unittest import skipUnless
|
||||
|
||||
from ..authentication import OAuth2AuthenticationAllowInactiveUser
|
||||
|
||||
factory = APIRequestFactory() # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class MockView(APIView): # pylint: disable=missing-docstring
|
||||
permission_classes = (IsAuthenticated,)
|
||||
|
||||
def get(self, request): # pylint: disable=missing-docstring,unused-argument
|
||||
return HttpResponse({'a': 1, 'b': 2, 'c': 3})
|
||||
|
||||
def post(self, request): # pylint: disable=missing-docstring,unused-argument
|
||||
return HttpResponse({'a': 1, 'b': 2, 'c': 3})
|
||||
|
||||
def put(self, request): # pylint: disable=missing-docstring,unused-argument
|
||||
return HttpResponse({'a': 1, 'b': 2, 'c': 3})
|
||||
|
||||
|
||||
# This is the a change we've made from the django-rest-framework-oauth version
|
||||
# of these tests. We're subclassing our custom OAuth2AuthenticationAllowInactiveUser
|
||||
# instead of OAuth2Authentication.
|
||||
class OAuth2AuthenticationDebug(OAuth2AuthenticationAllowInactiveUser): # pylint: disable=missing-docstring
|
||||
class OAuth2AuthAllowInactiveUserDebug(OAuth2AuthenticationAllowInactiveUser):
|
||||
"""
|
||||
A debug class analogous to the OAuth2AuthenticationDebug class that tests
|
||||
the OAuth2 flow with the access token sent in a query param."""
|
||||
allow_query_params_token = True
|
||||
|
||||
|
||||
urlpatterns = patterns(
|
||||
'',
|
||||
url(r'^oauth2/', include('provider.oauth2.urls', namespace='oauth2')),
|
||||
url(r'^oauth2-test/$', MockView.as_view(authentication_classes=[OAuth2AuthenticationAllowInactiveUser])),
|
||||
url(r'^oauth2-test-debug/$', MockView.as_view(authentication_classes=[OAuth2AuthenticationDebug])),
|
||||
url(
|
||||
r'^oauth2-with-scope-test/$',
|
||||
MockView.as_view(
|
||||
authentication_classes=[OAuth2AuthenticationAllowInactiveUser],
|
||||
permission_classes=[permissions.TokenHasReadWriteScope]
|
||||
# The following patch overrides the URL patterns for the MockView class used in
|
||||
# rest_framework.tests.test_authentication so that the corresponding AllowInactiveUser
|
||||
# classes are tested instead.
|
||||
@skipUnless(settings.FEATURES.get('ENABLE_OAUTH2_PROVIDER'), 'OAuth2 not enabled')
|
||||
@patch.object(
|
||||
test_authentication,
|
||||
'urlpatterns',
|
||||
patterns(
|
||||
'',
|
||||
url(
|
||||
r'^oauth2-test/$',
|
||||
test_authentication.MockView.as_view(authentication_classes=[OAuth2AuthenticationAllowInactiveUser])
|
||||
),
|
||||
url(
|
||||
r'^oauth2-test-debug/$',
|
||||
test_authentication.MockView.as_view(authentication_classes=[OAuth2AuthAllowInactiveUserDebug])
|
||||
),
|
||||
url(
|
||||
r'^oauth2-with-scope-test/$',
|
||||
test_authentication.MockView.as_view(
|
||||
authentication_classes=[OAuth2AuthenticationAllowInactiveUser],
|
||||
permission_classes=[permissions.TokenHasReadWriteScope]
|
||||
)
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class OAuth2Tests(TestCase):
|
||||
"""OAuth 2.0 authentication"""
|
||||
urls = 'openedx.core.lib.api.tests.test_authentication'
|
||||
|
||||
class OAuth2AuthenticationAllowInactiveUserTestCase(test_authentication.OAuth2Tests):
|
||||
"""
|
||||
Tests the OAuth2AuthenticationAllowInactiveUser class by running all the existing tests in
|
||||
OAuth2Tests but with the is_active flag on the user set to False.
|
||||
"""
|
||||
def setUp(self):
|
||||
self.csrf_client = APIClient(enforce_csrf_checks=True)
|
||||
self.username = 'john'
|
||||
self.email = 'lennon@thebeatles.com'
|
||||
self.password = 'password'
|
||||
self.user = User.objects.create_user(self.username, self.email, self.password)
|
||||
super(OAuth2AuthenticationAllowInactiveUserTestCase, self).setUp()
|
||||
|
||||
self.CLIENT_ID = 'client_key' # pylint: disable=invalid-name
|
||||
self.CLIENT_SECRET = 'client_secret' # pylint: disable=invalid-name
|
||||
self.ACCESS_TOKEN = "access_token" # pylint: disable=invalid-name
|
||||
self.REFRESH_TOKEN = "refresh_token" # pylint: disable=invalid-name
|
||||
|
||||
self.oauth2_client = oauth2_provider.oauth2.models.Client.objects.create(
|
||||
client_id=self.CLIENT_ID,
|
||||
client_secret=self.CLIENT_SECRET,
|
||||
redirect_uri='',
|
||||
client_type=0,
|
||||
name='example',
|
||||
user=None,
|
||||
)
|
||||
|
||||
self.access_token = oauth2_provider.oauth2.models.AccessToken.objects.create(
|
||||
token=self.ACCESS_TOKEN,
|
||||
client=self.oauth2_client,
|
||||
user=self.user,
|
||||
)
|
||||
self.refresh_token = oauth2_provider.oauth2.models.RefreshToken.objects.create(
|
||||
user=self.user,
|
||||
access_token=self.access_token,
|
||||
client=self.oauth2_client
|
||||
)
|
||||
|
||||
# This is the a change we've made from the django-rest-framework-oauth version
|
||||
# of these tests.
|
||||
# set the user's is_active flag to False.
|
||||
self.user.is_active = False
|
||||
self.user.save()
|
||||
|
||||
# This is the a change we've made from the django-rest-framework-oauth version
|
||||
# of these tests.
|
||||
# Override the SCOPE_NAME_DICT setting for tests for oauth2-with-scope-test. This is
|
||||
# needed to support READ and WRITE scopes as they currently aren't supported by the
|
||||
# edx-auth2-provider, and their scope values collide with other scopes defined in the
|
||||
# edx-auth2-provider.
|
||||
scope.SCOPE_NAME_DICT = {'read': constants.READ, 'write': constants.WRITE}
|
||||
|
||||
def _create_authorization_header(self, token=None): # pylint: disable=missing-docstring
|
||||
return "Bearer {0}".format(token or self.access_token.token)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_with_wrong_authorization_header_token_type_failing(self):
|
||||
"""Ensure that a wrong token type lead to the correct HTTP error status code"""
|
||||
auth = "Wrong token-type-obviously"
|
||||
response = self.csrf_client.get('/oauth2-test/', {}, HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
response = self.csrf_client.get('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_with_wrong_authorization_header_token_format_failing(self):
|
||||
"""Ensure that a wrong token format lead to the correct HTTP error status code"""
|
||||
auth = "Bearer wrong token format"
|
||||
response = self.csrf_client.get('/oauth2-test/', {}, HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
response = self.csrf_client.get('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_with_wrong_authorization_header_token_failing(self):
|
||||
"""Ensure that a wrong token lead to the correct HTTP error status code"""
|
||||
auth = "Bearer wrong-token"
|
||||
response = self.csrf_client.get('/oauth2-test/', {}, HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
response = self.csrf_client.get('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_with_wrong_authorization_header_token_missing(self):
|
||||
"""Ensure that a missing token lead to the correct HTTP error status code"""
|
||||
auth = "Bearer"
|
||||
response = self.csrf_client.get('/oauth2-test/', {}, HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
response = self.csrf_client.get('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 401)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_passing_auth(self):
|
||||
"""Ensure GETing form over OAuth with correct client credentials succeed"""
|
||||
auth = self._create_authorization_header()
|
||||
response = self.csrf_client.get('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_passing_auth_url_transport(self):
|
||||
"""Ensure GETing form over OAuth with correct client credentials in form data succeed"""
|
||||
response = self.csrf_client.post(
|
||||
'/oauth2-test/',
|
||||
data={'access_token': self.access_token.token}
|
||||
)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_passing_auth_url_transport(self):
|
||||
"""Ensure GETing form over OAuth with correct client credentials in query succeed when DEBUG is True"""
|
||||
query = urlencode({'access_token': self.access_token.token})
|
||||
response = self.csrf_client.get('/oauth2-test-debug/?%s' % query)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_get_form_failing_auth_url_transport(self):
|
||||
"""Ensure GETing form over OAuth with correct client credentials in query fails when DEBUG is False"""
|
||||
query = urlencode({'access_token': self.access_token.token})
|
||||
response = self.csrf_client.get('/oauth2-test/?%s' % query)
|
||||
self.assertIn(response.status_code, (status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN))
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_passing_auth(self):
|
||||
"""Ensure POSTing form over OAuth with correct credentials passes and does not require CSRF"""
|
||||
auth = self._create_authorization_header()
|
||||
response = self.csrf_client.post('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_token_removed_failing_auth(self):
|
||||
"""Ensure POSTing when there is no OAuth access token in db fails"""
|
||||
self.access_token.delete()
|
||||
auth = self._create_authorization_header()
|
||||
response = self.csrf_client.post('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertIn(response.status_code, (status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN))
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_with_refresh_token_failing_auth(self):
|
||||
"""Ensure POSTing with refresh token instead of access token fails"""
|
||||
auth = self._create_authorization_header(token=self.refresh_token.token)
|
||||
response = self.csrf_client.post('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertIn(response.status_code, (status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN))
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_with_expired_access_token_failing_auth(self):
|
||||
"""Ensure POSTing with expired access token fails with an 'Invalid token' error"""
|
||||
self.access_token.expires = datetime.datetime.now() - datetime.timedelta(seconds=10) # 10 seconds late
|
||||
self.access_token.save()
|
||||
auth = self._create_authorization_header()
|
||||
response = self.csrf_client.post('/oauth2-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertIn(response.status_code, (status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN))
|
||||
self.assertIn('Invalid token', response.content)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_with_invalid_scope_failing_auth(self):
|
||||
"""Ensure POSTing with a readonly scope instead of a write scope fails"""
|
||||
read_only_access_token = self.access_token
|
||||
read_only_access_token.scope = oauth2_provider_scope.SCOPE_NAME_DICT['read']
|
||||
read_only_access_token.save()
|
||||
auth = self._create_authorization_header(token=read_only_access_token.token)
|
||||
response = self.csrf_client.get('/oauth2-with-scope-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
response = self.csrf_client.post('/oauth2-with-scope-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
||||
|
||||
@unittest.skipUnless(oauth2_provider, 'django-oauth2-provider not installed')
|
||||
def test_post_form_with_valid_scope_passing_auth(self):
|
||||
"""Ensure POSTing with a write scope succeed"""
|
||||
read_write_access_token = self.access_token
|
||||
read_write_access_token.scope = oauth2_provider_scope.SCOPE_NAME_DICT['write']
|
||||
read_write_access_token.save()
|
||||
auth = self._create_authorization_header(token=read_write_access_token.token)
|
||||
response = self.csrf_client.post('/oauth2-with-scope-test/', HTTP_AUTHORIZATION=auth)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
@@ -9,7 +9,6 @@ from django.utils.translation import ugettext as _
|
||||
from rest_framework import status, response
|
||||
from rest_framework.exceptions import APIException
|
||||
from rest_framework.permissions import IsAuthenticated
|
||||
from rest_framework.request import clone_request
|
||||
from rest_framework.response import Response
|
||||
from rest_framework.mixins import RetrieveModelMixin, UpdateModelMixin
|
||||
from rest_framework.generics import GenericAPIView
|
||||
@@ -194,23 +193,3 @@ class RetrievePatchAPIView(RetrieveModelMixin, UpdateModelMixin, GenericAPIView)
|
||||
add_serializer_errors(serializer, patch, field_errors)
|
||||
|
||||
return field_errors
|
||||
|
||||
def get_object_or_none(self):
|
||||
"""
|
||||
Retrieve an object or return None if the object can't be found.
|
||||
|
||||
NOTE: This replaces functionality that was removed in Django Rest Framework v3.1.
|
||||
"""
|
||||
try:
|
||||
return self.get_object()
|
||||
except Http404:
|
||||
if self.request.method == 'PUT':
|
||||
# For PUT-as-create operation, we need to ensure that we have
|
||||
# relevant permissions, as if this was a POST request. This
|
||||
# will either raise a PermissionDenied exception, or simply
|
||||
# return None.
|
||||
self.check_permissions(clone_request(self.request, 'POST'))
|
||||
else:
|
||||
# PATCH requests where the object does not exist should still
|
||||
# return a 404 response.
|
||||
raise
|
||||
|
||||
Reference in New Issue
Block a user