feat: add secure portal self-registration API (#80)
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -24,11 +25,14 @@ type Config struct {
|
||||
Environment Environment
|
||||
DBDSN string
|
||||
ListenAddress string
|
||||
TrustedProxyCIDRs []netip.Prefix
|
||||
StorageRoot string
|
||||
SessionKey string
|
||||
SessionTTL time.Duration
|
||||
LoginAttempts int
|
||||
LoginWindow time.Duration
|
||||
RegistrationAttempts uint64
|
||||
RegistrationWindow time.Duration
|
||||
MaxPromptBytes int
|
||||
MaxImages int
|
||||
MaxImageBytes int64
|
||||
@@ -47,6 +51,7 @@ type Config struct {
|
||||
ProviderRateLimitCapacity uint64
|
||||
ProviderRateLimitWindow time.Duration
|
||||
TestDisableWorker bool
|
||||
registrationExplicit bool
|
||||
}
|
||||
|
||||
func Load() (Config, error) {
|
||||
@@ -98,6 +103,14 @@ func loadFromLookup(lookup func(string) (string, bool), defaults runtimeSettings
|
||||
StorageRoot: read("CHORUS_STORAGE_ROOT"),
|
||||
SessionKey: read("CHORUS_SESSION_KEY"),
|
||||
}
|
||||
var err error
|
||||
trustedProxyCIDRs := read("CHORUS_TRUSTED_PROXY_CIDRS")
|
||||
if trustedProxyCIDRs == "" {
|
||||
trustedProxyCIDRs = strings.Join(defaults.Server.TrustedProxyCIDRs, ",")
|
||||
}
|
||||
if cfg.TrustedProxyCIDRs, err = parseTrustedProxyCIDRs(trustedProxyCIDRs); err != nil {
|
||||
return Config{}, fmt.Errorf("CHORUS_TRUSTED_PROXY_CIDRS is invalid")
|
||||
}
|
||||
if environment != Production {
|
||||
if cfg.StorageRoot == "" {
|
||||
cfg.StorageRoot = "var/storage"
|
||||
@@ -132,11 +145,16 @@ func loadFromLookup(lookup func(string) (string, bool), defaults runtimeSettings
|
||||
missing = append(missing, name)
|
||||
}
|
||||
}
|
||||
if defaults.Registration.Attempts == nil && read("CHORUS_REGISTRATION_ATTEMPTS") == "" {
|
||||
missing = append(missing, "CHORUS_REGISTRATION_ATTEMPTS")
|
||||
}
|
||||
if defaults.Registration.WindowSeconds == nil && read("CHORUS_REGISTRATION_WINDOW_SECONDS") == "" {
|
||||
missing = append(missing, "CHORUS_REGISTRATION_WINDOW_SECONDS")
|
||||
}
|
||||
}
|
||||
if len(missing) > 0 {
|
||||
return Config{}, fmt.Errorf("missing required configuration: %s", strings.Join(missing, ", "))
|
||||
}
|
||||
var err error
|
||||
if cfg.SessionTTL, err = durationConfig(read("CHORUS_SESSION_TTL_MINUTES"), 480, time.Minute); err != nil {
|
||||
return Config{}, fmt.Errorf("CHORUS_SESSION_TTL_MINUTES is invalid")
|
||||
}
|
||||
@@ -146,6 +164,26 @@ func loadFromLookup(lookup func(string) (string, bool), defaults runtimeSettings
|
||||
if cfg.LoginAttempts, err = intConfig(read("CHORUS_LOGIN_ATTEMPTS"), 5); err != nil {
|
||||
return Config{}, fmt.Errorf("CHORUS_LOGIN_ATTEMPTS is invalid")
|
||||
}
|
||||
registrationAttemptsRaw := read("CHORUS_REGISTRATION_ATTEMPTS")
|
||||
registrationWindowRaw := read("CHORUS_REGISTRATION_WINDOW_SECONDS")
|
||||
cfg.registrationExplicit = registrationAttemptsRaw != "" && registrationWindowRaw != ""
|
||||
registrationAttemptsFallback := 5
|
||||
registrationWindowFallback := 900
|
||||
if defaults.Registration.Attempts != nil {
|
||||
registrationAttemptsFallback = *defaults.Registration.Attempts
|
||||
cfg.registrationExplicit = cfg.registrationExplicit || defaults.Registration.WindowSeconds != nil
|
||||
}
|
||||
if defaults.Registration.WindowSeconds != nil {
|
||||
registrationWindowFallback = *defaults.Registration.WindowSeconds
|
||||
}
|
||||
registrationAttempts, err := intConfig(registrationAttemptsRaw, registrationAttemptsFallback)
|
||||
if err != nil {
|
||||
return Config{}, fmt.Errorf("CHORUS_REGISTRATION_ATTEMPTS is invalid")
|
||||
}
|
||||
cfg.RegistrationAttempts = uint64(registrationAttempts)
|
||||
if cfg.RegistrationWindow, err = durationConfig(registrationWindowRaw, registrationWindowFallback, time.Second); err != nil {
|
||||
return Config{}, fmt.Errorf("CHORUS_REGISTRATION_WINDOW_SECONDS is invalid")
|
||||
}
|
||||
if cfg.MaxPromptBytes, err = intConfig(read("CHORUS_MAX_PROMPT_BYTES"), 8000); err != nil {
|
||||
return Config{}, fmt.Errorf("CHORUS_MAX_PROMPT_BYTES is invalid")
|
||||
}
|
||||
@@ -251,12 +289,34 @@ func (c Config) ValidateProduction() error {
|
||||
if c.Environment != Production {
|
||||
return nil
|
||||
}
|
||||
if c.DBDSN == "" || len(c.SessionKey) < 32 || c.StorageRoot == "" || c.SessionTTL <= 0 || c.MaxImageBytes <= 0 || c.MaxUploadBytes < c.MaxImageBytes || c.WorkerLeaseDuration <= c.ProviderHTTPTimeout+5*time.Second || c.ProviderMaxResponseBytes <= 0 || len(c.ProviderAllowedPorts) == 0 || c.UserRateLimitCapacity == 0 || c.UserRateLimitWindow <= 0 || c.APIKeyRateLimitCapacity == 0 || c.APIKeyRateLimitWindow <= 0 || c.ProviderRateLimitCapacity == 0 || c.ProviderRateLimitWindow <= 0 {
|
||||
if c.DBDSN == "" || len(c.SessionKey) < 32 || c.StorageRoot == "" || c.SessionTTL <= 0 || c.MaxImageBytes <= 0 || c.MaxUploadBytes < c.MaxImageBytes || c.WorkerLeaseDuration <= c.ProviderHTTPTimeout+5*time.Second || c.ProviderMaxResponseBytes <= 0 || len(c.ProviderAllowedPorts) == 0 || c.UserRateLimitCapacity == 0 || c.UserRateLimitWindow <= 0 || c.APIKeyRateLimitCapacity == 0 || c.APIKeyRateLimitWindow <= 0 || c.ProviderRateLimitCapacity == 0 || c.ProviderRateLimitWindow <= 0 || c.RegistrationAttempts == 0 || c.RegistrationWindow <= 0 || !c.registrationExplicit {
|
||||
return errors.New("production configuration is incomplete")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseTrustedProxyCIDRs(raw string) ([]netip.Prefix, error) {
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return nil, nil
|
||||
}
|
||||
parts := strings.Split(raw, ",")
|
||||
prefixes := make([]netip.Prefix, 0, len(parts))
|
||||
seen := make(map[netip.Prefix]struct{}, len(parts))
|
||||
for _, part := range parts {
|
||||
prefix, err := netip.ParsePrefix(strings.TrimSpace(part))
|
||||
if err != nil {
|
||||
return nil, errors.New("invalid trusted proxy CIDR")
|
||||
}
|
||||
prefix = prefix.Masked()
|
||||
if _, exists := seen[prefix]; exists {
|
||||
continue
|
||||
}
|
||||
seen[prefix] = struct{}{}
|
||||
prefixes = append(prefixes, prefix)
|
||||
}
|
||||
return prefixes, nil
|
||||
}
|
||||
|
||||
func ParseProviderAllowedPorts(raw string) ([]uint16, error) {
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return []uint16{80, 443}, nil
|
||||
|
||||
@@ -12,6 +12,10 @@ func TestLoadFileUsesYAMLAndEnvironmentOverrides(t *testing.T) {
|
||||
path := writeSettings(t, `schema_version: 1
|
||||
server:
|
||||
listen_address: 127.0.0.1:9080
|
||||
trusted_proxy_cidrs: [127.0.0.0/8, 10.0.0.0/8]
|
||||
registration:
|
||||
attempts: 5
|
||||
window_seconds: 900
|
||||
provider:
|
||||
http_timeout_seconds: 120
|
||||
max_response_bytes: 1048576
|
||||
@@ -30,6 +34,9 @@ worker:
|
||||
if cfg.ListenAddress != "127.0.0.1:9080" {
|
||||
t.Fatalf("listen address = %q", cfg.ListenAddress)
|
||||
}
|
||||
if cfg.RegistrationAttempts != 5 || cfg.RegistrationWindow != 15*time.Minute || len(cfg.TrustedProxyCIDRs) != 2 {
|
||||
t.Fatalf("registration settings were not loaded: %#v", cfg)
|
||||
}
|
||||
if got := cfg.ProviderAllowedPorts; len(got) != 3 || got[2] != 8080 {
|
||||
t.Fatalf("provider ports = %v", got)
|
||||
}
|
||||
@@ -42,6 +49,9 @@ worker:
|
||||
"CHORUS_PROVIDER_ALLOWED_PORTS": "80,443,8443",
|
||||
"CHORUS_WORKER_LEASE_SECONDS": "90",
|
||||
"CHORUS_WORKER_POLL_MILLISECONDS": "100",
|
||||
"CHORUS_TRUSTED_PROXY_CIDRS": "192.0.2.0/24",
|
||||
"CHORUS_REGISTRATION_ATTEMPTS": "7",
|
||||
"CHORUS_REGISTRATION_WINDOW_SECONDS": "600",
|
||||
}
|
||||
cfg, err = LoadFileFromLookup(path, lookup(overrides))
|
||||
if err != nil {
|
||||
@@ -53,6 +63,9 @@ worker:
|
||||
if cfg.ListenAddress != "127.0.0.1:9081" {
|
||||
t.Fatalf("environment listen address override = %q", cfg.ListenAddress)
|
||||
}
|
||||
if cfg.RegistrationAttempts != 7 || cfg.RegistrationWindow != 10*time.Minute || len(cfg.TrustedProxyCIDRs) != 1 {
|
||||
t.Fatalf("registration overrides were not applied: %#v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadFileRejectsInvalidSettings(t *testing.T) {
|
||||
@@ -83,6 +96,11 @@ worker: {lease_seconds: 60, poll_milliseconds: 250}
|
||||
server: {listen_address: "http://127.0.0.1:8080"}
|
||||
provider: {http_timeout_seconds: 45, max_response_bytes: 1024, allowed_ports: [80]}
|
||||
worker: {lease_seconds: 60, poll_milliseconds: 250}
|
||||
`,
|
||||
"partial registration": `schema_version: 1
|
||||
registration: {attempts: 5}
|
||||
provider: {http_timeout_seconds: 45, max_response_bytes: 1024, allowed_ports: [80]}
|
||||
worker: {lease_seconds: 60, poll_milliseconds: 250}
|
||||
`,
|
||||
}
|
||||
for name, body := range tests {
|
||||
@@ -222,6 +240,7 @@ func completeProductionValues() map[string]string {
|
||||
"CHORUS_USER_RATE_LIMIT_CAPACITY": "60", "CHORUS_USER_RATE_LIMIT_WINDOW_SECONDS": "60",
|
||||
"CHORUS_API_KEY_RATE_LIMIT_CAPACITY": "120", "CHORUS_API_KEY_RATE_LIMIT_WINDOW_SECONDS": "60",
|
||||
"CHORUS_PROVIDER_RATE_LIMIT_CAPACITY": "60", "CHORUS_PROVIDER_RATE_LIMIT_WINDOW_SECONDS": "60",
|
||||
"CHORUS_REGISTRATION_ATTEMPTS": "5", "CHORUS_REGISTRATION_WINDOW_SECONDS": "900",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -230,6 +249,7 @@ func TestProductionRequiresEveryRateLimitValue(t *testing.T) {
|
||||
"CHORUS_USER_RATE_LIMIT_CAPACITY", "CHORUS_USER_RATE_LIMIT_WINDOW_SECONDS",
|
||||
"CHORUS_API_KEY_RATE_LIMIT_CAPACITY", "CHORUS_API_KEY_RATE_LIMIT_WINDOW_SECONDS",
|
||||
"CHORUS_PROVIDER_RATE_LIMIT_CAPACITY", "CHORUS_PROVIDER_RATE_LIMIT_WINDOW_SECONDS",
|
||||
"CHORUS_REGISTRATION_ATTEMPTS", "CHORUS_REGISTRATION_WINDOW_SECONDS",
|
||||
} {
|
||||
values := completeProductionValues()
|
||||
delete(values, name)
|
||||
@@ -244,6 +264,7 @@ func TestLoadRejectsInvalidRateLimits(t *testing.T) {
|
||||
"CHORUS_USER_RATE_LIMIT_CAPACITY", "CHORUS_USER_RATE_LIMIT_WINDOW_SECONDS",
|
||||
"CHORUS_API_KEY_RATE_LIMIT_CAPACITY", "CHORUS_API_KEY_RATE_LIMIT_WINDOW_SECONDS",
|
||||
"CHORUS_PROVIDER_RATE_LIMIT_CAPACITY", "CHORUS_PROVIDER_RATE_LIMIT_WINDOW_SECONDS",
|
||||
"CHORUS_REGISTRATION_ATTEMPTS", "CHORUS_REGISTRATION_WINDOW_SECONDS",
|
||||
} {
|
||||
values := map[string]string{"CHORUS_DSN": "test-dsn", name: "0"}
|
||||
if _, err := LoadFromLookup(lookup(values)); err == nil || !strings.Contains(err.Error(), name) {
|
||||
@@ -252,6 +273,16 @@ func TestLoadRejectsInvalidRateLimits(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadTrustedProxyCIDRs(t *testing.T) {
|
||||
cfg, err := LoadFromLookup(lookup(map[string]string{"CHORUS_DSN": "test-dsn", "CHORUS_TRUSTED_PROXY_CIDRS": "127.0.0.0/8, 10.0.0.0/8,127.0.0.0/8"}))
|
||||
if err != nil || len(cfg.TrustedProxyCIDRs) != 2 {
|
||||
t.Fatalf("trusted proxy CIDRs = %v, %v", cfg.TrustedProxyCIDRs, err)
|
||||
}
|
||||
if _, err := LoadFromLookup(lookup(map[string]string{"CHORUS_DSN": "test-dsn", "CHORUS_TRUSTED_PROXY_CIDRS": "not-a-cidr"})); err == nil {
|
||||
t.Fatal("invalid trusted proxy CIDR accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsUnsafeWorkerLease(t *testing.T) {
|
||||
_, err := LoadFromLookup(lookup(map[string]string{
|
||||
"CHORUS_DSN": "test-dsn", "CHORUS_WORKER_LEASE_SECONDS": "50", "CHORUS_PROVIDER_HTTP_TIMEOUT_SECONDS": "45",
|
||||
|
||||
+18
-5
@@ -15,14 +15,21 @@ import (
|
||||
const maxSettingsBytes = 1 << 20
|
||||
|
||||
type runtimeSettings struct {
|
||||
SchemaVersion int `yaml:"schema_version"`
|
||||
Server serverSettings `yaml:"server"`
|
||||
Provider providerSettings `yaml:"provider"`
|
||||
Worker workerSettings `yaml:"worker"`
|
||||
SchemaVersion int `yaml:"schema_version"`
|
||||
Server serverSettings `yaml:"server"`
|
||||
Registration registrationSettings `yaml:"registration"`
|
||||
Provider providerSettings `yaml:"provider"`
|
||||
Worker workerSettings `yaml:"worker"`
|
||||
}
|
||||
|
||||
type serverSettings struct {
|
||||
ListenAddress *string `yaml:"listen_address"`
|
||||
ListenAddress *string `yaml:"listen_address"`
|
||||
TrustedProxyCIDRs []string `yaml:"trusted_proxy_cidrs"`
|
||||
}
|
||||
|
||||
type registrationSettings struct {
|
||||
Attempts *int `yaml:"attempts"`
|
||||
WindowSeconds *int `yaml:"window_seconds"`
|
||||
}
|
||||
|
||||
type providerSettings struct {
|
||||
@@ -89,6 +96,12 @@ func loadRuntimeSettings(path string) (runtimeSettings, error) {
|
||||
return runtimeSettings{}, err
|
||||
}
|
||||
}
|
||||
if (settings.Registration.Attempts == nil) != (settings.Registration.WindowSeconds == nil) {
|
||||
return runtimeSettings{}, errors.New("portal settings registration attempts and window_seconds must be configured together")
|
||||
}
|
||||
if settings.Registration.Attempts != nil && (*settings.Registration.Attempts <= 0 || *settings.Registration.WindowSeconds <= 0) {
|
||||
return runtimeSettings{}, errors.New("portal settings registration values must be positive")
|
||||
}
|
||||
if settings.Provider.HTTPTimeoutSeconds <= 0 || settings.Provider.MaxResponseBytes <= 0 || settings.Worker.LeaseSeconds <= 0 || settings.Worker.PollMilliseconds <= 0 {
|
||||
return runtimeSettings{}, errors.New("portal settings values must be positive")
|
||||
}
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
package password
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
@@ -16,7 +18,7 @@ func Encode(value string) (string, error) {
|
||||
if len(value) < minLength || len(value) > 1024 {
|
||||
return "", ErrInvalidEncoding
|
||||
}
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(value), bcrypt.DefaultCost)
|
||||
hash, err := bcrypt.GenerateFromPassword(bcryptInput(value), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -27,5 +29,14 @@ func Verify(encoded, value string) bool {
|
||||
if !strings.HasPrefix(encoded, prefix) || len(value) > 1024 {
|
||||
return false
|
||||
}
|
||||
return bcrypt.CompareHashAndPassword([]byte(strings.TrimPrefix(encoded, prefix)), []byte(value)) == nil
|
||||
return bcrypt.CompareHashAndPassword([]byte(strings.TrimPrefix(encoded, prefix)), bcryptInput(value)) == nil
|
||||
}
|
||||
|
||||
func bcryptInput(value string) []byte {
|
||||
raw := []byte(value)
|
||||
if len(raw) <= 72 {
|
||||
return raw
|
||||
}
|
||||
digest := sha256.Sum256(raw)
|
||||
return []byte("sha256:" + hex.EncodeToString(digest[:]))
|
||||
}
|
||||
|
||||
@@ -28,4 +28,12 @@ func TestPasswordLength(t *testing.T) {
|
||||
if Verify("bcrypt:v2:anything", "value") {
|
||||
t.Fatal("unknown version accepted")
|
||||
}
|
||||
long := strings.Repeat("密", 300)
|
||||
encodedLong, err := Encode(long)
|
||||
if err != nil || !Verify(encodedLong, long) || Verify(encodedLong, long+"x") {
|
||||
t.Fatalf("long password round trip failed: %v", err)
|
||||
}
|
||||
if _, err := Encode(strings.Repeat("x", 1025)); err == nil {
|
||||
t.Fatal("password longer than 1024 bytes accepted")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
coreregistration "git.ilapage.cn/OPC/chorus/internal/core/registration"
|
||||
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
|
||||
mysqlDriver "github.com/go-sql-driver/mysql"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidRegistration = errors.New("registration input is invalid")
|
||||
ErrAccountUnavailable = errors.New("registration account is unavailable")
|
||||
ErrRegistrationClosed = errors.New("registration is closed")
|
||||
ErrRegistrationService = errors.New("registration service is unavailable")
|
||||
registrationAccountRule = regexp.MustCompile(`^[a-z0-9][a-z0-9._-]{2,63}$`)
|
||||
)
|
||||
|
||||
type RegistrationInput struct {
|
||||
Username string
|
||||
DisplayName string
|
||||
Password string
|
||||
PasswordConfirmation string
|
||||
RequestID string
|
||||
}
|
||||
|
||||
type portalUserInsert struct {
|
||||
ID uint64 `gorm:"column:id;primaryKey;autoIncrement"`
|
||||
Username string `gorm:"column:username"`
|
||||
Email *string `gorm:"column:email"`
|
||||
PasswordHash string `gorm:"column:password_hash"`
|
||||
DisplayName string `gorm:"column:display_name"`
|
||||
Status string `gorm:"column:status"`
|
||||
CreatedAt time.Time `gorm:"column:created_at"`
|
||||
UpdatedAt time.Time `gorm:"column:updated_at"`
|
||||
}
|
||||
|
||||
func (portalUserInsert) TableName() string { return "users" }
|
||||
|
||||
func (s *Service) RegistrationPolicy(ctx context.Context) (coreregistration.Policy, error) {
|
||||
policy, err := s.registration.Policy(ctx)
|
||||
if err != nil {
|
||||
return coreregistration.Policy{}, ErrRegistrationService
|
||||
}
|
||||
return policy, nil
|
||||
}
|
||||
|
||||
func (s *Service) Register(ctx context.Context, input RegistrationInput) (model.User, error) {
|
||||
input.Username = strings.ToLower(strings.TrimSpace(input.Username))
|
||||
input.DisplayName = strings.TrimSpace(input.DisplayName)
|
||||
input.RequestID = strings.TrimSpace(input.RequestID)
|
||||
if !validRegistrationInput(input) {
|
||||
return model.User{}, s.rejectRegistration(ctx, coreregistration.EventRejected, "invalid_request", input.RequestID, ErrInvalidRegistration)
|
||||
}
|
||||
encoded, err := passwordpkg.Encode(input.Password)
|
||||
if err != nil {
|
||||
return model.User{}, s.rejectRegistration(ctx, coreregistration.EventRejected, "invalid_request", input.RequestID, ErrInvalidRegistration)
|
||||
}
|
||||
|
||||
now := time.Now().UTC()
|
||||
row := portalUserInsert{
|
||||
Username: input.Username, Email: nil, PasswordHash: encoded,
|
||||
DisplayName: input.DisplayName, Status: "active", CreatedAt: now, UpdatedAt: now,
|
||||
}
|
||||
err = s.registration.WithEnabledPolicy(ctx, func(tx *gorm.DB) error {
|
||||
if createErr := tx.Create(&row).Error; createErr != nil {
|
||||
if isDuplicateKey(createErr) {
|
||||
return ErrAccountUnavailable
|
||||
}
|
||||
return ErrRegistrationService
|
||||
}
|
||||
return coreregistration.AppendInTransaction(tx, coreregistration.AuthEvent{
|
||||
EventType: coreregistration.EventSucceeded, Outcome: coreregistration.OutcomeSucceeded,
|
||||
ReasonCode: "created", UserID: &row.ID, RequestID: input.RequestID, CreatedAt: now,
|
||||
})
|
||||
})
|
||||
if err != nil {
|
||||
switch {
|
||||
case errors.Is(err, coreregistration.ErrRegistrationClosed):
|
||||
return model.User{}, s.rejectRegistration(ctx, coreregistration.EventRejected, "registration_closed", input.RequestID, ErrRegistrationClosed)
|
||||
case errors.Is(err, ErrAccountUnavailable):
|
||||
return model.User{}, s.rejectRegistration(ctx, coreregistration.EventRejected, "account_unavailable", input.RequestID, ErrAccountUnavailable)
|
||||
default:
|
||||
return model.User{}, ErrRegistrationService
|
||||
}
|
||||
}
|
||||
return model.User{
|
||||
ID: row.ID, Username: row.Username, PasswordHash: row.PasswordHash,
|
||||
DisplayName: row.DisplayName, Status: row.Status, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) RecordRegistrationRateLimited(ctx context.Context, requestID string) error {
|
||||
return s.registration.Append(ctx, coreregistration.AuthEvent{
|
||||
EventType: coreregistration.EventRateLimited, Outcome: coreregistration.OutcomeRejected,
|
||||
ReasonCode: "registration_rate_limited", RequestID: strings.TrimSpace(requestID),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) RecordRegistrationRejected(ctx context.Context, reasonCode, requestID string) error {
|
||||
return s.registration.Append(ctx, coreregistration.AuthEvent{
|
||||
EventType: coreregistration.EventRejected, Outcome: coreregistration.OutcomeRejected,
|
||||
ReasonCode: strings.TrimSpace(reasonCode), RequestID: strings.TrimSpace(requestID),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) rejectRegistration(ctx context.Context, eventType, reasonCode, requestID string, cause error) error {
|
||||
if err := s.registration.Append(ctx, coreregistration.AuthEvent{
|
||||
EventType: eventType, Outcome: coreregistration.OutcomeRejected,
|
||||
ReasonCode: reasonCode, RequestID: requestID,
|
||||
}); err != nil {
|
||||
return ErrRegistrationService
|
||||
}
|
||||
return cause
|
||||
}
|
||||
|
||||
func validRegistrationInput(input RegistrationInput) bool {
|
||||
return registrationAccountRule.MatchString(input.Username) &&
|
||||
utf8.ValidString(input.DisplayName) && utf8.RuneCountInString(input.DisplayName) >= 1 && utf8.RuneCountInString(input.DisplayName) <= 120 &&
|
||||
utf8.ValidString(input.Password) && len(input.Password) >= 6 && len(input.Password) <= 1024 &&
|
||||
input.Password == input.PasswordConfirmation && input.RequestID != "" && len(input.RequestID) <= 128
|
||||
}
|
||||
|
||||
func isDuplicateKey(err error) bool {
|
||||
var mysqlError *mysqlDriver.MySQLError
|
||||
return errors.As(err, &mysqlError) && mysqlError.Number == 1062
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidRegistrationInput(t *testing.T) {
|
||||
valid := RegistrationInput{Username: "test_user", DisplayName: "测试用户", Password: "test123", PasswordConfirmation: "test123", RequestID: "request-1"}
|
||||
if !validRegistrationInput(valid) {
|
||||
t.Fatal("valid registration input was rejected")
|
||||
}
|
||||
for _, mutate := range []func(*RegistrationInput){
|
||||
func(input *RegistrationInput) { input.Username = "AB" },
|
||||
func(input *RegistrationInput) { input.Username = "email@example.com" },
|
||||
func(input *RegistrationInput) { input.DisplayName = "" },
|
||||
func(input *RegistrationInput) { input.Password = "12345"; input.PasswordConfirmation = "12345" },
|
||||
func(input *RegistrationInput) { input.PasswordConfirmation = "different" },
|
||||
func(input *RegistrationInput) {
|
||||
input.Password = strings.Repeat("x", 1025)
|
||||
input.PasswordConfirmation = input.Password
|
||||
},
|
||||
func(input *RegistrationInput) { input.RequestID = "" },
|
||||
} {
|
||||
input := valid
|
||||
mutate(&input)
|
||||
if validRegistrationInput(input) {
|
||||
t.Fatalf("invalid registration input accepted: %#v", input)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDuplicateKeyClassification(t *testing.T) {
|
||||
if isDuplicateKey(errors.New("Duplicate entry secret-account@example.invalid")) {
|
||||
t.Fatal("unstructured database error was treated as duplicate and could be exposed")
|
||||
}
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"time"
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/registration"
|
||||
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
@@ -17,15 +18,20 @@ var ErrUnavailable = errors.New("authentication service is unavailable")
|
||||
var dummyHash, _ = passwordpkg.Encode("dummy-password-not-used")
|
||||
|
||||
type Service struct {
|
||||
db *gorm.DB
|
||||
limiter *Limiter
|
||||
db *gorm.DB
|
||||
limiter *Limiter
|
||||
registration *registration.Repository
|
||||
}
|
||||
|
||||
func NewService(db *gorm.DB, attempts int, window time.Duration) (*Service, error) {
|
||||
if db == nil || attempts <= 0 || window <= 0 {
|
||||
return nil, errors.New("authentication configuration is invalid")
|
||||
}
|
||||
return &Service{db: db, limiter: NewLimiter(attempts, window)}, nil
|
||||
repository, err := registration.NewRepository(db)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Service{db: db, limiter: NewLimiter(attempts, window), registration: repository}, nil
|
||||
}
|
||||
|
||||
func (s *Service) Login(ctx context.Context, remoteKey, account, password string) (model.User, error) {
|
||||
|
||||
@@ -2,6 +2,11 @@ schema_version: 1
|
||||
|
||||
server:
|
||||
listen_address: 127.0.0.1:8080
|
||||
trusted_proxy_cidrs: []
|
||||
|
||||
registration:
|
||||
attempts: 5
|
||||
window_seconds: 900
|
||||
|
||||
provider:
|
||||
http_timeout_seconds: 45
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type clientIPResolver struct{ trusted []netip.Prefix }
|
||||
|
||||
func newClientIPResolver(trusted []netip.Prefix) *clientIPResolver {
|
||||
return &clientIPResolver{trusted: append([]netip.Prefix(nil), trusted...)}
|
||||
}
|
||||
|
||||
func (r *clientIPResolver) Key(request *http.Request) string {
|
||||
peer, ok := remoteAddress(request.RemoteAddr)
|
||||
if !ok {
|
||||
return strings.TrimSpace(request.RemoteAddr)
|
||||
}
|
||||
if !r.isTrusted(peer) {
|
||||
return peer.Unmap().String()
|
||||
}
|
||||
raw := strings.TrimSpace(request.Header.Get("X-Forwarded-For"))
|
||||
if raw == "" {
|
||||
return peer.Unmap().String()
|
||||
}
|
||||
parts := strings.Split(raw, ",")
|
||||
chain := make([]netip.Addr, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
address, err := netip.ParseAddr(strings.TrimSpace(part))
|
||||
if err != nil {
|
||||
return peer.Unmap().String()
|
||||
}
|
||||
chain = append(chain, address.Unmap())
|
||||
}
|
||||
for index := len(chain) - 1; index >= 0; index-- {
|
||||
if !r.isTrusted(chain[index]) {
|
||||
return chain[index].String()
|
||||
}
|
||||
}
|
||||
return chain[0].String()
|
||||
}
|
||||
|
||||
func (r *clientIPResolver) isTrusted(address netip.Addr) bool {
|
||||
for _, prefix := range r.trusted {
|
||||
if prefix.Contains(address) || prefix.Contains(address.Unmap()) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func remoteAddress(value string) (netip.Addr, bool) {
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(value))
|
||||
if err == nil {
|
||||
address, parseErr := netip.ParseAddr(host)
|
||||
return address, parseErr == nil
|
||||
}
|
||||
address, parseErr := netip.ParseAddr(strings.TrimSpace(value))
|
||||
return address, parseErr == nil
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestClientIPResolver(t *testing.T) {
|
||||
trusted := []netip.Prefix{netip.MustParsePrefix("127.0.0.0/8"), netip.MustParsePrefix("10.0.0.0/8")}
|
||||
resolver := newClientIPResolver(trusted)
|
||||
tests := []struct {
|
||||
name string
|
||||
remoteAddr string
|
||||
forwarded string
|
||||
want string
|
||||
}{
|
||||
{"direct ignores spoofed header", "198.51.100.8:3210", "203.0.113.9", "198.51.100.8"},
|
||||
{"trusted proxy uses client", "127.0.0.1:8080", "203.0.113.9", "203.0.113.9"},
|
||||
{"trusted chain skips trusted hop", "127.0.0.1:8080", "203.0.113.9, 10.0.0.5", "203.0.113.9"},
|
||||
{"malformed chain fails closed", "127.0.0.1:8080", "203.0.113.9, invalid", "127.0.0.1"},
|
||||
{"trusted proxy without header uses peer", "127.0.0.1:8080", "", "127.0.0.1"},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
request := httptest.NewRequest("GET", "http://portal.invalid/", nil)
|
||||
request.RemoteAddr = test.remoteAddr
|
||||
request.Header.Set("X-Forwarded-For", test.forwarded)
|
||||
if got := resolver.Key(request); got != test.want {
|
||||
t.Fatalf("client key = %q, want %q", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/queue"
|
||||
coreregistration "git.ilapage.cn/OPC/chorus/internal/core/registration"
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
|
||||
platformapikey "git.ilapage.cn/OPC/chorus/internal/platform/apikey"
|
||||
@@ -32,6 +33,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
type apiClient struct {
|
||||
@@ -94,7 +96,7 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
t.Skip("CHORUS_TEST_DSN is not set")
|
||||
}
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{})
|
||||
db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -153,13 +155,89 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
sessions, _ := session.New([]byte(strings.Repeat("s", 32)), time.Hour, false)
|
||||
authService, _ := auth.NewService(db, 20, time.Minute)
|
||||
router, err := NewRouter(sessions, authService, generationService, 1<<20, RateLimitConfig{
|
||||
Limiter: ratelimit.New(),
|
||||
UserPolicy: ratelimit.Policy{Capacity: 10_000, Window: time.Minute},
|
||||
APIKeyPolicy: ratelimit.Policy{Capacity: 10_000, Window: time.Minute},
|
||||
Limiter: ratelimit.New(),
|
||||
UserPolicy: ratelimit.Policy{Capacity: 10_000, Window: time.Minute},
|
||||
APIKeyPolicy: ratelimit.Policy{Capacity: 10_000, Window: time.Minute},
|
||||
RegistrationPolicy: ratelimit.Policy{Capacity: 5, Window: 15 * time.Minute},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Model(&coreregistration.Policy{}).Where("id = ?", 1).Updates(map[string]any{"enabled": false, "version": gorm.Expr("version + 1")}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
registrationClient := &apiClient{router: router}
|
||||
registrationClient.start(t)
|
||||
registrationBody, _ := json.Marshal(map[string]string{
|
||||
"username": "self_" + suffix, "display_name": "Self Registered",
|
||||
"password": "test123", "password_confirmation": "test123",
|
||||
})
|
||||
registrationRequestIDs := make([]string, 0, 6)
|
||||
closedRegistration := registrationClient.do(http.MethodPost, "/api/session/register", registrationBody, "application/json")
|
||||
registrationRequestIDs = append(registrationRequestIDs, closedRegistration.Header().Get("X-Request-ID"))
|
||||
if closedRegistration.Code != http.StatusForbidden || !strings.Contains(closedRegistration.Body.String(), `"code":"registration_closed"`) {
|
||||
t.Fatalf("closed registration response=%d %s", closedRegistration.Code, closedRegistration.Body.String())
|
||||
}
|
||||
closedPage := registrationClient.do(http.MethodGet, "/register", nil, "")
|
||||
if closedPage.Code != http.StatusNotFound {
|
||||
t.Fatalf("closed registration page=%d", closedPage.Code)
|
||||
}
|
||||
closedLogin := registrationClient.do(http.MethodGet, "/login", nil, "")
|
||||
if strings.Contains(closedLogin.Body.String(), `href="/register"`) {
|
||||
t.Fatal("closed login page exposed registration link")
|
||||
}
|
||||
if err := db.Model(&coreregistration.Policy{}).Where("id = ?", 1).Updates(map[string]any{"enabled": true, "version": gorm.Expr("version + 1")}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
openPage := registrationClient.do(http.MethodGet, "/register", nil, "")
|
||||
if openPage.Code != http.StatusOK || !strings.Contains(openPage.Body.String(), `id="register-form"`) {
|
||||
t.Fatalf("open registration page=%d %s", openPage.Code, openPage.Body.String())
|
||||
}
|
||||
openLogin := registrationClient.do(http.MethodGet, "/login", nil, "")
|
||||
if !strings.Contains(openLogin.Body.String(), `href="/register"`) {
|
||||
t.Fatal("open login page did not expose registration link")
|
||||
}
|
||||
createdRegistration := registrationClient.do(http.MethodPost, "/api/session/register", registrationBody, "application/json")
|
||||
registrationRequestIDs = append(registrationRequestIDs, createdRegistration.Header().Get("X-Request-ID"))
|
||||
if createdRegistration.Code != http.StatusCreated || !strings.Contains(createdRegistration.Body.String(), `"authenticated":true`) {
|
||||
t.Fatalf("created registration response=%d %s", createdRegistration.Code, createdRegistration.Body.String())
|
||||
}
|
||||
var registeredUser model.User
|
||||
if err := db.Where("username = ?", "self_"+suffix).Take(®isteredUser).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var emailIsNull bool
|
||||
if err := db.Raw("SELECT email IS NULL FROM users WHERE id = ?", registeredUser.ID).Scan(&emailIsNull).Error; err != nil || !emailIsNull {
|
||||
t.Fatalf("self-registered email is NULL = %t, error=%v", emailIsNull, err)
|
||||
}
|
||||
conflictClient := &apiClient{router: router}
|
||||
conflictClient.start(t)
|
||||
conflictRegistration := conflictClient.do(http.MethodPost, "/api/session/register", registrationBody, "application/json")
|
||||
registrationRequestIDs = append(registrationRequestIDs, conflictRegistration.Header().Get("X-Request-ID"))
|
||||
if conflictRegistration.Code != http.StatusConflict || !strings.Contains(conflictRegistration.Body.String(), `"code":"account_unavailable"`) {
|
||||
t.Fatalf("conflict registration response=%d %s", conflictRegistration.Code, conflictRegistration.Body.String())
|
||||
}
|
||||
for attempt := 0; attempt < 2; attempt++ {
|
||||
invalid := conflictClient.do(http.MethodPost, "/api/session/register", []byte(`{}`), "application/json")
|
||||
registrationRequestIDs = append(registrationRequestIDs, invalid.Header().Get("X-Request-ID"))
|
||||
if invalid.Code != http.StatusBadRequest {
|
||||
t.Fatalf("invalid registration response=%d %s", invalid.Code, invalid.Body.String())
|
||||
}
|
||||
}
|
||||
rateLimited := conflictClient.do(http.MethodPost, "/api/session/register", registrationBody, "application/json")
|
||||
registrationRequestIDs = append(registrationRequestIDs, rateLimited.Header().Get("X-Request-ID"))
|
||||
if rateLimited.Code != http.StatusTooManyRequests || rateLimited.Header().Get("Retry-After") == "" || !strings.Contains(rateLimited.Body.String(), `"code":"registration_rate_limited"`) {
|
||||
t.Fatalf("rate-limited registration response=%d %s", rateLimited.Code, rateLimited.Body.String())
|
||||
}
|
||||
defer func() {
|
||||
db.Where("request_id IN ?", registrationRequestIDs).Delete(&coreregistration.AuthEvent{})
|
||||
db.Delete(®isteredUser)
|
||||
db.Model(&coreregistration.Policy{}).Where("id = ?", 1).Updates(map[string]any{"enabled": false, "version": gorm.Expr("version + 1")})
|
||||
}()
|
||||
var registrationEvents int64
|
||||
if err := db.Model(&coreregistration.AuthEvent{}).Where("request_id IN ?", registrationRequestIDs).Count(®istrationEvents).Error; err != nil || registrationEvents != 6 {
|
||||
t.Fatalf("registration event count=%d error=%v", registrationEvents, err)
|
||||
}
|
||||
staticRequest := httptest.NewRequest(http.MethodGet, "/static/app.css", nil)
|
||||
staticResponse := httptest.NewRecorder()
|
||||
router.ServeHTTP(staticResponse, staticRequest)
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
@@ -10,6 +12,29 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func (h *Handler) limitRegistration(c *gin.Context) {
|
||||
decision := h.rateLimiter.Take(h.now(), ratelimit.Bucket{
|
||||
Key: "registration:client:" + h.clientIPs.Key(c.Request), Policy: h.registrationRatePolicy,
|
||||
})
|
||||
if !decision.Allowed {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
if err := h.auth.RecordRegistrationRateLimited(ctx, requestID(c)); err != nil {
|
||||
log.Printf("registration audit write failed request_id=%s: %v", requestID(c), err)
|
||||
}
|
||||
retryAfter := decision.RetryAt.Sub(h.now())
|
||||
seconds := int64((retryAfter + time.Second - 1) / time.Second)
|
||||
if seconds < 1 {
|
||||
seconds = 1
|
||||
}
|
||||
c.Header("Retry-After", strconv.FormatInt(seconds, 10))
|
||||
writeError(c, http.StatusTooManyRequests, "registration_rate_limited", "registration rate limit exceeded")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
|
||||
func (h *Handler) limitPortalSubmission(c *gin.Context) {
|
||||
state := currentSession(c)
|
||||
decision := h.rateLimiter.Take(h.now(), ratelimit.Bucket{
|
||||
|
||||
+92
-24
@@ -7,8 +7,8 @@ import (
|
||||
"html/template"
|
||||
"io"
|
||||
"mime"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"path"
|
||||
"strconv"
|
||||
@@ -26,26 +26,30 @@ import (
|
||||
)
|
||||
|
||||
type Handler struct {
|
||||
sessions *session.Manager
|
||||
auth *auth.Service
|
||||
service *service.Service
|
||||
maxUploadBytes int64
|
||||
renderer *web.Renderer
|
||||
rateLimiter *ratelimit.Limiter
|
||||
userRatePolicy ratelimit.Policy
|
||||
apiKeyRatePolicy ratelimit.Policy
|
||||
now func() time.Time
|
||||
sessions *session.Manager
|
||||
auth *auth.Service
|
||||
service *service.Service
|
||||
maxUploadBytes int64
|
||||
renderer *web.Renderer
|
||||
rateLimiter *ratelimit.Limiter
|
||||
userRatePolicy ratelimit.Policy
|
||||
apiKeyRatePolicy ratelimit.Policy
|
||||
registrationRatePolicy ratelimit.Policy
|
||||
clientIPs *clientIPResolver
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
type RateLimitConfig struct {
|
||||
Limiter *ratelimit.Limiter
|
||||
UserPolicy ratelimit.Policy
|
||||
APIKeyPolicy ratelimit.Policy
|
||||
Now func() time.Time
|
||||
Limiter *ratelimit.Limiter
|
||||
UserPolicy ratelimit.Policy
|
||||
APIKeyPolicy ratelimit.Policy
|
||||
RegistrationPolicy ratelimit.Policy
|
||||
TrustedProxyCIDRs []netip.Prefix
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
func NewRouter(sessions *session.Manager, authService *auth.Service, generationService *service.Service, maxUploadBytes int64, rateLimits RateLimitConfig) (*gin.Engine, error) {
|
||||
if sessions == nil || authService == nil || generationService == nil || maxUploadBytes <= 0 || rateLimits.Limiter == nil || !rateLimits.UserPolicy.Valid() || !rateLimits.APIKeyPolicy.Valid() {
|
||||
if sessions == nil || authService == nil || generationService == nil || maxUploadBytes <= 0 || rateLimits.Limiter == nil || !rateLimits.UserPolicy.Valid() || !rateLimits.APIKeyPolicy.Valid() || !rateLimits.RegistrationPolicy.Valid() {
|
||||
return nil, errors.New("portal handler configuration is invalid")
|
||||
}
|
||||
if rateLimits.Now == nil {
|
||||
@@ -55,6 +59,8 @@ func NewRouter(sessions *session.Manager, authService *auth.Service, generationS
|
||||
sessions: sessions, auth: authService, service: generationService, maxUploadBytes: maxUploadBytes,
|
||||
rateLimiter: rateLimits.Limiter, userRatePolicy: rateLimits.UserPolicy,
|
||||
apiKeyRatePolicy: rateLimits.APIKeyPolicy, now: rateLimits.Now,
|
||||
registrationRatePolicy: rateLimits.RegistrationPolicy,
|
||||
clientIPs: newClientIPResolver(rateLimits.TrustedProxyCIDRs),
|
||||
}
|
||||
renderer, err := web.NewRenderer()
|
||||
if err != nil {
|
||||
@@ -78,6 +84,7 @@ func NewRouter(sessions *session.Manager, authService *auth.Service, generationS
|
||||
static.StaticFS("/", http.FS(staticFS))
|
||||
sessionRoutes := router.Group("", handler.sessionMiddleware)
|
||||
sessionRoutes.GET("/login", handler.loginPage)
|
||||
sessionRoutes.GET("/register", handler.registrationPage)
|
||||
sessionRoutes.GET("/", handler.appPage)
|
||||
sessionRoutes.GET("/api-keys", handler.apiKeysPage)
|
||||
sessionRoutes.GET("/generations/:id", handler.appPage)
|
||||
@@ -85,6 +92,7 @@ func NewRouter(sessions *session.Manager, authService *auth.Service, generationS
|
||||
api := sessionRoutes.Group("/api")
|
||||
api.GET("/session", handler.sessionState)
|
||||
api.POST("/session/login", handler.csrf, handler.login)
|
||||
api.POST("/session/register", handler.csrf, handler.limitRegistration, handler.register)
|
||||
api.POST("/session/logout", handler.requireAuth, handler.csrf, handler.logout)
|
||||
apiKeys := api.Group("/api-keys", func(c *gin.Context) {
|
||||
noStore(c)
|
||||
@@ -165,7 +173,7 @@ func (h *Handler) login(c *gin.Context) {
|
||||
writeError(c, 400, "invalid_request", "account and password are required")
|
||||
return
|
||||
}
|
||||
user, err := h.auth.Login(c.Request.Context(), remoteKey(c.Request), input.Account, input.Password)
|
||||
user, err := h.auth.Login(c.Request.Context(), h.clientIPs.Key(c.Request), input.Account, input.Password)
|
||||
if err != nil {
|
||||
if !errors.Is(err, auth.ErrInvalidCredentials) {
|
||||
writeError(c, http.StatusInternalServerError, "internal_error", "request could not be completed")
|
||||
@@ -181,6 +189,47 @@ func (h *Handler) login(c *gin.Context) {
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"authenticated": true, "csrf_token": state.CSRFToken, "user": gin.H{"id": user.ID, "display_name": user.DisplayName}})
|
||||
}
|
||||
|
||||
func (h *Handler) register(c *gin.Context) {
|
||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 64<<10)
|
||||
var input struct {
|
||||
Username string `json:"username"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Password string `json:"password"`
|
||||
PasswordConfirmation string `json:"password_confirmation"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&input); err != nil {
|
||||
if auditErr := h.auth.RecordRegistrationRejected(c.Request.Context(), "invalid_request", requestID(c)); auditErr != nil {
|
||||
writeError(c, http.StatusServiceUnavailable, "registration_unavailable", "registration is temporarily unavailable")
|
||||
return
|
||||
}
|
||||
writeError(c, http.StatusBadRequest, "invalid_request", "registration information is invalid")
|
||||
return
|
||||
}
|
||||
user, err := h.auth.Register(c.Request.Context(), auth.RegistrationInput{
|
||||
Username: input.Username, DisplayName: input.DisplayName, Password: input.Password,
|
||||
PasswordConfirmation: input.PasswordConfirmation, RequestID: requestID(c),
|
||||
})
|
||||
if err != nil {
|
||||
switch {
|
||||
case errors.Is(err, auth.ErrInvalidRegistration):
|
||||
writeError(c, http.StatusBadRequest, "invalid_request", "registration information is invalid")
|
||||
case errors.Is(err, auth.ErrRegistrationClosed):
|
||||
writeError(c, http.StatusForbidden, "registration_closed", "registration is not available")
|
||||
case errors.Is(err, auth.ErrAccountUnavailable):
|
||||
writeError(c, http.StatusConflict, "account_unavailable", "account is unavailable")
|
||||
default:
|
||||
writeError(c, http.StatusServiceUnavailable, "registration_unavailable", "registration is temporarily unavailable")
|
||||
}
|
||||
return
|
||||
}
|
||||
state, err := h.sessions.Authenticate(c.Writer, c.Request, user.ID)
|
||||
if err != nil {
|
||||
writeError(c, http.StatusInternalServerError, "internal_error", "request could not be completed")
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, gin.H{"authenticated": true, "csrf_token": state.CSRFToken, "user": gin.H{"id": user.ID, "display_name": user.DisplayName}})
|
||||
}
|
||||
func (h *Handler) logout(c *gin.Context) {
|
||||
state, err := h.sessions.Logout(c.Writer, c.Request)
|
||||
if err != nil {
|
||||
@@ -480,13 +529,6 @@ func uintParam(c *gin.Context, name string) (uint64, bool) {
|
||||
}
|
||||
return value, true
|
||||
}
|
||||
func remoteKey(request *http.Request) string {
|
||||
host, _, err := net.SplitHostPort(request.RemoteAddr)
|
||||
if err == nil {
|
||||
return host
|
||||
}
|
||||
return strings.TrimSpace(request.RemoteAddr)
|
||||
}
|
||||
func isHTMX(c *gin.Context) bool { return strings.EqualFold(c.GetHeader("HX-Request"), "true") }
|
||||
|
||||
func safeName(value string) string {
|
||||
@@ -514,7 +556,33 @@ func (h *Handler) loginPage(c *gin.Context) {
|
||||
}
|
||||
c.Header("Cache-Control", "no-store")
|
||||
c.Header("Content-Type", "text/html; charset=utf-8")
|
||||
if err := h.renderer.Render(c.Writer, "login", web.Page{Title: "登录", CSRFToken: state.CSRFToken, ReturnTo: returnTo}); err != nil {
|
||||
registrationEnabled := false
|
||||
if policy, err := h.auth.RegistrationPolicy(c.Request.Context()); err == nil {
|
||||
registrationEnabled = policy.Enabled
|
||||
}
|
||||
if err := h.renderer.Render(c.Writer, "login", web.Page{Title: "登录", CSRFToken: state.CSRFToken, ReturnTo: returnTo, RegistrationEnabled: registrationEnabled}); err != nil {
|
||||
c.Status(http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) registrationPage(c *gin.Context) {
|
||||
state := currentSession(c)
|
||||
if state.UserID != 0 {
|
||||
c.Redirect(http.StatusSeeOther, "/")
|
||||
return
|
||||
}
|
||||
policy, err := h.auth.RegistrationPolicy(c.Request.Context())
|
||||
if err != nil {
|
||||
c.Status(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
if !policy.Enabled {
|
||||
c.Status(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
c.Header("Cache-Control", "no-store")
|
||||
c.Header("Content-Type", "text/html; charset=utf-8")
|
||||
if err := h.renderer.Render(c.Writer, "register", web.Page{Title: "创建账号", CSRFToken: state.CSRFToken, RegistrationEnabled: true}); err != nil {
|
||||
c.Status(http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
+5
-3
@@ -127,9 +127,11 @@ func run(configPath string) error {
|
||||
return err
|
||||
}
|
||||
router, err := handler.NewRouter(sessions, authService, generationService, cfg.MaxUploadBytes, handler.RateLimitConfig{
|
||||
Limiter: rateLimiter,
|
||||
UserPolicy: ratelimit.Policy{Capacity: cfg.UserRateLimitCapacity, Window: cfg.UserRateLimitWindow},
|
||||
APIKeyPolicy: ratelimit.Policy{Capacity: cfg.APIKeyRateLimitCapacity, Window: cfg.APIKeyRateLimitWindow},
|
||||
Limiter: rateLimiter,
|
||||
UserPolicy: ratelimit.Policy{Capacity: cfg.UserRateLimitCapacity, Window: cfg.UserRateLimitWindow},
|
||||
APIKeyPolicy: ratelimit.Policy{Capacity: cfg.APIKeyRateLimitCapacity, Window: cfg.APIKeyRateLimitWindow},
|
||||
RegistrationPolicy: ratelimit.Policy{Capacity: cfg.RegistrationAttempts, Window: cfg.RegistrationWindow},
|
||||
TrustedProxyCIDRs: cfg.TrustedProxyCIDRs,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
@@ -48,6 +48,7 @@
|
||||
<div class="form-field">
|
||||
<button class="button button-primary button-block" id="login-button" type="submit"><span>登录</span></button>
|
||||
</div>
|
||||
{{if .RegistrationEnabled}}<p class="auth-switch">还没有账号?<a href="/register">创建账号</a></p>{{end}}
|
||||
</form>
|
||||
</section>
|
||||
</main>
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
{{define "register"}}<!doctype html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1">
|
||||
<meta name="color-scheme" content="light">
|
||||
<meta name="csrf-token" content="{{.CSRFToken}}">
|
||||
<title>创建 Chorus 账号</title>
|
||||
<link rel="stylesheet" href="/static/app.css">
|
||||
<script defer src="/static/vendor/icons.min.js"></script>
|
||||
<script defer src="/static/app.js"></script>
|
||||
</head>
|
||||
<body data-page="register">
|
||||
<main class="login-view">
|
||||
<section class="login-context" aria-labelledby="register-context-title">
|
||||
<a class="brand" href="/login"><span class="brand-mark" aria-hidden="true">C</span><span>CHORUS</span></a>
|
||||
<div class="login-copy"><h1 id="register-context-title">创建终端用户账号</h1><p>注册成功后将直接进入现有工作台。</p></div>
|
||||
</section>
|
||||
<section class="login-form-wrap" aria-labelledby="register-title">
|
||||
<form class="login-form" id="register-form" novalidate>
|
||||
<h2 id="register-title">创建账号</h2>
|
||||
<div class="form-field"><label class="field-label" for="username">账号</label><input class="input" id="username" name="username" autocomplete="username" minlength="3" maxlength="64" required></div>
|
||||
<div class="form-field"><label class="field-label" for="display-name">昵称</label><input class="input" id="display-name" name="display_name" autocomplete="nickname" maxlength="120" required></div>
|
||||
<div class="form-field"><label class="field-label" for="register-password">密码</label><input class="input" id="register-password" name="password" type="password" autocomplete="new-password" minlength="6" maxlength="1024" required><p class="field-help">密码至少 6 位</p></div>
|
||||
<div class="form-field"><label class="field-label" for="password-confirmation">确认密码</label><input class="input" id="password-confirmation" name="password_confirmation" type="password" autocomplete="new-password" minlength="6" maxlength="1024" required></div>
|
||||
<p class="field-error hidden" id="register-error" role="alert" aria-live="polite"></p>
|
||||
<button class="button button-primary button-block" id="register-button" type="submit"><span>创建并登录</span></button>
|
||||
<a class="button button-ghost button-block" href="/login">返回登录</a>
|
||||
</form>
|
||||
</section>
|
||||
</main>
|
||||
</body>
|
||||
</html>{{end}}
|
||||
@@ -31,6 +31,7 @@ type Page struct {
|
||||
MaxImages int
|
||||
ImageGenerateAvailable bool
|
||||
ImageEditAvailable bool
|
||||
RegistrationEnabled bool
|
||||
}
|
||||
|
||||
type Renderer struct{ templates *template.Template }
|
||||
|
||||
Reference in New Issue
Block a user