Files
NixMsg/internal/app/identity/register.go
T

418 lines
12 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package identity
import (
"context"
"crypto/rand"
"crypto/subtle"
"database/sql"
"errors"
"io"
"log/slog"
"net"
"net/http"
"strings"
"time"
"unicode/utf8"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
const (
maxRegisterBodyBytes = 4 * 1024
settingRegistrationEnabled = "registration_enabled"
settingRegistrationCode = "registration_code"
sourceSelf = "self"
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
)
// APIError 是注册 HTTP/业务错误,带 HTTP 状态与协议错误码。
type APIError struct {
Status int
Code string
Message string
}
func (e *APIError) Error() string {
if e == nil {
return ""
}
if e.Message == "" {
return e.Code
}
return e.Code + ": " + e.Message
}
func apiErr(status int, code, msg string) *APIError {
return &APIError{Status: status, Code: code, Message: msg}
}
// RegisterConfig 是可挂载注册处理器的依赖。
// Hash / Locks 用 auth 接口;P3 未合入时测试可注入 StubHashPool 与可锁定的 LoginLocks。
type RegisterConfig struct {
DB *store.DB
Hash auth.HashPool
Locks auth.LoginLocks
Logger *slog.Logger
// Now 可测;nil 则用 time.Now。
Now func() time.Time
// ClientIP 可测;nil 则从 RemoteAddr 取 host。
ClientIP func(*http.Request) string
}
// RegisterHandler 处理 POST/OPTIONS /api/client/register(可挂到任意 ServeMux)。
type RegisterHandler struct {
cfg RegisterConfig
}
// NewRegisterHandler 构造可挂载的注册 Handler。DB/Hash/Locks 必填。
func NewRegisterHandler(cfg RegisterConfig) *RegisterHandler {
if cfg.Logger == nil {
cfg.Logger = slog.Default()
}
if cfg.Now == nil {
cfg.Now = time.Now
}
if cfg.ClientIP == nil {
cfg.ClientIP = clientIPFromRemoteAddr
}
return &RegisterHandler{cfg: cfg}
}
// ServeHTTP 实现 http.Handler。
func (h *RegisterHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
setCORS(w)
switch r.Method {
case http.MethodOptions:
w.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Content-Type")
w.WriteHeader(http.StatusNoContent)
return
case http.MethodPost:
h.handlePost(w, r)
default:
writeRegisterError(w, apiErr(http.StatusMethodNotAllowed, protocol.CodeBadRequest, "method not allowed"))
}
}
func (h *RegisterHandler) handlePost(w http.ResponseWriter, r *http.Request) {
ip := h.cfg.ClientIP(r)
r.Body = http.MaxBytesReader(w, r.Body, maxRegisterBodyBytes)
body, err := io.ReadAll(r.Body)
if err != nil {
var maxErr *http.MaxBytesError
if errors.As(err, &maxErr) || errors.Is(err, io.ErrUnexpectedEOF) || isBodyTooLarge(err) {
h.logResult("bad_request", "", ip)
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "body too large"))
return
}
h.logResult("bad_request", "", ip)
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "read body failed"))
return
}
req, err := protocol.DecodeRegister(body)
if err != nil {
h.logResult("bad_request", "", ip)
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "invalid json"))
return
}
result, apiErr := h.register(r.Context(), req, ip)
if apiErr != nil {
id := ""
if req != nil {
id = req.ID
}
h.logResult(apiErr.Code, id, ip)
writeRegisterError(w, apiErr)
return
}
h.logResult("ok", result.ID, ip)
writeRegisterOK(w, result)
}
func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRequest, ip string) (RegisterResult, *APIError) {
enabled, storedCode, err := h.loadRegistrationSettings(ctx)
if err != nil {
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "settings unavailable")
}
if !enabled {
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
}
lockKey := auth.LockKey{Kind: auth.LockRegisterIP, IP: ip}
if locked, _ := h.cfg.Locks.Check(lockKey); locked {
return RegisterResult{}, apiErr(http.StatusTooManyRequests, protocol.CodeRateLimited, "rate limited")
}
if !constantTimeEqual(req.RegistrationCode, storedCode) {
h.cfg.Locks.Fail(lockKey)
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationCodeInvalid, "registration code invalid")
}
if valErr := req.Validate(); valErr != nil {
code, msg := protocol.CodeBadRequest, valErr.Error()
var pe *protocol.Error
if errors.As(valErr, &pe) && pe != nil {
code, msg = pe.Code, pe.Message
}
return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg)
}
id := req.ID
loginPassword := req.LoginPassword
passwordGenerated := false
if loginPassword == "" {
pw, genErr := generateLoginPassword()
if genErr != nil {
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate password failed")
}
loginPassword = pw
passwordGenerated = true
}
loginHash, err := h.cfg.Hash.Hash(ctx, auth.PasswordLogin, loginPassword)
if err != nil {
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed")
}
var talkHash sql.NullString
if req.TalkPassword != "" {
th, hashErr := h.cfg.Hash.Hash(ctx, auth.PasswordTalk, req.TalkPassword)
if hashErr != nil {
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed")
}
talkHash = sql.NullString{String: th, Valid: true}
}
nowMs := h.cfg.Now().UnixMilli()
const maxIDAttempts = 8
for attempt := 0; attempt < maxIDAttempts; attempt++ {
useID := id
if useID == "" {
genID, genErr := generateEndpointID()
if genErr != nil {
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id failed")
}
useID = genID
}
insertErr := h.cfg.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, execErr := tx.ExecContext(ctx, `
INSERT INTO endpoints(
id, name, remark, source, login_hash, talk_hash, talk_version,
default_delay_ms, enabled, created_at
) VALUES (?, ?, '', ?, ?, ?, 0, 0, 1, ?)`,
useID, req.Name, sourceSelf, loginHash, talkHash, nowMs,
)
return execErr
})
if insertErr == nil {
out := RegisterResult{ID: useID}
if passwordGenerated {
out.LoginPassword = loginPassword
}
return out, nil
}
if isUniqueConstraint(insertErr) {
if id != "" {
return RegisterResult{}, apiErr(http.StatusConflict, protocol.CodeIDTaken, "id taken")
}
continue
}
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "insert failed")
}
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id exhausted")
}
func (h *RegisterHandler) loadRegistrationSettings(ctx context.Context) (enabled bool, code string, err error) {
var enabledVal, codeVal sql.NullString
row := h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationEnabled)
if scanErr := row.Scan(&enabledVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) {
return false, "", scanErr
}
row = h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationCode)
if scanErr := row.Scan(&codeVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) {
return false, "", scanErr
}
return settingTruthy(enabledVal.String), codeVal.String, nil
}
func (h *RegisterHandler) logResult(result, id, ip string) {
h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip)
}
// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。
type Server struct {
handler *RegisterHandler
}
// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。
func NewServer(cfg RegisterConfig) *Server {
return &Server{handler: NewRegisterHandler(cfg)}
}
// Handler 返回可挂载的注册 HTTP 处理器。
func (s *Server) Handler() http.Handler { return s.handler }
// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。
func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
preq := &protocol.RegisterRequest{
RegistrationCode: req.RegistrationCode,
ID: req.ID,
LoginPassword: req.LoginPassword,
Name: req.Name,
TalkPassword: req.TalkPassword,
}
ip := req.RemoteIP
if ip == "" {
ip = "0.0.0.0"
}
result, err := s.handler.register(ctx, preq, ip)
if err != nil {
return RegisterResult{}, err
}
return result, nil
}
func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) {
return SelfInfo{}, ErrNotImplemented
}
func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error {
return ErrNotImplemented
}
func (s *Server) SelfSetTalkPassword(context.Context, string, string) error {
return ErrNotImplemented
}
func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
return "", ErrNotImplemented
}
func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented }
func (s *Server) UnlockTalk(context.Context, string, string, string) error {
return ErrNotImplemented
}
func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) {
return false, nil
}
func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented }
func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented }
func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented }
var _ Service = (*Server)(nil)
var _ http.Handler = (*RegisterHandler)(nil)
func setCORS(w http.ResponseWriter) {
w.Header().Set("Access-Control-Allow-Origin", "*")
}
func writeRegisterOK(w http.ResponseWriter, result RegisterResult) {
setCORS(w)
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(http.StatusOK)
_ = protocol.Encode(w, protocol.RegisterResponse{
OK: true,
Data: protocol.RegisterData{
ID: result.ID,
LoginPassword: result.LoginPassword,
},
})
}
func writeRegisterError(w http.ResponseWriter, err *APIError) {
setCORS(w)
w.Header().Set("Content-Type", "application/json; charset=utf-8")
status := http.StatusInternalServerError
code, msg := protocol.CodeBusy, "internal error"
if err != nil {
status = err.Status
code, msg = err.Code, err.Message
}
w.WriteHeader(status)
_ = protocol.Encode(w, protocol.RegisterResponse{
OK: false,
Error: &protocol.ErrorBody{Code: code, Message: msg},
})
}
func clientIPFromRemoteAddr(r *http.Request) string {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return host
}
func settingTruthy(v string) bool {
switch strings.TrimSpace(strings.ToLower(v)) {
case "1", "true", "yes", "on":
return true
default:
return false
}
}
func constantTimeEqual(a, b string) bool {
// 长度不同时 ConstantTimeCompare 直接失败;先按较短侧对齐比较,再核对长度,避免过早返回。
ab := []byte(a)
bb := []byte(b)
if len(ab) != len(bb) {
dummy := make([]byte, len(ab))
subtle.ConstantTimeCompare(ab, dummy)
return false
}
return subtle.ConstantTimeCompare(ab, bb) == 1
}
func generateEndpointID() (string, error) {
b := make([]byte, 8)
if _, err := rand.Read(b); err != nil {
return "", err
}
out := make([]byte, 8)
for i := range b {
out[i] = idAlphabet[int(b[i])%len(idAlphabet)]
}
return "e_" + string(out), nil
}
func generateLoginPassword() (string, error) {
// 20 字节可读字符,满足 8–128,且不以 nst_ 开头(字母数字混合,冲突概率极低;若撞前缀则重抽)。
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
for range 8 {
b := make([]byte, 20)
if _, err := rand.Read(b); err != nil {
return "", err
}
out := make([]byte, 20)
for i := range b {
out[i] = alphabet[int(b[i])%len(alphabet)]
}
pw := string(out)
if !strings.HasPrefix(pw, protocol.SessionTokenPrefix) && utf8.RuneCountInString(pw) >= protocol.MinLoginPasswordLen {
return pw, nil
}
}
return "", errors.New("identity: generate login password failed")
}
func isUniqueConstraint(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "unique constraint") || strings.Contains(msg, "constraint failed")
}
func isBodyTooLarge(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "request body too large") || strings.Contains(msg, "http: request body too large")
}