418 lines
12 KiB
Go
418 lines
12 KiB
Go
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")
|
||
}
|