3636from .. import settings as app_settings
3737from ..base .forms import PasswordResetForm
3838from ..counters .exceptions import SkipCheck
39- from ..registration import REGISTRATION_METHOD_CHOICES
4039from ..utils import (
4140 get_group_checks ,
4241 get_organization_radius_settings ,
@@ -571,9 +570,13 @@ class RegisterSerializer(
571570 'verification in its "Organization RADIUS Settings."'
572571 ),
573572 default = "" ,
574- choices = REGISTRATION_METHOD_CHOICES ,
573+ choices = () ,
575574 )
576575
576+ def __init__ (self , * args , ** kwargs ):
577+ super ().__init__ (* args , ** kwargs )
578+ self .fields ["method" ].choices = app_settings .USER_SETTABLE_REGISTRATION_METHODS
579+
577580 def validate_phone_number (self , phone_number ):
578581 org = self .context ["view" ].organization
579582 if get_organization_radius_settings (org , "sms_verification" ):
@@ -688,9 +691,11 @@ def save(self, request):
688691 # the custom_signup method contains the openwisp specific logic
689692 self .custom_signup (request , user )
690693 # create a RegisteredUser object for every user that registers through API
691- RegisteredUser .objects .create (
694+ org = self .context ["view" ].organization
695+ RegisteredUser .get_or_create_for_user_and_org (
692696 user = user ,
693- method = self .validated_data ["method" ],
697+ organization = org ,
698+ defaults = {"method" : self .validated_data ["method" ]},
694699 )
695700 setup_user_email (request , user , [])
696701 return user
@@ -753,20 +758,64 @@ def save(self):
753758 # yet, tha will be done by the phone token validation view
754759 # once the phone number has been validated
755760 # at this point we flag the user as unverified again
756- self .user .registered_user .is_verified = False
757- self .user .registered_user .save ()
761+ org = self .context ["view" ].organization
762+ reg_user , _ = RegisteredUser .get_or_create_for_user_and_org (
763+ user = self .user ,
764+ organization = org ,
765+ defaults = {"is_verified" : False , "method" : "" },
766+ )
767+ reg_user .is_verified = False
768+ reg_user .save ()
769+
770+
771+ class UpdateRegisteredUserMethodSerializer (ValidatedModelSerializer ):
772+ method = serializers .ChoiceField (
773+ choices = app_settings .USER_SETTABLE_REGISTRATION_METHODS ,
774+ help_text = _ (
775+ "The registration method to set for the user. "
776+ "Cannot be 'pending_verification'."
777+ ),
778+ )
779+
780+ class Meta :
781+ model = RegisteredUser
782+ fields = ["method" ]
783+
784+ def __init__ (self , * args , ** kwargs ):
785+ super ().__init__ (* args , ** kwargs )
786+ self .fields ["method" ].choices = app_settings .USER_SETTABLE_REGISTRATION_METHODS
787+
788+ def validate_method (self , value ):
789+ if value == "pending_verification" :
790+ raise serializers .ValidationError (
791+ _ ("'pending_verification' cannot be set as a registration method." )
792+ )
793+ return value
794+
795+ def validate (self , attrs ):
796+ if self .instance .method != "pending_verification" :
797+ raise serializers .ValidationError (
798+ {
799+ "method" : _ (
800+ "Method can only be updated from pending verification state."
801+ )
802+ }
803+ )
804+ return attrs
805+
806+ def update (self , instance , validated_data ):
807+ instance .method = validated_data ["method" ]
808+ instance .save ()
809+ return instance
758810
759811
760812class RadiusUserSerializer (serializers .ModelSerializer ):
761813 """
762814 Used to return information about the logged in user
763815 """
764816
765- is_verified = serializers .BooleanField (source = "registered_user.is_verified" )
766- method = serializers .CharField (
767- source = "registered_user.method" ,
768- allow_null = True ,
769- )
817+ is_verified = serializers .SerializerMethodField ()
818+ method = serializers .SerializerMethodField ()
770819 password_expired = serializers .BooleanField (source = "has_password_expired" )
771820 radius_user_token = serializers .CharField (source = "radius_token.key" , default = None )
772821
@@ -786,3 +835,30 @@ class Meta:
786835 "password_expired" ,
787836 "radius_user_token" ,
788837 ]
838+
839+ def _get_registered_user (self , obj ):
840+ if not hasattr (self , "_registered_user_cache" ):
841+ self ._registered_user_cache = {}
842+ if obj .pk not in self ._registered_user_cache :
843+ view = self .context .get ("view" )
844+ organization = getattr (view , "organization" , None )
845+ reg_user = None
846+ # We iterate over .all() instead of using .filter() because callers
847+ # of this serializer (e.g. validate_auth_token) prefetch
848+ # "registered_users" via prefetch_related. Using .all() hits the
849+ # in-memory prefetch cache (0 DB queries), whereas .filter() would
850+ # bypass the cache and issue a new query every time.
851+ for ru in obj .registered_users .all ():
852+ if organization and ru .organization_id == organization .pk :
853+ reg_user = ru
854+ break
855+ self ._registered_user_cache [obj .pk ] = reg_user
856+ return self ._registered_user_cache [obj .pk ]
857+
858+ def get_is_verified (self , obj ):
859+ reg_user = self ._get_registered_user (obj )
860+ return reg_user .is_verified if reg_user else None
861+
862+ def get_method (self , obj ):
863+ reg_user = self ._get_registered_user (obj )
864+ return reg_user .method if reg_user else None
0 commit comments