package api_impl import ( "database/sql" "fmt" "net/http" "atlas9.dev/c/core" "atlas9.dev/c/core/dbi" "atlas9.dev/c/core/iam" "atlas9.dev/c/demo/api" "atlas9.dev/c/demo/lib/access" "atlas9.dev/c/demo/lib/mfa" "atlas9.dev/c/demo/lib/sms" ) type MfaImpl struct { DB *sql.DB MfaEnrollments dbi.Factory[mfa.EnrollmentStore] MfaChallenges dbi.Factory[mfa.ChallengeStore] MfaPolicy dbi.Factory[mfa.PolicyStore] Audit dbi.Factory[iam.AuditStore] SMS *sms.Sender } func (s *MfaImpl) ServeMux(mux *http.ServeMux) { mux.HandleFunc(api.Path_Mfa_GetStatus, s.GetStatus) mux.HandleFunc(api.Path_Mfa_StartEnrollment, s.StartEnrollment) mux.HandleFunc(api.Path_Mfa_ConfirmEnrollment, s.ConfirmEnrollment) mux.HandleFunc(api.Path_Mfa_Disable, s.Disable) mux.HandleFunc(api.Path_Mfa_GetPolicy, s.GetPolicy) mux.HandleFunc(api.Path_Mfa_SetPolicy, s.SetPolicy) } func (s *MfaImpl) GetStatus(w http.ResponseWriter, r *http.Request) { var req api.Mfa_GetStatusReq if read(w, r, &req) { return } ctx := r.Context() userID, err := core.ParseID(iam.GetPrincipal(ctx).Subject) if err != nil { write(ctx, w, err, nil) return } // Whether the user is *required* to use MFA spans their tenant memberships, // a system-plane read; grant it to self for this request. ctx = access.Put(ctx, access.Get(ctx).WithSystem(mfa.Cap_Mfa_Read)) var res api.Mfa_GetStatusRes err = dbi.ReadOnly(ctx, s.DB, func(tx dbi.DBI) error { enrolled, phone, err := enrollmentState(ctx, s.MfaEnrollments(tx), userID) if err != nil { return err } res.Enrolled = enrolled res.Phone = phone required, err := s.MfaPolicy(tx).RequiredForUser(ctx, userID) if err != nil { return err } res.Required = required return nil }) write(ctx, w, err, &res) } func (s *MfaImpl) StartEnrollment(w http.ResponseWriter, r *http.Request) { var req api.Mfa_StartEnrollmentReq if read(w, r, &req) { return } ctx := r.Context() userID, err := core.ParseID(iam.GetPrincipal(ctx).Subject) if err != nil { write(ctx, w, err, nil) return } // Issuing a challenge and texting its code are system-plane effects; add them // to the user's own access for this request. ctx = access.Put(ctx, access.Get(ctx).WithSystem(mfa.Cap_Mfa_Challenge, sms.CapSmsSend)) var challengeID, code string err = dbi.ReadWrite(ctx, s.DB, func(tx dbi.DBI) error { // Store the candidate phone unverified. It becomes a real second factor // only when ConfirmEnrollment flips verified. err := s.MfaEnrollments(tx).Save(ctx, &mfa.Enrollment{UserID: userID, Phone: req.Phone}) if err != nil { return err } id, c, err := s.MfaChallenges(tx).Create(ctx, userID) if err != nil { return err } challengeID, code = id, c return audit(ctx, s.Audit(tx), iam.AuditEntry{Subject: userID, Action: "Mfa_StartEnrollment"}) }) if err != nil { write(ctx, w, err, nil) return } // Text the code after commit; the send is inline (the user is waiting on the // code screen). if err := s.SMS.SendMFACode(ctx, req.Phone, code); err != nil { write(ctx, w, err, nil) return } write(ctx, w, nil, &api.Mfa_StartEnrollmentRes{ChallengeID: challengeID}) } func (s *MfaImpl) ConfirmEnrollment(w http.ResponseWriter, r *http.Request) { var req api.Mfa_ConfirmEnrollmentReq if read(w, r, &req) { return } ctx := r.Context() userID, err := core.ParseID(iam.GetPrincipal(ctx).Subject) if err != nil { write(ctx, w, err, nil) return } ctx = access.Put(ctx, access.Get(ctx).WithSystem(mfa.Cap_Mfa_Challenge)) err = dbi.ReadWrite(ctx, s.DB, func(tx dbi.DBI) error { challengeUser, err := s.MfaChallenges(tx).Verify(ctx, req.ChallengeID, req.Code) if err != nil { return err } // A challenge is bound to the user it was issued for; refuse to confirm // someone else's. if challengeUser != userID { return iam.ErrForbidden } enr, err := s.MfaEnrollments(tx).Get(ctx, userID) if err != nil { return err } enr.Verified = true if err := s.MfaEnrollments(tx).Save(ctx, enr); err != nil { return err } return audit(ctx, s.Audit(tx), iam.AuditEntry{Subject: userID, Action: "Mfa_Enabled"}) }) write(ctx, w, err, &api.Mfa_ConfirmEnrollmentRes{}) } func (s *MfaImpl) Disable(w http.ResponseWriter, r *http.Request) { var req api.Mfa_DisableReq if read(w, r, &req) { return } ctx := r.Context() userID, err := core.ParseID(iam.GetPrincipal(ctx).Subject) if err != nil { write(ctx, w, err, nil) return } err = dbi.ReadWrite(ctx, s.DB, func(tx dbi.DBI) error { if err := s.MfaEnrollments(tx).Delete(ctx, userID); err != nil { return err } return audit(ctx, s.Audit(tx), iam.AuditEntry{Subject: userID, Action: "Mfa_Disabled"}) }) write(ctx, w, err, &api.Mfa_DisableRes{}) } func (s *MfaImpl) GetPolicy(w http.ResponseWriter, r *http.Request) { var req api.Mfa_GetPolicyReq if read(w, r, &req) { return } ctx := r.Context() var res api.Mfa_GetPolicyRes err := dbi.ReadOnly(ctx, s.DB, func(tx dbi.DBI) error { // The policy store enforces Cap_Mfa_PolicyRead scoped to the tenant. required, err := s.MfaPolicy(tx).Get(ctx, req.Tenant) if err != nil { return err } res.Required = required return nil }) write(ctx, w, err, &res) } func (s *MfaImpl) SetPolicy(w http.ResponseWriter, r *http.Request) { var req api.Mfa_SetPolicyReq if read(w, r, &req) { return } ctx := r.Context() err := dbi.ReadWrite(ctx, s.DB, func(tx dbi.DBI) error { // The policy store enforces Cap_Mfa_PolicyWrite scoped to the tenant, held // by owners. if err := s.MfaPolicy(tx).Set(ctx, req.Tenant, req.Required); err != nil { return err } return audit(ctx, s.Audit(tx), iam.AuditEntry{ Tenant: req.Tenant, Action: "Mfa_SetPolicy", Detail: fmt.Sprintf("required=%t", req.Required), }) }) write(ctx, w, err, &api.Mfa_SetPolicyRes{}) }