From d617ae087fc673e9515af0ec5841decdd5465788 Mon Sep 17 00:00:00 2001 From: ila Date: Sat, 29 Aug 2026 21:19:48 +0800 Subject: [PATCH] feat: add secure portal self-registration API (#80) --- internal/config/config.go | 64 +++++++++- internal/config/config_test.go | 31 +++++ internal/config/file.go | 23 +++- internal/platform/password/password.go | 15 ++- internal/platform/password/password_test.go | 8 ++ portal/auth/registration.go | 134 ++++++++++++++++++++ portal/auth/registration_test.go | 38 ++++++ portal/auth/service.go | 12 +- portal/config/settings.example.yml | 5 + portal/handler/client_ip.go | 62 +++++++++ portal/handler/client_ip_test.go | 34 +++++ portal/handler/mysql_integration_test.go | 86 ++++++++++++- portal/handler/rate_limit.go | 25 ++++ portal/handler/router.go | 116 +++++++++++++---- portal/main.go | 8 +- portal/web/templates/login.html | 1 + portal/web/templates/register.html | 33 +++++ portal/web/web.go | 1 + 18 files changed, 653 insertions(+), 43 deletions(-) create mode 100644 portal/auth/registration.go create mode 100644 portal/auth/registration_test.go create mode 100644 portal/handler/client_ip.go create mode 100644 portal/handler/client_ip_test.go create mode 100644 portal/web/templates/register.html diff --git a/internal/config/config.go b/internal/config/config.go index 5e62363..e915d6c 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -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 diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 2cd5abb..3e6b226 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -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", diff --git a/internal/config/file.go b/internal/config/file.go index a821cfd..f451795 100644 --- a/internal/config/file.go +++ b/internal/config/file.go @@ -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") } diff --git a/internal/platform/password/password.go b/internal/platform/password/password.go index b5b8b4f..8ec6b0d 100644 --- a/internal/platform/password/password.go +++ b/internal/platform/password/password.go @@ -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[:])) } diff --git a/internal/platform/password/password_test.go b/internal/platform/password/password_test.go index 9b21001..3f36908 100644 --- a/internal/platform/password/password_test.go +++ b/internal/platform/password/password_test.go @@ -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") + } } diff --git a/portal/auth/registration.go b/portal/auth/registration.go new file mode 100644 index 0000000..4bbff65 --- /dev/null +++ b/portal/auth/registration.go @@ -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 +} diff --git a/portal/auth/registration_test.go b/portal/auth/registration_test.go new file mode 100644 index 0000000..6c4210b --- /dev/null +++ b/portal/auth/registration_test.go @@ -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") + } +} diff --git a/portal/auth/service.go b/portal/auth/service.go index 711a1c8..682da4a 100644 --- a/portal/auth/service.go +++ b/portal/auth/service.go @@ -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) { diff --git a/portal/config/settings.example.yml b/portal/config/settings.example.yml index b50ce67..d2b8d24 100644 --- a/portal/config/settings.example.yml +++ b/portal/config/settings.example.yml @@ -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 diff --git a/portal/handler/client_ip.go b/portal/handler/client_ip.go new file mode 100644 index 0000000..54920a8 --- /dev/null +++ b/portal/handler/client_ip.go @@ -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 +} diff --git a/portal/handler/client_ip_test.go b/portal/handler/client_ip_test.go new file mode 100644 index 0000000..b3cf7f5 --- /dev/null +++ b/portal/handler/client_ip_test.go @@ -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) + } + }) + } +} diff --git a/portal/handler/mysql_integration_test.go b/portal/handler/mysql_integration_test.go index 32038db..51c0127 100644 --- a/portal/handler/mysql_integration_test.go +++ b/portal/handler/mysql_integration_test.go @@ -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) diff --git a/portal/handler/rate_limit.go b/portal/handler/rate_limit.go index f38baab..fc33dc4 100644 --- a/portal/handler/rate_limit.go +++ b/portal/handler/rate_limit.go @@ -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{ diff --git a/portal/handler/router.go b/portal/handler/router.go index 443aaad..10ca22b 100644 --- a/portal/handler/router.go +++ b/portal/handler/router.go @@ -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) } } diff --git a/portal/main.go b/portal/main.go index 7c7348a..2669786 100644 --- a/portal/main.go +++ b/portal/main.go @@ -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 diff --git a/portal/web/templates/login.html b/portal/web/templates/login.html index 820f04a..ddf979f 100644 --- a/portal/web/templates/login.html +++ b/portal/web/templates/login.html @@ -48,6 +48,7 @@
+ {{if .RegistrationEnabled}}

还没有账号?创建账号

{{end}} diff --git a/portal/web/templates/register.html b/portal/web/templates/register.html new file mode 100644 index 0000000..80c0475 --- /dev/null +++ b/portal/web/templates/register.html @@ -0,0 +1,33 @@ +{{define "register"}} + + + + + + + 创建 Chorus 账号 + + + + + +
+ + +
+ +{{end}} diff --git a/portal/web/web.go b/portal/web/web.go index bc0843b..1d93bec 100644 --- a/portal/web/web.go +++ b/portal/web/web.go @@ -31,6 +31,7 @@ type Page struct { MaxImages int ImageGenerateAvailable bool ImageEditAvailable bool + RegistrationEnabled bool } type Renderer struct{ templates *template.Template }