Files
goauto/server/app/goauto/device/handler.go
T
QiuSWandClaude Opus 5 3bf428acd7 fix(server): read purchaser identity from JWT claims for owned devices (#333)
go-admin's Authorizator runs per request with the IdentityHandler map,
which carries no user entry, so c.Get("userId") was always 0 and every
purchaser got an empty device list.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NTDbDcwbDw1TSAcE6wfh2F
2026-09-22 10:48:11 +08:00

256 lines
6.9 KiB
Go

package device
import (
stdcontext "context"
"encoding/json"
"errors"
"io"
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"github.com/go-admin-team/go-admin-core/sdk/pkg"
jwt "github.com/go-admin-team/go-admin-core/sdk/pkg/jwtauth"
"gorm.io/gorm"
)
type Handler struct {
DB *gorm.DB
}
func (handler Handler) List(context *gin.Context) {
request, err := parseListRequest(context)
if err != nil {
writeError(context, err)
return
}
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
if role, _ := jwt.ExtractClaims(context)["rolekey"].(string); role != "admin" {
id := currentUserID(context)
if id > 0 {
request.OwnerUserID = &id
} else {
request.OwnerUserID = new(uint64)
}
}
response, err := NewService(db).List(context.Request.Context(), request)
if err != nil {
writeError(context, err)
return
}
context.JSON(http.StatusOK, gin.H{"code": http.StatusOK, "data": response})
}
func currentUserID(c *gin.Context) uint64 {
// #333: go-admin's Authorizator runs on every request with the
// IdentityHandler map, which has no "user" entry, so c.Get("userId") is
// always 0 there. The JWT "identity" claim is the authenticated user id.
switch id := jwt.ExtractClaims(c)["identity"].(type) {
case float64:
if id > 0 {
return uint64(id)
}
case int:
if id > 0 {
return uint64(id)
}
case int64:
if id > 0 {
return uint64(id)
}
case uint64:
return id
}
return 0
}
func (handler Handler) Owners(context *gin.Context) {
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
rows, err := NewService(db).Owners(context.Request.Context())
if err != nil {
writeError(context, err)
return
}
context.JSON(http.StatusOK, gin.H{"code": http.StatusOK, "data": rows})
}
func (handler Handler) SetOwner(context *gin.Context) {
deviceID, err := strconv.ParseUint(context.Param("deviceId"), 10, 64)
if err != nil || deviceID == 0 {
writeError(context, invalidRequest("deviceId 无效"))
return
}
var req struct {
OwnerUserID *uint64 `json:"ownerUserId"`
}
if err := decodeJSON(context, &req); err != nil {
writeError(context, invalidRequest("请求 JSON 无效"))
return
}
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
if err := NewService(db).SetOwner(context.Request.Context(), deviceID, req.OwnerUserID); err != nil {
writeError(context, err)
return
}
context.JSON(http.StatusOK, gin.H{"code": http.StatusOK, "data": gin.H{"deviceId": deviceID, "ownerUserId": req.OwnerUserID}})
}
func (handler Handler) Register(context *gin.Context) {
request, err := decodeRegisterRequest(context)
if err != nil {
writeError(context, invalidRequest("请求 JSON 无效"))
return
}
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
response, err := NewService(db).Register(context.Request.Context(), request, bearerToken(context.GetHeader("Authorization")), strings.TrimSpace(context.GetHeader("X-GoAuto-Device-Recovery-Code")))
if err != nil {
writeError(context, err)
return
}
context.Header("Cache-Control", "no-store")
context.JSON(http.StatusOK, gin.H{"data": response})
}
func (handler Handler) Heartbeat(context *gin.Context) {
var request HeartbeatRequest
if err := decodeJSON(context, &request); err != nil {
writeError(context, invalidRequest("请求 JSON 无效"))
return
}
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
response, err := NewService(db).Heartbeat(context.Request.Context(), request, bearerToken(context.GetHeader("Authorization")))
if err != nil {
writeError(context, err)
return
}
context.JSON(http.StatusOK, gin.H{"data": response})
}
func (handler Handler) Disable(context *gin.Context) {
handler.adminAction(context, (*Service).Disable)
}
func (handler Handler) RevokeToken(context *gin.Context) {
handler.adminAction(context, (*Service).RevokeToken)
}
func (handler Handler) ResetIdentity(context *gin.Context) {
deviceID, err := strconv.ParseUint(context.Param("deviceId"), 10, 64)
if err != nil || deviceID == 0 {
writeError(context, invalidRequest("deviceId 无效"))
return
}
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
response, err := NewService(db).ResetIdentity(context.Request.Context(), deviceID)
if err != nil {
writeError(context, err)
return
}
context.Header("Cache-Control", "no-store")
context.JSON(http.StatusOK, gin.H{"code": http.StatusOK, "data": response})
}
func (handler Handler) adminAction(context *gin.Context, action func(*Service, stdcontext.Context, uint64) error) {
deviceID, err := strconv.ParseUint(context.Param("deviceId"), 10, 64)
if err != nil || deviceID == 0 {
writeError(context, invalidRequest("deviceId 无效"))
return
}
db, err := handler.database(context)
if err != nil {
writeError(context, internalError(err))
return
}
if err := action(NewService(db), context.Request.Context(), deviceID); err != nil {
writeError(context, err)
return
}
context.JSON(http.StatusOK, gin.H{"code": http.StatusOK, "data": gin.H{"deviceId": deviceID}})
}
func (handler Handler) database(context *gin.Context) (*gorm.DB, error) {
if handler.DB != nil {
return handler.DB, nil
}
return pkg.GetOrm(context)
}
func decodeRegisterRequest(context *gin.Context) (RegisterRequest, error) {
var request RegisterRequest
if err := decodeJSON(context, &request); err != nil {
return RegisterRequest{}, err
}
return request, nil
}
func decodeJSON(context *gin.Context, destination any) error {
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, 64<<10)
decoder := json.NewDecoder(context.Request.Body)
decoder.DisallowUnknownFields()
if err := decoder.Decode(destination); err != nil {
return err
}
if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
return errors.New("request body must contain one JSON object")
}
return nil
}
func bearerToken(header string) string {
parts := strings.Fields(header)
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
return ""
}
return parts[1]
}
func writeError(context *gin.Context, err error) {
var serviceError *ServiceError
if !errors.As(err, &serviceError) {
serviceError = internalError(err).(*ServiceError)
}
status := http.StatusInternalServerError
switch serviceError.Code {
case CodeInvalidRequest:
status = http.StatusUnprocessableEntity
case CodeInstallIDConflict:
status = http.StatusConflict
case CodeTokenInvalid:
status = http.StatusUnauthorized
case CodeDeviceDisabled:
status = http.StatusForbidden
case CodeDeviceTaskMismatch:
status = http.StatusConflict
case CodeDeviceNotFound:
status = http.StatusNotFound
}
context.JSON(status, gin.H{
"code": serviceError.Code, "message": serviceError.Message, "retryable": serviceError.Retryable,
})
}