import itertools import functools from flask import render_template, redirect, url_for, request, flash from flask.views import MethodView from sqlalchemy import or_ from webargs import fields from webargs.flaskparser import use_kwargs from atlas_um.auth import claims_required, Role from atlas_um.settings import Settings from atlas_um.helpers.ordering import QueryOrdering from atlas_um.helpers.pagination import QueryPagination from atlas_um.pgdb import ClaimName, DNAAccount, ResourceGroup, ClaimValue class AdminSettingsBaseView(MethodView): """ Base class-based view for admin settings pages. Configurable attributes: - template - title - model - service_class - form_class - url_prefix - filter_args - ordering_fields - search_fields """ ordering_class = QueryOrdering pagination_class = QueryPagination decorators = [claims_required([Role(Role.Values.admin)])] success_message = "Saved!" template = None title = None model = None affected_by_field = None holder_model = None service_class = None form_class = None url_prefix = None filter_args = None holder_arg = None parent_arg = None ordering_fields = None search_fields = None @classmethod def register(cls, root_prefix, blueprint): base_url = f"{root_prefix}/{cls.url_prefix}" if cls.filter_args: base_url = f"{base_url}/{cls.filter_args}" blueprint.add_url_rule(base_url, view_func=cls.as_view(cls.url_prefix)) blueprint.add_url_rule( f"{base_url}/", view_func=cls.as_view(f"{cls.url_prefix}_update"), ) @use_kwargs( { "page": fields.Int(load_default=1), "per_page": fields.Int(load_default=Settings.ITEMS_PER_PAGE), "order": fields.Str(load_default=""), "item_search_term": fields.Str(load_default=""), }, location="query", ) def get(self, page, per_page, order, item_search_term, **kwargs): query = self.get_query() query = self.add_search(query, item_search_term) ordering = self.get_ordering(order=order, query=query) pagination = self.get_pagination( page=page, per_page=per_page, query=ordering.query ) objects = pagination.items form = self.get_form() return render_template( self.template, objects=objects, ordering=ordering, pagination=pagination, title=self.get_title(), form=form, resource_groups=ResourceGroup.query.active().order_by("name"), holder=self.get_holder(), global_claim_names=ClaimName.query.filter_global() .active() .order_by("friendly"), affected_count=self.get_affected_count(), parents=self.get_parents(), parent_id=self.get_parent_id(), url_for_create=functools.partial(self.url_for_create, pagination), url_for_update=functools.partial(self.url_for_update), item_search_term=item_search_term, ) def post(self, **kwargs): id = kwargs.pop("id", None) obj = self.get_object(id) if id else None form = self.get_form(obj) args = request.args.copy() if form.validate_on_submit(): service = self.service_class service.execute(obj, **form.data) flash(self.success_message, "positive") else: args["page"] = args.get("initial_page") message = ", ".join(itertools.chain(*form.errors.values())) flash(message, "negative") return redirect( url_for( request.endpoint.replace("_update", ""), **kwargs, **args, ) ) def get_title(self): return self.title def get_ordering(self, order, query): ordering = self.ordering_class( query, self.ordering_fields or [self.model.id, self.model.name, self.model.slug], order, ) return ordering def get_pagination(self, page, per_page, query): pagination = self.pagination_class( page=page, per_page=per_page, query=query ) return pagination def get_object(self, id): return self.get_query().filter_by(id=id).first_or_404() def get_holder(self): if not any((self.holder_model, self.holder_arg)): return None id = request.view_args.get(self.holder_arg) return self.holder_model.query.filter_by(id=id).first() def get_query(self): query = self.model.query if self.filter_args: query = query.filter_by(**request.view_args) if parent_id := self.get_parent_id(): query = query.filter_by(**{self.parent_arg: parent_id}) return query def add_search(self, query, search_term): if self.search_fields and search_term: searches = [] for field in self.search_fields: searches.append(field.ilike(f"%{search_term}%")) query = query.filter(or_(*searches)) return query def get_form(self, obj=None): extra_kwargs = request.view_args if self.filter_args else {} return self.form_class(obj=obj, **extra_kwargs) def get_affected_count(self): if self.affected_by_field: return { k: v for k, v in DNAAccount.query.affected_by_count( self.affected_by_field ) } return {} def get_parent_id(self): try: return int(request.args.get(self.parent_arg)) except (TypeError, ValueError): return def get_parents(self): parents = None holder = self.get_holder() if self.parent_arg and holder and holder.parent_id: parents = ClaimValue.query.parents_by_claim_name_id(holder.id) return parents @staticmethod def url_for_create(pagination, **kwargs): kwargs.update(request.view_args) kwargs.update(request.args) kwargs["page"] = pagination.total // pagination.per_page + 1 return url_for(request.endpoint, **kwargs) @staticmethod def url_for_update(**kwargs): kwargs.update(request.view_args) kwargs.update(request.args) return url_for(request.endpoint + "_update", **kwargs)