Files
NianAIGC/backend/internal/localstore/accounts.go

293 lines
9.5 KiB
Go

package localstore
import (
"context"
"crypto/subtle"
"encoding/hex"
"fmt"
"sort"
"strings"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/administration"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity"
"golang.org/x/crypto/scrypt"
)
func (s *Store) ListAccounts(_ context.Context, f administration.AccountFilters) ([]administration.Account, error) {
s.mu.RLock()
defer s.mu.RUnlock()
out := []administration.Account{}
for _, a := range s.accounts {
if f.OrganizationID != "" && a.OrganizationID != f.OrganizationID || f.Role != "" && a.Role != f.Role || !f.IncludeDisabled && a.Status == administration.StatusDisabled {
continue
}
out = append(out, a)
}
sort.Slice(out, func(i, j int) bool { return out[i].CreatedAt.Before(out[j].CreatedAt) })
return out, nil
}
func (s *Store) GetAccount(_ context.Context, id string) (administration.Account, bool, error) {
s.mu.RLock()
defer s.mu.RUnlock()
a, ok := s.accounts[id]
return a, ok, nil
}
func (s *Store) CreateAccount(_ context.Context, a administration.Account) (administration.Account, error) {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.accounts[a.ID]; ok {
return administration.Account{}, conflict("账号已存在。")
}
for _, v := range s.accounts {
if administration.NormalizePhone(v.Phone) == administration.NormalizePhone(a.Phone) {
return administration.Account{}, conflict("手机号已存在。")
}
}
s.accounts[a.ID] = a
return a, nil
}
func (s *Store) UpdateAccount(_ context.Context, a administration.Account) (administration.Account, error) {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.accounts[a.ID]; !ok {
return administration.Account{}, notFound("账号不存在。")
}
for id, v := range s.accounts {
if id != a.ID && administration.NormalizePhone(v.Phone) == administration.NormalizePhone(a.Phone) {
return administration.Account{}, conflict("手机号已存在。")
}
}
s.accounts[a.ID] = a
return a, nil
}
func (s *Store) DeleteAccount(_ context.Context, id, archive string) error {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.accounts[id]; !ok {
return notFound("账号不存在。")
}
if archive != "" {
for key, a := range s.assets {
if a.OwnerID == id {
a.OwnerID = archive
s.assets[key] = a
}
}
for key, j := range s.jobs {
if j.OwnerID == id {
j.OwnerID = archive
s.jobs[key] = j
}
}
for key, t := range s.templates {
if t.OwnerID == id {
t.OwnerID = archive
s.templates[key] = t
}
}
for key, event := range s.usageEvents {
if event.OwnerID == id {
event.OwnerID = archive
s.usageEvents[key] = event
}
}
}
delete(s.accounts, id)
return nil
}
func (s *Store) ListOrganizations(_ context.Context, include bool) ([]administration.Organization, error) {
s.mu.RLock()
defer s.mu.RUnlock()
out := []administration.Organization{}
for _, o := range s.organizations {
if include || o.Status != administration.StatusDisabled {
out = append(out, o)
}
}
sort.Slice(out, func(i, j int) bool { return out[i].CreatedAt.Before(out[j].CreatedAt) })
return out, nil
}
func (s *Store) GetOrganization(_ context.Context, id string) (administration.Organization, bool, error) {
s.mu.RLock()
defer s.mu.RUnlock()
o, ok := s.organizations[id]
return o, ok, nil
}
func (s *Store) CreateOrganization(_ context.Context, o administration.Organization) (administration.Organization, error) {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.organizations[o.ID]; ok {
return administration.Organization{}, conflict("组织已存在。")
}
s.organizations[o.ID] = o
return o, nil
}
func (s *Store) UpdateOrganization(_ context.Context, o administration.Organization) (administration.Organization, error) {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.organizations[o.ID]; !ok {
return administration.Organization{}, notFound("组织不存在。")
}
s.organizations[o.ID] = o
return o, nil
}
func (s *Store) DeleteOrganization(_ context.Context, id string) error {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.organizations[id]; !ok {
return notFound("组织不存在。")
}
for _, a := range s.accounts {
if a.OrganizationID == id {
return conflict("组织仍有成员。")
}
}
delete(s.organizations, id)
delete(s.wallets, id)
return nil
}
func (s *Store) CountOrganizationMembers(_ context.Context, id string) (int, error) {
s.mu.RLock()
defer s.mu.RUnlock()
n := 0
for _, a := range s.accounts {
if a.OrganizationID == id {
n++
}
}
return n, nil
}
func conflict(message string) error {
return &administration.Error{Kind: administration.ErrorConflict, Message: message}
}
func notFound(message string) error {
return &administration.Error{Kind: administration.ErrorNotFound, Message: message}
}
func (s *Store) FindAuthorizationSnapshot(_ context.Context, key string) (identity.AuthorizationSnapshot, bool, error) {
s.mu.RLock()
defer s.mu.RUnlock()
for _, a := range s.accounts {
if a.ID == key || a.LegacySubject == key {
return s.snapshotLocked(a), true, nil
}
}
return identity.AuthorizationSnapshot{}, false, nil
}
func (s *Store) snapshotLocked(a administration.Account) identity.AuthorizationSnapshot {
out := identity.AuthorizationSnapshot{Account: identity.AccountSnapshot{ID: a.ID, Phone: a.Phone, DisplayName: a.DisplayName, Role: string(a.Role), OrganizationID: a.OrganizationID, Status: string(a.Status), SessionVersion: a.SessionVersion}}
if o, ok := s.organizations[a.OrganizationID]; ok {
out.Organization = &identity.OrganizationSnapshot{ID: o.ID, Name: o.Name, Status: string(o.Status)}
}
return out
}
const maxFailures = 5
const lockTime = 15 * time.Minute
func (s *Store) AttemptPasswordLogin(_ context.Context, phone, password string, now time.Time) (identity.LoginAccount, error) {
s.mu.Lock()
defer s.mu.Unlock()
var a administration.Account
found := false
phone = administration.NormalizePhone(phone)
for _, v := range s.accounts {
if administration.NormalizePhone(v.Phone) == phone {
a = v
found = true
break
}
}
if !found {
return identity.LoginAccount{}, identity.NewPasswordLoginError(identity.LoginFailureInvalidCredentials)
}
if a.Status != administration.StatusActive {
return identity.LoginAccount{}, identity.NewPasswordLoginError(identity.LoginFailureAccountDisabled)
}
if a.LockedUntil != nil && a.LockedUntil.After(now) {
return identity.LoginAccount{}, identity.NewPasswordLoginError(identity.LoginFailureAccountLocked)
}
if !verifyPassword(password, a.PasswordHash, a.PasswordSalt) {
a.FailedLoginCount++
reason := identity.LoginFailureInvalidCredentials
if a.FailedLoginCount >= maxFailures {
a.FailedLoginCount = 0
locked := now.Add(lockTime)
a.LockedUntil = &locked
reason = identity.LoginFailureAccountLocked
}
a.UpdatedAt = now
s.accounts[a.ID] = a
return identity.LoginAccount{}, identity.NewPasswordLoginError(reason)
}
a.FailedLoginCount = 0
a.LockedUntil = nil
loginAt := now
a.LastLoginAt = &loginAt
a.UpdatedAt = now
s.accounts[a.ID] = a
snapshot := s.snapshotLocked(a)
return identity.LoginAccount{Account: snapshot.Account, Organization: snapshot.Organization}, nil
}
func (s *Store) ChangeOwnPassword(_ context.Context, id, current, next string, now time.Time) (identity.AuthorizationSnapshot, error) {
s.mu.Lock()
defer s.mu.Unlock()
a, ok := s.accounts[id]
if !ok || a.Status != administration.StatusActive {
return identity.AuthorizationSnapshot{}, &identity.PasswordChangeError{Reason: identity.PasswordChangeNotFound}
}
if !verifyPassword(current, a.PasswordHash, a.PasswordSalt) {
return identity.AuthorizationSnapshot{}, &identity.PasswordChangeError{Reason: identity.PasswordChangeCurrentIncorrect}
}
hashed, err := administration.HashPassword(next)
if err != nil {
return identity.AuthorizationSnapshot{}, fmt.Errorf("hash new password: %w", err)
}
a.PasswordHash = hashed.Hash
a.PasswordSalt = hashed.Salt
a.SessionVersion++
a.UpdatedAt = now
s.accounts[id] = a
return s.snapshotLocked(a), nil
}
func verifyPassword(password, hash, salt string) bool {
if saltBytes, err := hex.DecodeString(salt); err != nil || len(saltBytes) == 0 {
return false
}
expected, err := hex.DecodeString(hash)
if err != nil || len(expected) == 0 {
return false
}
derived, err := scrypt.Key([]byte(strings.TrimSpace(password)), []byte(salt), 16384, 8, 1, len(expected))
return err == nil && subtle.ConstantTimeCompare(derived, expected) == 1
}
func (s *Store) BillingOrganizations(_ context.Context) ([]billing.Organization, error) {
s.mu.RLock()
defer s.mu.RUnlock()
out := []billing.Organization{}
for _, o := range s.organizations {
out = append(out, billing.Organization{ID: o.ID, Name: o.Name, Status: string(o.Status), ArchiveOwnerID: o.ArchiveOwnerID, CreatedAt: o.CreatedAt.Format(time.RFC3339Nano), UpdatedAt: o.UpdatedAt.Format(time.RFC3339Nano)})
}
sort.Slice(out, func(i, j int) bool { return out[i].CreatedAt < out[j].CreatedAt })
return out, nil
}
func (s *Store) BillingMembers(_ context.Context) ([]billing.Member, error) {
s.mu.RLock()
defer s.mu.RUnlock()
out := []billing.Member{}
for _, a := range s.accounts {
out = append(out, billing.Member{ID: a.ID, DisplayName: a.DisplayName, Phone: a.Phone, Role: string(a.Role), OrganizationID: a.OrganizationID, Status: string(a.Status)})
}
sort.Slice(out, func(i, j int) bool { return out[i].ID < out[j].ID })
return out, nil
}
func (s *Store) BillingOrganizationExists(_ context.Context, id string) (bool, error) {
s.mu.RLock()
defer s.mu.RUnlock()
_, ok := s.organizations[id]
return ok, nil
}