Skip to content
Merged
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: 2 additions & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,8 @@ clean:
$(COMPOSE) down --volumes --remove-orphans

format: build
$(COMPOSE) run --rm test uv run ruff format .
$(COMPOSE) run --rm test uv run ruff format . \
--exclude "*/migrations/*.py"

lock:
$(COMPOSE) run --rm test uv lock
Expand Down
24 changes: 15 additions & 9 deletions citation/admin.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,13 @@
from django.contrib import admin
from django.contrib.admin.helpers import ActionForm
from django.contrib.auth.models import User
from django.db.models import OuterRef, Exists
from django.db import transaction
from django.db.models import Exists, OuterRef
from django.utils.translation import gettext_lazy as _

from .models import (
Author,
AuditCommand,
Author,
CodeArchiveUrl,
Container,
ModelDocumentation,
Expand All @@ -18,6 +19,7 @@
SuggestedMerge,
Tag,
)
from .signals import notify_publications_changed


class PublicationStatusListFilter(admin.SimpleListFilter):
Expand Down Expand Up @@ -46,16 +48,20 @@ def queryset(self, request, queryset):

def assign_curator(modeladmin, request, queryset):
assigned_curator_id = request.POST["assigned_curator_id"]
publication_ids = list(queryset.values_list("pk", flat=True))
audit_command = AuditCommand(
creator=request.user,
action=AuditCommand.Action.MANUAL,
)

user = request.user
audit_command = AuditCommand(creator=user, action=AuditCommand.Action.MANUAL)

# Does not seem to be a Haystack method to update records based on a queryset so records are updated one at a time
# to keep the Solr index in sync
for publication in queryset:
publication.log_update(
with transaction.atomic():
queryset.log_update(
audit_command=audit_command, assigned_curator_id=assigned_curator_id
)
notify_publications_changed(
sender=Publication,
publication_ids=publication_ids,
)


assign_curator.short_description = "Assign Curator to Publications"
Expand Down
3 changes: 3 additions & 0 deletions citation/apps.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
from pathlib import Path

from django.apps import AppConfig


class CitationConfig(AppConfig):
name = "citation"
path = str(Path(__file__).resolve().parent)
default_auto_field = "django.db.models.BigAutoField"

def ready(self):
Expand Down
223 changes: 139 additions & 84 deletions citation/export_data.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,54 @@
import csv
import pathlib

from citation.models import Publication, Platform, Sponsor
import numpy as np
import pandas as pd
from django.contrib.postgres.aggregates import ArrayAgg, StringAgg
from django.core.exceptions import FieldDoesNotExist
from django.db import models
from django.db.models import Count, F, OuterRef, Q, Value
from django.db.models.functions import Concat, Trim

from citation.models import (
Author,
CodeArchiveUrl,
ModelDocumentation,
Platform,
Publication,
PublicationAuthors,
PublicationCitations,
PublicationModelDocumentations,
PublicationPlatforms,
PublicationSponsors,
Sponsor,
)


CSV_DEFAULT_HEADER = [
"id",
"title",
"abstract",
"short_title",
"contact_email",
"email_sent_count",
"contact_author_name",
"is_primary",
"doi",
"series_text",
"series_title",
"series",
"issue",
"volume",
"pages",
"author_names",
"container__issn",
"container__name",
"year_published",
]


# Streaming CSV taken from
# https://docs.djangoproject.com/en/2.1/howto/outputting-csv/
# Streaming CSV follows Django's documented pseudo-buffer pattern:
# https://docs.djangoproject.com/en/5.2/howto/outputting-csv/


class Echo:
Expand All @@ -19,37 +63,56 @@ def write(self, value):

class CategoricalVariable:
def __init__(self, levels):
self._levels = levels
self._levels = tuple(levels)

def dense_encode(self, values):
return [(level in values) for level in self._levels]
values = set(values)
return [level in values for level in self._levels]

def __iter__(self):
return iter(self._levels)


class PublicationCSVExporter:
def __init__(self, attributes=None):
self.m2m_attributes = [fields.name for fields in Publication._meta.many_to_many]
self.platforms = CategoricalVariable(
Platform.objects.all().values_list("name", flat=True).order_by("name")
)
self.sponsors = CategoricalVariable(
Sponsor.objects.all().values_list("name", flat=True).order_by("name")
)
if attributes is None:
self.attributes = CSV_DEFAULT_HEADER
else:
self.attributes = attributes
annotated_attributes = frozenset({"author_names"})
categorical_models = {"platforms": Platform, "sponsors": Sponsor}

def __init__(self, attributes=None):
self.attributes = list(CSV_DEFAULT_HEADER if attributes is None else attributes)
self.m2m_attributes = {field.name for field in Publication._meta.many_to_many}
self.categorical_variables = {}
self.verify_attributes()

@classmethod
def attribute_exists(cls, name):
if name in cls.annotated_attributes:
return True

model = Publication
attributes = name.split("__")
for position, attribute in enumerate(attributes):
if not hasattr(model, attribute):
return False
if position < len(attributes) - 1:
try:
field = model._meta.get_field(attribute)
except FieldDoesNotExist:
return False
if field.many_to_many or field.one_to_many:
return False
model = field.related_model
if model is None:
return False
return True

def verify_attributes(self):
for name in self.attributes:
if not hasattr(Publication, name):
if not self.attribute_exists(name):
raise AttributeError(
"Publication model doesn't have attribute :" + name
f"Publication model doesn't have attribute: {name}"
)
if name in self.m2m_attributes and name not in self.categorical_models:
raise AttributeError(f"Unsupported many-to-many attribute: {name}")

def get_header(self):
header = []
Expand All @@ -59,90 +122,81 @@ def get_header(self):
header.append(name)
header.extend(self.get_all_m2m_levels(name))
else:
header.append(name.strip().replace("_", " "))
header.append(name)
return header

def get_all_m2m_levels(self, name):
if name in ["sponsors", "platforms"]:
return getattr(self, name)
else:
raise AttributeError("Forgot to declare " + name + " m2m attribute")
try:
model = self.categorical_models[name]
except KeyError as error:
raise AttributeError(
f"Unsupported many-to-many attribute: {name}"
) from error

if name not in self.categorical_variables:
levels = model.objects.values_list("name", flat=True).order_by("name")
self.categorical_variables[name] = CategoricalVariable(levels)
return self.categorical_variables[name]

@staticmethod
def get_attribute(pub, name):
value = pub
for attribute in name.split("__"):
if value is None:
return ""
value = getattr(value, attribute)
return value

def get_row(self, pub):
row = []
for name in self.attributes:
if name in self.m2m_attributes:
source = getattr(pub, name)
pub_m2m_data_list = source.all().values_list("name", flat=True)
row.append(pub_m2m_data_list)
pub_m2m_data_list = sorted(item.name for item in source.all())
row.append("; ".join(pub_m2m_data_list))
row.extend(
self.get_all_m2m_levels(name).dense_encode(pub_m2m_data_list)
)
else:
row.append(getattr(pub, name))
row.append(self.get_attribute(pub, name))
return row

def get_publications(self):
publications = Publication.api.primary()
if any(name.startswith("container__") for name in self.attributes):
publications = publications.select_related("container")
m2m_attributes = self.m2m_attributes.intersection(self.attributes)
if m2m_attributes:
publications = publications.prefetch_related(*m2m_attributes)
if "author_names" in self.attributes:
publications = publications.annotate(
author_names=StringAgg(
Trim(
Concat(
F("creators__given_name"),
Value(" "),
F("creators__family_name"),
)
),
delimiter="; ",
order_by=("creators__family_name", "creators__given_name"),
)
)
return publications

def rows(self):
yield self.get_header()
for publication in self.get_publications():
yield self.get_row(publication)

def write_all(self, file):
writer = csv.writer(file, delimiter=",")
writer.writerow(self.get_header())
publications = Publication.api.primary()
for pub in publications:
writer.writerow(self.get_row(pub))
writer.writerows(self.rows())
return writer

def stream(self):
pseudo_buffer = Echo()
writer = csv.writer(pseudo_buffer, delimiter=",")
writer.writerow(self.get_header())
publications = Publication.api.primary()
for pub in publications:
yield writer.writerow(self.get_row(pub))


import numpy as np
import pandas as pd
import pathlib

from django.contrib.postgres.aggregates import ArrayAgg
from django.db import models
from django.db.models import F, Count, Q, Value, Sum, OuterRef
from django.db.models.functions import Concat

from citation.models import (
Publication,
Platform,
Sponsor,
PublicationCitations,
PublicationAuthors,
Author,
CodeArchiveUrl,
PublicationPlatforms,
PublicationSponsors,
PublicationModelDocumentations,
ModelDocumentation,
)

CSV_DEFAULT_HEADER = [
"id",
"title",
"abstract",
"short_title",
"contact_email",
"email_sent_count",
"contact_author_name",
"is_primary",
"doi",
"series_text",
"series_title",
"series",
"issue",
"volume",
"pages",
"author_names",
"container__issn",
"container__name",
"year_published",
]
writer = csv.writer(Echo(), delimiter=",")
return (writer.writerow(row) for row in self.rows())


def get_queryset():
Expand Down Expand Up @@ -449,7 +503,8 @@ def get_publication_row(publication):


def export(path):
remove_recoded = lambda df: df[["raw_name"]].rename(columns={"raw_name": "name"})
def remove_recoded(df):
return df[["raw_name"]].rename(columns={"raw_name": "name"})

path = pathlib.Path(path)
publications = get_queryset()
Expand Down
Loading
Loading