feat: add secure portal self-registration API (#80)

This commit is contained in:
ila
2026-08-29 21:19:48 +08:00
parent a079a1d7bb
commit d617ae087f
18 changed files with 653 additions and 43 deletions
+62 -2
View File
@@ -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
+31
View File
@@ -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
View File
@@ -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")
}
+13 -2
View File
@@ -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")
}
}
+134
View File
@@ -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
}
+38
View File
@@ -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")
}
}
+9 -3
View File
@@ -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) {
+5
View File
@@ -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
+62
View File
@@ -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
}
+34
View File
@@ -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)
}
})
}
}
+82 -4
View File
@@ -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(&registeredUser).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(&registeredUser)
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(&registrationEvents).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)
+25
View File
@@ -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
View File
@@ -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
View File
@@ -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
+1
View File
@@ -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>
+33
View File
@@ -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}}
+1
View File
@@ -31,6 +31,7 @@ type Page struct {
MaxImages int
ImageGenerateAvailable bool
ImageEditAvailable bool
RegistrationEnabled bool
}
type Renderer struct{ templates *template.Template }