Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions admin/preprints/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -289,6 +289,9 @@ def _copy_primary_file(self, preprint, file_guid):
latest_version = source_file.versions.order_by('-created').first()
if latest_version is None:
raise ValueError(f'File "{file_guid}" has no versions to copy.')
if latest_version.purged:
raise ValueError(f'File "{file_guid}" latest version was purged from storage and cannot be copied.')
preprint.set_storage_region(latest_version.region_id)
copied = copy_files(source_file, target_node=preprint, identifier=latest_version.identifier)
preprint.set_primary_file(copied, auth=self.request, save=True, ignore_permission=True)

Expand Down
39 changes: 39 additions & 0 deletions admin_tests/preprints/test_views.py
Original file line number Diff line number Diff line change
Expand Up @@ -1110,6 +1110,45 @@ def test_copies_primary_file_when_admin_is_not_a_contributor(self):
assert recovered.deleted is None
assert recovered.primary_file.copied_from_id == source.primary_file.id

def _source_in_other_region(self):
from osf_tests.factories import RegionFactory
other_region = RegionFactory()
source = PreprintFactory(provider=self.provider)
Preprint.objects.filter(id=source.id).update(region=other_region)
source.primary_file.versions.update(region=other_region)
return source, other_region

def test_recovered_preprint_uses_region_of_source_file(self):
source, other_region = self._source_in_other_region()
assert self.user.get_addon('osfstorage').default_region_id != other_region.id

response = self._post(self._base_data(file_guid=source._id))
assert response.status_code == 302

recovered = Preprint.load('abcde')
assert recovered.region_id == other_region.id
copied_versions = recovered.primary_file.versions.all()
assert copied_versions
for version in copied_versions:
assert version.region_id == other_region.id
assert version.location == source.primary_file.versions.get(identifier=version.identifier).location
assert not source.primary_file.versions.exclude(region=other_region).exists()

def test_second_recovered_version_keeps_region_of_previous_version(self):
source, other_region = self._source_in_other_region()
assert self._post(self._base_data(file_guid=source._id)).status_code == 302
assert self._post(self._base_data()).status_code == 302

assert Preprint.load('abcde_v2').region_id == other_region.id

def test_purged_source_version_is_rejected(self):
source = PreprintFactory(provider=self.provider)
source.primary_file.versions.update(purged=timezone.now())

response = self._post(self._base_data(file_guid=source._id))
assert response.status_code == 302
assert Preprint.load('abcde') is None

def test_unknown_source_guid_shows_error(self):
response = self._post(self._base_data(file_guid='zzzzz_v1'))
assert response.status_code == 302
Expand Down
14 changes: 11 additions & 3 deletions api/users/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,8 @@
from api.addons.views import AddonSettingsMixin
from api.base import permissions as base_permissions
from api.users.permissions import UserMessagePermissions
from api.base.exceptions import Conflict, UserGone
from api.base.exceptions import Conflict, ServiceUnavailableError, UserGone
from osf.exceptions import OrcidRevocationError
from api.base.filters import ListFilterMixin, PreprintFilterMixin
from api.base.parsers import (
JSONAPIRelationshipParser,
Expand Down Expand Up @@ -618,13 +619,20 @@ def get_object(self):

def perform_destroy(self, instance):
user = self.get_user()
identity_id = self.kwargs['identity_id']
provider = self.kwargs['identity_id']
try:
user.external_identity.pop(identity_id)
identity_ids = list(user.external_identity[provider].keys())
except KeyError:
raise NotFound('Requested external identity could not be found.')
if not user.has_usable_password():
user.set_password(str(uuid.uuid4()))

try:
for identity_id in identity_ids:
user.disconnect_external_identity(provider, identity_id)
except OrcidRevocationError:
raise ServiceUnavailableError(detail='Unable to revoke ORCiD access at this time. Please try again later.')

user.save()


Expand Down
65 changes: 65 additions & 0 deletions api_tests/users/views/test_user_external_identities.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
from unittest import mock

import pytest
import requests

from osf_tests.factories import AuthUserFactory
from api.base.settings.defaults import API_BASE
Expand Down Expand Up @@ -87,6 +90,68 @@ def test_delete_204(self, app, user, url):
}
}

def test_delete_revokes_orcid_token_and_removes_it(self, app, user, url):
user.external_identity_tokens = {
'ORCID': {'0000-0001-9143-4653': {'access_token': 'fake-orcid-token'}},
}
user.save()

with mock.patch('osf.models.user.requests_retry_session') as mock_retry_session:
mock_session = mock.Mock()
mock_session.post.return_value = mock.Mock(status_code=200, text='')
mock_retry_session.return_value = mock_session

res = app.delete(url, auth=user.auth)

assert res.status_code == 204
mock_session.post.assert_called_once()
assert mock_session.post.call_args.kwargs['data']['token'] == 'fake-orcid-token'

user.refresh_from_db()
assert user.external_identity == {
'LOTUS': {
'0000-0001-9143-4652': 'LINK'
}
}
assert 'ORCID' not in user.external_identity_tokens

def test_delete_orcid_fails_and_keeps_data_when_revoke_fails(self, app, user, url):
user.external_identity_tokens = {
'ORCID': {'0000-0001-9143-4653': {'access_token': 'fake-orcid-token'}},
}
user.save()

with mock.patch('osf.models.user.requests_retry_session') as mock_retry_session:
mock_session = mock.Mock()
mock_session.post.return_value = mock.Mock(status_code=401, text='invalid_client')
mock_session.post.return_value.raise_for_status.side_effect = requests.exceptions.HTTPError('401')
mock_retry_session.return_value = mock_session

res = app.delete(url, auth=user.auth, expect_errors=True)

assert res.status_code == 503

user.refresh_from_db()
assert user.external_identity['ORCID'] == {'0000-0001-9143-4653': 'VERIFIED'}
assert user.external_identity_tokens['ORCID'] == {'0000-0001-9143-4653': {'access_token': 'fake-orcid-token'}}

def test_delete_non_orcid_identity_does_not_call_orcid_api(self, app, user):
url = f'/{API_BASE}users/{user._id}/settings/identities/LOTUS/'

with mock.patch('osf.models.user.requests_retry_session') as mock_retry_session:
res = app.delete(url, auth=user.auth)

assert res.status_code == 204
mock_retry_session.assert_not_called()

user.refresh_from_db()
assert 'LOTUS' not in user.external_identity
assert user.external_identity == {
'ORCID': {
'0000-0001-9143-4653': 'VERIFIED'
}
}

def test_anonymous_gets_401(self, app, url):
res = app.get(url, expect_errors=True)
assert res.status_code == 401
Expand Down
5 changes: 5 additions & 0 deletions osf/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,11 @@ class UserStateError(OSFError):
pass


class OrcidRevocationError(OSFError):
"""Raised when ORCiD fails to revoke an access token, e.g. when disconnecting an ORCiD identity."""
pass


class InstitutionAffiliationStateError(OSFError):
pass

Expand Down
5 changes: 5 additions & 0 deletions osf/management/commands/repair_recovered_preprint_files.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,11 @@ def repair_recovered_preprint_files(dry_run=False, user_id=None):

auth = Auth(acting_user or donor.creator)
latest_version = donor.primary_file.versions.order_by('-created').first()
if latest_version.purged:
logger.warning(f'{broken._id}: donor latest version {latest_version.id} is purged, skipping')
continue
# keep the donor's storage region, the copied blob does not move between buckets
broken.set_storage_region(latest_version.region_id)
copied = copy_files(donor.primary_file, target_node=broken, identifier=latest_version.identifier)
broken.set_primary_file(copied, auth=auth, save=True)
logger.info(f'{broken._id}: attached primary_file {copied._id} (copied from {donor._id})')
Expand Down
129 changes: 129 additions & 0 deletions osf/management/commands/repair_recovered_preprint_regions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,129 @@
import logging

from django.core.management.base import BaseCommand

from addons.osfstorage import settings as osfstorage_settings
from addons.osfstorage.models import Region
from osf.models import Preprint
from osf.models.admin_log_entry import AdminLogEntry, PREPRINT_RECOVERED, PREPRINT_RESTORED

logger = logging.getLogger(__name__)


def _bucket_of(region):
return (region.waterbutler_settings or {}).get('storage', {}).get('bucket')


def find_recovered_preprints(guids=None):
if guids:
object_ids = list(guids)
else:
object_ids = list(
AdminLogEntry.objects.filter(
action_flag__in=[PREPRINT_RECOVERED, PREPRINT_RESTORED],
).values_list('object_id', flat=True).distinct()
)
for object_id in object_ids:
preprint = Preprint.load(object_id)
if preprint is None:
logger.warning(f'{object_id}: preprint not found, skipping')
continue
yield preprint


def resolve_version_region(version, regions, client=None):
location = version.location or {}
bucket = location.get('bucket') or location.get(osfstorage_settings.WATERBUTLER_RESOURCE)
if bucket:
matches = [r for r in regions if _bucket_of(r) == bucket]
if len(matches) == 1:
return matches[0], 'location bucket'

blob_name = location.get('object')
if not client or not blob_name:
return None, (
f'location bucket "{bucket}" does not identify exactly one region; '
'rerun with --probe to look the blob up in each region bucket'
)
holders = [r for r in regions if _bucket_of(r) and client.bucket(_bucket_of(r)).blob(blob_name).exists()]
if not holders:
return None, 'blob not found in any region bucket (purged?)'
if version.region in holders:
return version.region, 'probe (current region holds the blob)'
if len(holders) > 1:
return None, 'blob found in several region buckets; ambiguous'
return holders[0], 'probe'


def repair_preprint_regions(preprints, dry_run=False, client=None):
regions = list(Region.objects.all())
stats = {'checked': 0, 'versions_fixed': 0, 'preprints_fixed': 0, 'unresolved': 0}

for preprint in preprints:
primary_file = preprint.primary_file
if primary_file is None:
logger.info(f'{preprint._id}: no primary file, skipping')
continue
stats['checked'] += 1

for version in primary_file.versions.select_related('region').order_by('created'):
if version.purged:
logger.warning(f'{preprint._id}: version {version.identifier} (FV {version.id}) is purged, skipping')
stats['unresolved'] += 1
continue
true_region, reason = resolve_version_region(version, regions, client=client)
if true_region is None:
logger.warning(f'{preprint._id}: version {version.identifier} (FV {version.id}) unresolved: {reason}')
stats['unresolved'] += 1
continue
if version.region_id == true_region.id:
continue
logger.info(
f'{preprint._id}: FV {version.id} region {version.region and version.region._id} -> '
f'{true_region._id} ({reason})'
)
stats['versions_fixed'] += 1
if not dry_run:
version.region = true_region
version.save()

latest = primary_file.versions.select_related('region').order_by('-created').first()
if latest and latest.region_id and latest.region_id != preprint.region_id:
logger.info(f'{preprint._id}: preprint region {preprint.region_id} -> {latest.region_id}')
stats['preprints_fixed'] += 1
if not dry_run:
preprint.set_storage_region(latest.region_id)

return stats


class Command(BaseCommand):
help = (
'Repair FileVersion/Preprint storage regions of admin-recovered preprints whose file copy was relabeled '
'to another region without moving the blob (Waterbutler NoSuchKey, "Missing PDF file").'
)

def add_arguments(self, parser):
parser.add_argument('--dry_run', action='store_true', default=False, help='Log changes without saving')
parser.add_argument('--guids', nargs='*', default=None, help='Only these preprint guids (e.g. stzmk_v1)')
parser.add_argument(
'--probe',
action='store_true',
default=False,
help='When a version location has no bucket, look the blob up in every region bucket via GCS',
)

def handle(self, *args, **options):
client = None
if options['probe']:
from google.cloud.storage.client import Client
from google.oauth2.service_account import Credentials
from website.settings import GCS_CREDS
client = Client(credentials=Credentials.from_service_account_file(GCS_CREDS))

if options['dry_run']:
logger.info('DRY RUN. Data will not be saved.')
stats = repair_preprint_regions(
find_recovered_preprints(options['guids']), dry_run=options['dry_run'], client=client,
)
logger.info(f'Done: {stats}')
17 changes: 17 additions & 0 deletions osf/models/preprint.py
Original file line number Diff line number Diff line change
Expand Up @@ -518,6 +518,10 @@ def create_version(cls, create_from_guid, auth, assign_version_number=None, igno
guid_version.save()
preprint.save(guid_ready=True, first_save=True, set_creator_as_contributor=False)

# The new version's file is copied from the previous version, so keep the same storage region
# rather than the acting user's default (the blob does not move between regions on copy).
preprint.set_storage_region(latest_version.region_id)

# Add contributors
for contributor in latest_version.contributor_set.all():
try:
Expand Down Expand Up @@ -1074,6 +1078,19 @@ def _set_default_region(self):
self.region_id = user_settings.default_region_id
self.save()

def set_storage_region(self, region_id):
"""Point this preprint at the storage region that actually holds its files.

Waterbutler takes the bucket from the file version's region, and files copied from another
preprint keep their blobs in the source region, so recreated versions must share that region
instead of the creating user's default. Uses a queryset update because `save()` refuses a
non-initial version that has no primary file yet.
"""
if not region_id or region_id == self.region_id:
return
self.region_id = region_id
Preprint.objects.filter(pk=self.pk).update(region_id=region_id)

def _add_creator_as_contributor(self):
self.add_contributor(self.creator, permissions=ADMIN, visible=True, log=False, save=True)

Expand Down
Loading
Loading