feat: post handler to sync provider_data (#30107)
* feat: post handler to sync provider_data this allows us to read provider_data metadata from a remote metadata url. reuses code from the task that currently processes all proiderconfigs in a batch ENT-5482 * feat: lint fixes * test: add test for sync_provider_data * test: add case for update * fix: lint fix * fix: lint fix * feat: use exc_info to report error better * feat: update log message
This commit is contained in:
@@ -1,18 +1,20 @@
|
||||
# pylint: disable=missing-module-docstring
|
||||
import copy
|
||||
import pytz
|
||||
from uuid import uuid4 # lint-amnesty, pylint: disable=wrong-import-order
|
||||
from datetime import datetime # lint-amnesty, pylint: disable=wrong-import-order
|
||||
from unittest import mock
|
||||
from uuid import uuid4 # lint-amnesty, pylint: disable=wrong-import-order
|
||||
|
||||
import pytz
|
||||
from django.contrib.sites.models import Site
|
||||
from django.urls import reverse
|
||||
from django.utils.http import urlencode
|
||||
from enterprise.constants import ENTERPRISE_ADMIN_ROLE, ENTERPRISE_LEARNER_ROLE
|
||||
from enterprise.models import EnterpriseCustomer, EnterpriseCustomerIdentityProvider
|
||||
from rest_framework import status
|
||||
from rest_framework.test import APITestCase
|
||||
|
||||
from enterprise.models import EnterpriseCustomer, EnterpriseCustomerIdentityProvider
|
||||
from enterprise.constants import ENTERPRISE_ADMIN_ROLE, ENTERPRISE_LEARNER_ROLE
|
||||
from common.djangoapps.student.tests.factories import UserFactory
|
||||
from common.djangoapps.third_party_auth.models import SAMLProviderData, SAMLProviderConfig
|
||||
from common.djangoapps.third_party_auth.models import SAMLProviderConfig, SAMLProviderData
|
||||
from common.djangoapps.third_party_auth.tests.samlutils import set_jwt_cookie
|
||||
from common.djangoapps.third_party_auth.tests.utils import skip_unless_thirdpartyauth
|
||||
from common.djangoapps.third_party_auth.utils import convert_saml_slug_provider_id
|
||||
@@ -180,3 +182,35 @@ class SAMLProviderDataTests(APITestCase):
|
||||
set_jwt_cookie(self.client, self.user, [(ENTERPRISE_ADMIN_ROLE, BAD_ENTERPRISE_ID)])
|
||||
response = self.client.get(url, format='json')
|
||||
assert response.status_code == status.HTTP_403_FORBIDDEN
|
||||
|
||||
@mock.patch('common.djangoapps.third_party_auth.samlproviderdata.views.fetch_metadata_xml')
|
||||
@mock.patch('common.djangoapps.third_party_auth.samlproviderdata.views.parse_metadata_xml')
|
||||
def test_sync_one_provider_data_success(self, mock_parse, mock_fetch):
|
||||
"""
|
||||
POST auth/saml/v0/provider_data/sync_provider_data -d data
|
||||
"""
|
||||
mock_fetch.return_value = '<?xml><a>tag</a>'
|
||||
public_key = 'askdjf;sakdjfs;adkfjas;dkfjas;dkfjas;dlkfj'
|
||||
sso_url = 'https://fake-test.id'
|
||||
expires_at = datetime.now()
|
||||
mock_parse.return_value = (public_key, sso_url, expires_at)
|
||||
url = reverse('saml_provider_data-sync-provider-data')
|
||||
data = {
|
||||
'entity_id': 'http://entity-id-1',
|
||||
'metadata_url': 'http://a-url',
|
||||
'enterprise_customer_uuid': ENTERPRISE_ID,
|
||||
}
|
||||
SAMLProviderData.objects.all().delete()
|
||||
orig_count = SAMLProviderData.objects.count()
|
||||
|
||||
response = self.client.post(url, data)
|
||||
|
||||
assert response.status_code == status.HTTP_201_CREATED
|
||||
assert response.data == " Created new record for SAMLProviderData for entityID http://entity-id-1"
|
||||
assert SAMLProviderData.objects.count() == orig_count + 1
|
||||
|
||||
# should only update this time
|
||||
response = self.client.post(url, data)
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
assert response.data == (" Updated existing SAMLProviderData for entityID http://entity-id-1")
|
||||
assert SAMLProviderData.objects.count() == orig_count + 1
|
||||
|
||||
@@ -1,21 +1,32 @@
|
||||
"""
|
||||
Viewset for auth/saml/v0/samlproviderdata
|
||||
"""
|
||||
import logging
|
||||
|
||||
from django.shortcuts import get_object_or_404
|
||||
from django.http import Http404
|
||||
from django.shortcuts import get_object_or_404
|
||||
from edx_rbac.mixins import PermissionRequiredMixin
|
||||
from edx_rest_framework_extensions.auth.jwt.authentication import JwtAuthentication
|
||||
from rest_framework import permissions, viewsets
|
||||
from rest_framework.authentication import SessionAuthentication
|
||||
from rest_framework.exceptions import ParseError
|
||||
|
||||
from enterprise.models import EnterpriseCustomerIdentityProvider
|
||||
from common.djangoapps.third_party_auth.utils import validate_uuid4_string, convert_saml_slug_provider_id
|
||||
from rest_framework import permissions, status, viewsets
|
||||
from rest_framework.authentication import SessionAuthentication
|
||||
from rest_framework.decorators import action
|
||||
from rest_framework.exceptions import ParseError
|
||||
from rest_framework.response import Response
|
||||
|
||||
from common.djangoapps.third_party_auth.utils import (
|
||||
convert_saml_slug_provider_id,
|
||||
create_or_update_saml_provider_data,
|
||||
fetch_metadata_xml,
|
||||
parse_metadata_xml,
|
||||
validate_uuid4_string
|
||||
)
|
||||
|
||||
from ..models import SAMLProviderConfig, SAMLProviderData
|
||||
from .serializers import SAMLProviderDataSerializer
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SAMLProviderDataMixin:
|
||||
authentication_classes = [JwtAuthentication, SessionAuthentication]
|
||||
@@ -36,6 +47,7 @@ class SAMLProviderDataViewSet(PermissionRequiredMixin, SAMLProviderDataMixin, vi
|
||||
POST /auth/saml/v0/provider_data/ -d postData (must contain 'enterprise_customer_uuid')
|
||||
DELETE /auth/saml/v0/provider_data/:pk -d postData (must contain 'enterprise_customer_uuid')
|
||||
PATCH /auth/saml/v0/provider_data/:pk -d postData (must contain 'enterprise_customer_uuid')
|
||||
POST /auth/saml/v0/provider_data/sync_provider_data (fetches metadata info from metadata url provided)
|
||||
|
||||
"""
|
||||
permission_required = 'enterprise.can_access_admin_dashboard'
|
||||
@@ -81,3 +93,37 @@ class SAMLProviderDataViewSet(PermissionRequiredMixin, SAMLProviderDataMixin, vi
|
||||
Retrieve an EnterpriseCustomer to do auth against
|
||||
"""
|
||||
return self.requested_enterprise_uuid
|
||||
|
||||
@action(detail=False, methods=['post'])
|
||||
def sync_provider_data(self, request):
|
||||
"""
|
||||
Creates or updates a SAMProviderData record using info fetched from remote SAML metadata
|
||||
For now we will require entityID but in future we will enhance this to try and extract entityID
|
||||
from the metadata file, and make entityId optional, and return error response if there are
|
||||
multiple entityIDs listed so that the user can choose and retry with a specified entityID
|
||||
"""
|
||||
entity_id = request.POST.get('entity_id')
|
||||
metadata_url = request.POST.get('metadata_url')
|
||||
if not entity_id:
|
||||
return Response('entity_id is required!', status.HTTP_400_BAD_REQUEST)
|
||||
if not metadata_url:
|
||||
return Response('metadata_url is required!', status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
# part 1: fetch information from remote metadata based on metadataUrl in samlproviderconfig
|
||||
xml = fetch_metadata_xml(metadata_url)
|
||||
|
||||
# part 2: create/update samlproviderdata
|
||||
log.info("Processing IdP with entityID %s", entity_id)
|
||||
public_key, sso_url, expires_at = parse_metadata_xml(xml, entity_id)
|
||||
changed = create_or_update_saml_provider_data(entity_id, public_key, sso_url, expires_at)
|
||||
if changed:
|
||||
str_message = f" Created new record for SAMLProviderData for entityID {entity_id}"
|
||||
log.info(str_message)
|
||||
response = str_message
|
||||
http_status = status.HTTP_201_CREATED
|
||||
else:
|
||||
str_message = f" Updated existing SAMLProviderData for entityID {entity_id}"
|
||||
log.info(str_message)
|
||||
response = str_message
|
||||
http_status = status.HTTP_200_OK
|
||||
return Response(response, status=http_status)
|
||||
|
||||
@@ -7,13 +7,16 @@ import logging
|
||||
|
||||
import requests
|
||||
from celery import shared_task
|
||||
from django.utils.timezone import now
|
||||
from edx_django_utils.monitoring import set_code_owner_attribute
|
||||
from lxml import etree
|
||||
from requests import exceptions
|
||||
|
||||
from common.djangoapps.third_party_auth.models import SAMLConfiguration, SAMLProviderConfig, SAMLProviderData
|
||||
from common.djangoapps.third_party_auth.utils import MetadataParseError, parse_metadata_xml
|
||||
from common.djangoapps.third_party_auth.models import SAMLConfiguration, SAMLProviderConfig
|
||||
from common.djangoapps.third_party_auth.utils import (
|
||||
MetadataParseError,
|
||||
create_or_update_saml_provider_data,
|
||||
parse_metadata_xml,
|
||||
)
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -85,7 +88,7 @@ def fetch_saml_metadata():
|
||||
for entity_id in entity_ids:
|
||||
log.info("Processing IdP with entityID %s", entity_id)
|
||||
public_key, sso_url, expires_at = parse_metadata_xml(xml, entity_id)
|
||||
changed = _update_data(entity_id, public_key, sso_url, expires_at)
|
||||
changed = create_or_update_saml_provider_data(entity_id, public_key, sso_url, expires_at)
|
||||
if changed:
|
||||
log.info(f"→ Created new record for SAMLProviderData for entityID {entity_id}")
|
||||
num_updated += 1
|
||||
@@ -124,28 +127,3 @@ def fetch_saml_metadata():
|
||||
|
||||
# Return counts for total, skipped, attempted, updated, and failed, along with any failure messages
|
||||
return num_total, num_skipped, num_attempted, num_updated, len(failure_messages), failure_messages
|
||||
|
||||
|
||||
def _update_data(entity_id, public_key, sso_url, expires_at):
|
||||
"""
|
||||
Update/Create the SAMLProviderData for the given entity ID.
|
||||
Return value:
|
||||
False if nothing has changed and existing data's "fetched at" timestamp is just updated.
|
||||
True if a new record was created. (Either this is a new provider or something changed.)
|
||||
"""
|
||||
data_obj = SAMLProviderData.current(entity_id)
|
||||
fetched_at = now()
|
||||
if data_obj and (data_obj.public_key == public_key and data_obj.sso_url == sso_url):
|
||||
data_obj.expires_at = expires_at
|
||||
data_obj.fetched_at = fetched_at
|
||||
data_obj.save()
|
||||
return False
|
||||
else:
|
||||
SAMLProviderData.objects.create(
|
||||
entity_id=entity_id,
|
||||
fetched_at=fetched_at,
|
||||
expires_at=expires_at,
|
||||
sso_url=sso_url,
|
||||
public_key=public_key,
|
||||
)
|
||||
return True
|
||||
|
||||
@@ -3,29 +3,68 @@ Utility functions for third_party_auth
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
from uuid import UUID
|
||||
|
||||
import dateutil.parser
|
||||
import pytz
|
||||
import requests
|
||||
from django.contrib.auth.models import User # lint-amnesty, pylint: disable=imported-auth-user
|
||||
from django.utils.timezone import now
|
||||
from enterprise.models import EnterpriseCustomerIdentityProvider, EnterpriseCustomerUser
|
||||
from lxml import etree
|
||||
from onelogin.saml2.utils import OneLogin_Saml2_Utils
|
||||
from requests import exceptions
|
||||
from social_core.pipeline.social_auth import associate_by_email
|
||||
|
||||
from common.djangoapps.third_party_auth.models import OAuth2ProviderConfig
|
||||
from common.djangoapps.third_party_auth.models import OAuth2ProviderConfig, SAMLProviderData
|
||||
from openedx.core.djangolib.markup import Text
|
||||
|
||||
from . import provider
|
||||
|
||||
SAML_XML_NS = 'urn:oasis:names:tc:SAML:2.0:metadata' # The SAML Metadata XML namespace
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MetadataParseError(Exception):
|
||||
""" An error occurred while parsing the SAML metadata from an IdP """
|
||||
pass # lint-amnesty, pylint: disable=unnecessary-pass
|
||||
|
||||
|
||||
def fetch_metadata_xml(url):
|
||||
"""
|
||||
Fetches IDP metadata from provider url
|
||||
Returns: xml document
|
||||
"""
|
||||
try:
|
||||
log.info("Fetching %s", url)
|
||||
if not url.lower().startswith('https'):
|
||||
log.warning("This SAML metadata URL is not secure! It should use HTTPS. (%s)", url)
|
||||
response = requests.get(url, verify=True) # May raise HTTPError or SSLError or ConnectionError
|
||||
response.raise_for_status() # May raise an HTTPError
|
||||
|
||||
try:
|
||||
parser = etree.XMLParser(remove_comments=True)
|
||||
xml = etree.fromstring(response.content, parser)
|
||||
except etree.XMLSyntaxError: # lint-amnesty, pylint: disable=try-except-raise
|
||||
raise
|
||||
# TODO: Can use OneLogin_Saml2_Utils to validate signed XML if anyone is using that
|
||||
return xml
|
||||
except (exceptions.SSLError, exceptions.HTTPError, exceptions.RequestException, MetadataParseError) as error:
|
||||
# Catch and process exception in case of errors during fetching and processing saml metadata.
|
||||
# Here is a description of each exception.
|
||||
# SSLError is raised in case of errors caused by SSL (e.g. SSL cer verification failure etc.)
|
||||
# HTTPError is raised in case of unexpected status code (e.g. 500 error etc.)
|
||||
# RequestException is the base exception for any request related error that "requests" lib raises.
|
||||
# MetadataParseError is raised if there is error in the fetched meta data (e.g. missing @entityID etc.)
|
||||
log.exception(str(error), exc_info=error)
|
||||
raise error
|
||||
except etree.XMLSyntaxError as error:
|
||||
log.exception(str(error), exc_info=error)
|
||||
raise error
|
||||
|
||||
|
||||
def parse_metadata_xml(xml, entity_id):
|
||||
"""
|
||||
Given an XML document containing SAML 2.0 metadata, parse it and return a tuple of
|
||||
@@ -125,6 +164,31 @@ def get_user_from_email(details):
|
||||
return None
|
||||
|
||||
|
||||
def create_or_update_saml_provider_data(entity_id, public_key, sso_url, expires_at):
|
||||
"""
|
||||
Update/Create the SAMLProviderData for the given entity ID.
|
||||
Return value:
|
||||
False if nothing has changed and existing data's "fetched at" timestamp is just updated.
|
||||
True if a new record was created. (Either this is a new provider or something changed.)
|
||||
"""
|
||||
data_obj = SAMLProviderData.current(entity_id)
|
||||
fetched_at = now()
|
||||
if data_obj and (data_obj.public_key == public_key and data_obj.sso_url == sso_url):
|
||||
data_obj.expires_at = expires_at
|
||||
data_obj.fetched_at = fetched_at
|
||||
data_obj.save()
|
||||
return False
|
||||
else:
|
||||
SAMLProviderData.objects.create(
|
||||
entity_id=entity_id,
|
||||
fetched_at=fetched_at,
|
||||
expires_at=expires_at,
|
||||
sso_url=sso_url,
|
||||
public_key=public_key,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def convert_saml_slug_provider_id(provider): # lint-amnesty, pylint: disable=redefined-outer-name
|
||||
"""
|
||||
Provider id is stored with the backend type prefixed to it (ie "saml-")
|
||||
|
||||
Reference in New Issue
Block a user