chore: initial commit of ai site platform
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
297
platform/internal/invitestore/store.go
Normal file
297
platform/internal/invitestore/store.go
Normal file
@@ -0,0 +1,297 @@
|
||||
package invitestore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"aijianzhan/platform/internal/authx"
|
||||
)
|
||||
|
||||
type Invite struct {
|
||||
InviteID int64 `json:"invite_id"`
|
||||
TenantID int64 `json:"tenant_id"`
|
||||
Code string `json:"code"`
|
||||
Role string `json:"role"`
|
||||
OrgUnitID int64 `json:"org_unit_id,omitempty"`
|
||||
CreatedBy int64 `json:"created_by"`
|
||||
MaxUses int `json:"max_uses"`
|
||||
UsedCount int `json:"used_count"`
|
||||
Status string `json:"status"`
|
||||
ExpiresAt *time.Time `json:"expires_at,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type CreateInput struct {
|
||||
Role string
|
||||
OrgUnitID int64
|
||||
MaxUses int
|
||||
ExpiresIn time.Duration // 0 = 不过期
|
||||
}
|
||||
|
||||
type Store interface {
|
||||
Create(ctx context.Context, tenantID, createdBy int64, in CreateInput) (*Invite, error)
|
||||
List(ctx context.Context, tenantID int64) ([]Invite, error)
|
||||
GetByCode(ctx context.Context, code string) (*Invite, error)
|
||||
Consume(ctx context.Context, inviteID int64) error
|
||||
Revoke(ctx context.Context, tenantID, inviteID int64) error
|
||||
}
|
||||
|
||||
type MemoryStore struct {
|
||||
mu sync.Mutex
|
||||
byID map[int64]*Invite
|
||||
code map[string]int64
|
||||
seq int64
|
||||
}
|
||||
|
||||
func NewMemoryStore() *MemoryStore {
|
||||
return &MemoryStore{byID: map[int64]*Invite{}, code: map[string]int64{}}
|
||||
}
|
||||
|
||||
func (s *MemoryStore) Create(_ context.Context, tenantID, createdBy int64, in CreateInput) (*Invite, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
code, err := randomCode()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.seq++
|
||||
inv := &Invite{
|
||||
InviteID: s.seq,
|
||||
TenantID: tenantID,
|
||||
Code: code,
|
||||
Role: normalizeRole(in.Role),
|
||||
OrgUnitID: in.OrgUnitID,
|
||||
CreatedBy: createdBy,
|
||||
MaxUses: maxUses(in.MaxUses),
|
||||
Status: "active",
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}
|
||||
if in.ExpiresIn > 0 {
|
||||
t := time.Now().UTC().Add(in.ExpiresIn)
|
||||
inv.ExpiresAt = &t
|
||||
}
|
||||
s.byID[inv.InviteID] = inv
|
||||
s.code[inv.Code] = inv.InviteID
|
||||
cp := *inv
|
||||
return &cp, nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) List(_ context.Context, tenantID int64) ([]Invite, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
out := make([]Invite, 0)
|
||||
for _, inv := range s.byID {
|
||||
if inv.TenantID == tenantID {
|
||||
out = append(out, *inv)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) GetByCode(_ context.Context, code string) (*Invite, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
id, ok := s.code[strings.TrimSpace(code)]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invite not found")
|
||||
}
|
||||
inv := s.byID[id]
|
||||
cp := *inv
|
||||
return &cp, nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) Consume(_ context.Context, inviteID int64) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
inv, ok := s.byID[inviteID]
|
||||
if !ok {
|
||||
return fmt.Errorf("invite not found")
|
||||
}
|
||||
if err := usable(inv); err != nil {
|
||||
return err
|
||||
}
|
||||
inv.UsedCount++
|
||||
if inv.UsedCount >= inv.MaxUses {
|
||||
inv.Status = "exhausted"
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) Revoke(_ context.Context, tenantID, inviteID int64) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
inv, ok := s.byID[inviteID]
|
||||
if !ok || inv.TenantID != tenantID {
|
||||
return fmt.Errorf("invite not found")
|
||||
}
|
||||
inv.Status = "revoked"
|
||||
return nil
|
||||
}
|
||||
|
||||
type PostgresStore struct {
|
||||
DB *sql.DB
|
||||
}
|
||||
|
||||
func NewPostgresStore(db *sql.DB) *PostgresStore {
|
||||
return &PostgresStore{DB: db}
|
||||
}
|
||||
|
||||
func (s *PostgresStore) Create(ctx context.Context, tenantID, createdBy int64, in CreateInput) (*Invite, error) {
|
||||
code, err := randomCode()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var expires any
|
||||
var expPtr *time.Time
|
||||
if in.ExpiresIn > 0 {
|
||||
t := time.Now().UTC().Add(in.ExpiresIn)
|
||||
expPtr = &t
|
||||
expires = t
|
||||
}
|
||||
inv := &Invite{
|
||||
TenantID: tenantID,
|
||||
Code: code,
|
||||
Role: normalizeRole(in.Role),
|
||||
OrgUnitID: in.OrgUnitID,
|
||||
CreatedBy: createdBy,
|
||||
MaxUses: maxUses(in.MaxUses),
|
||||
Status: "active",
|
||||
ExpiresAt: expPtr,
|
||||
}
|
||||
err = s.DB.QueryRowContext(ctx, `
|
||||
INSERT INTO platform_meta.tenant_invites(tenant_id, code, role, created_by, max_uses, expires_at, org_unit_id)
|
||||
VALUES($1,$2,$3,$4,$5,$6,$7)
|
||||
RETURNING invite_id, created_at`,
|
||||
tenantID, code, inv.Role, createdBy, inv.MaxUses, expires, in.OrgUnitID,
|
||||
).Scan(&inv.InviteID, &inv.CreatedAt)
|
||||
return inv, err
|
||||
}
|
||||
|
||||
func (s *PostgresStore) List(ctx context.Context, tenantID int64) ([]Invite, error) {
|
||||
rows, err := s.DB.QueryContext(ctx, `
|
||||
SELECT invite_id, tenant_id, code, role, created_by, max_uses, used_count, status, expires_at, created_at, COALESCE(org_unit_id,0)
|
||||
FROM platform_meta.tenant_invites WHERE tenant_id=$1 ORDER BY invite_id DESC`, tenantID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := make([]Invite, 0)
|
||||
for rows.Next() {
|
||||
var inv Invite
|
||||
var exp sql.NullTime
|
||||
if err := rows.Scan(&inv.InviteID, &inv.TenantID, &inv.Code, &inv.Role, &inv.CreatedBy,
|
||||
&inv.MaxUses, &inv.UsedCount, &inv.Status, &exp, &inv.CreatedAt, &inv.OrgUnitID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if exp.Valid {
|
||||
t := exp.Time.UTC()
|
||||
inv.ExpiresAt = &t
|
||||
}
|
||||
out = append(out, inv)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (s *PostgresStore) GetByCode(ctx context.Context, code string) (*Invite, error) {
|
||||
var inv Invite
|
||||
var exp sql.NullTime
|
||||
err := s.DB.QueryRowContext(ctx, `
|
||||
SELECT invite_id, tenant_id, code, role, created_by, max_uses, used_count, status, expires_at, created_at, COALESCE(org_unit_id,0)
|
||||
FROM platform_meta.tenant_invites WHERE code=$1`, strings.TrimSpace(code),
|
||||
).Scan(&inv.InviteID, &inv.TenantID, &inv.Code, &inv.Role, &inv.CreatedBy,
|
||||
&inv.MaxUses, &inv.UsedCount, &inv.Status, &exp, &inv.CreatedAt, &inv.OrgUnitID)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("invite not found")
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if exp.Valid {
|
||||
t := exp.Time.UTC()
|
||||
inv.ExpiresAt = &t
|
||||
}
|
||||
return &inv, nil
|
||||
}
|
||||
|
||||
func (s *PostgresStore) Consume(ctx context.Context, inviteID int64) error {
|
||||
res, err := s.DB.ExecContext(ctx, `
|
||||
UPDATE platform_meta.tenant_invites
|
||||
SET used_count = used_count + 1,
|
||||
status = CASE WHEN used_count + 1 >= max_uses THEN 'exhausted' ELSE status END
|
||||
WHERE invite_id=$1 AND status='active'
|
||||
AND (expires_at IS NULL OR expires_at > now())
|
||||
AND used_count < max_uses`, inviteID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return fmt.Errorf("invite not usable")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *PostgresStore) Revoke(ctx context.Context, tenantID, inviteID int64) error {
|
||||
res, err := s.DB.ExecContext(ctx, `
|
||||
UPDATE platform_meta.tenant_invites SET status='revoked'
|
||||
WHERE invite_id=$1 AND tenant_id=$2`, inviteID, tenantID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return fmt.Errorf("invite not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func randomCode() (string, error) {
|
||||
b := make([]byte, 12)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
func normalizeRole(role string) string {
|
||||
n := authx.NormalizeRole(role)
|
||||
if authx.ValidPlatformRole(n) {
|
||||
return n
|
||||
}
|
||||
return authx.Role编辑
|
||||
}
|
||||
|
||||
func maxUses(n int) int {
|
||||
if n <= 0 {
|
||||
return 1
|
||||
}
|
||||
if n > 1000 {
|
||||
return 1000
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func usable(inv *Invite) error {
|
||||
if inv.Status != "active" {
|
||||
return fmt.Errorf("invite not usable")
|
||||
}
|
||||
if inv.ExpiresAt != nil && time.Now().UTC().After(*inv.ExpiresAt) {
|
||||
return fmt.Errorf("invite expired")
|
||||
}
|
||||
if inv.UsedCount >= inv.MaxUses {
|
||||
return fmt.Errorf("invite exhausted")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateUsable 供业务层在 Consume 前检查。
|
||||
func ValidateUsable(inv *Invite) error {
|
||||
return usable(inv)
|
||||
}
|
||||
Reference in New Issue
Block a user