Files
goauto/server/app/goauto/device/service.go
T

391 lines
14 KiB
Go

package device
import (
"context"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
"go-admin/app/goauto/models"
"github.com/google/uuid"
"gorm.io/gorm"
)
const DefaultHeartbeatIntervalSeconds = 15
const (
CodeInvalidRequest = "INVALID_REQUEST"
CodeInstallIDConflict = "DEVICE_INSTALL_ID_CONFLICT"
CodeTokenInvalid = "DEVICE_TOKEN_INVALID"
CodeDeviceDisabled = "DEVICE_DISABLED"
CodeDeviceTaskMismatch = "DEVICE_TASK_MISMATCH"
CodeDeviceNotFound = "DEVICE_NOT_FOUND"
CodeInternal = "INTERNAL_ERROR"
CodeRecoveryInvalid = "DEVICE_RECOVERY_INVALID"
CodeRecoveryExpired = "DEVICE_RECOVERY_EXPIRED"
)
const deviceRecoveryLifetime = 10 * time.Minute
type ServiceError struct {
Code string
Message string
Retryable bool
Cause error
}
func (err *ServiceError) Error() string {
if err.Cause == nil {
return err.Message
}
return fmt.Sprintf("%s: %v", err.Message, err.Cause)
}
func (err *ServiceError) Unwrap() error { return err.Cause }
type RegisterRequest struct {
RequestID string `json:"requestId"`
InstallID string `json:"installId"`
Name string `json:"name"`
Manufacturer string `json:"manufacturer"`
Model string `json:"model"`
AndroidVersion string `json:"androidVersion"`
AgentVersion string `json:"agentVersion"`
PDDVersion string `json:"pddVersion"`
Capabilities []string `json:"capabilities,omitempty"`
}
type RegisterResponse struct {
RequestID string `json:"requestId"`
DeviceID uint64 `json:"deviceId"`
DeviceToken string `json:"deviceToken,omitempty"`
HeartbeatIntervalSeconds int `json:"heartbeatIntervalSeconds"`
Created bool `json:"created"`
Replayed bool `json:"replayed,omitempty"`
}
type ResetIdentityResponse struct {
DeviceID uint64 `json:"deviceId"`
ExpiresAt time.Time `json:"expiresAt"`
}
type Service struct {
DB *gorm.DB
Now func() time.Time
GenerateToken func() (string, error)
HeartbeatInterval int
}
func NewService(db *gorm.DB) *Service {
return &Service{
DB: db,
Now: func() time.Time { return time.Now().UTC() },
GenerateToken: generateToken,
HeartbeatInterval: DefaultHeartbeatIntervalSeconds,
}
}
func generateToken() (string, error) {
raw := make([]byte, 32)
if _, err := rand.Read(raw); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(raw), nil
}
func tokenDigest(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}
func tokenMatches(token, digest string) bool {
if token == "" || len(digest) != sha256.Size*2 {
return false
}
want, err := hex.DecodeString(digest)
if err != nil {
return false
}
got := sha256.Sum256([]byte(token))
return subtle.ConstantTimeCompare(got[:], want) == 1
}
func digestMatches(value string, digest *string) bool {
return digest != nil && tokenMatches(value, *digest)
}
// Authenticate returns the active device represented by a bearer token.
// Agent feature packages use this method so token verification stays in one
// place and raw tokens never leave request memory.
func (service *Service) Authenticate(ctx context.Context, token string) (models.AgentDevice, error) {
if token == "" {
return models.AgentDevice{}, tokenInvalidError()
}
var result models.AgentDevice
err := service.DB.WithContext(ctx).
Where("token_digest = ? AND token_revoked_at IS NULL", tokenDigest(token)).
First(&result).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return models.AgentDevice{}, tokenInvalidError()
}
if err != nil {
return models.AgentDevice{}, internalError(err)
}
if result.Status == models.DeviceStatusDisabled {
return models.AgentDevice{}, &ServiceError{Code: CodeDeviceDisabled, Message: "设备已停用", Retryable: false}
}
return result, nil
}
func (service *Service) Register(ctx context.Context, request RegisterRequest, presentedToken string, recoveryCodes ...string) (RegisterResponse, error) {
recoveryCode := ""
if len(recoveryCodes) > 0 {
recoveryCode = recoveryCodes[0]
}
request = normalizeRegisterRequest(request)
if err := validateRegisterRequest(request); err != nil {
return RegisterResponse{}, err
}
if service.DB == nil {
return RegisterResponse{}, internalError(errors.New("database is nil"))
}
interval := service.HeartbeatInterval
if interval <= 0 {
interval = DefaultHeartbeatIntervalSeconds
}
response := RegisterResponse{RequestID: request.RequestID, HeartbeatIntervalSeconds: interval}
db := service.DB.WithContext(ctx)
var existing models.AgentDevice
err := db.Where("install_id = ?", request.InstallID).First(&existing).Error
if err == nil {
if existing.LastRegisterRequestID != nil && *existing.LastRegisterRequestID == request.RequestID {
response.DeviceID = existing.ID
response.Replayed = true
return response, nil
}
if existing.TokenRevokedAt != nil {
return RegisterResponse{}, &ServiceError{
Code: CodeInstallIDConflict, Message: "installId 已注册,需要该设备的有效 Token", Retryable: false,
}
}
if existing.Status == models.DeviceStatusDisabled {
return RegisterResponse{}, &ServiceError{Code: CodeDeviceDisabled, Message: "设备已停用", Retryable: false}
}
usingRecovery := !tokenMatches(presentedToken, existing.TokenDigest)
if usingRecovery {
validCode := recoveryCode != "" && digestMatches(recoveryCode, existing.RecoveryCodeDigest)
autoRecovery := recoveryCode == "" && existing.RecoveryCodeDigest == nil && existing.RecoveryExpiresAt != nil && existing.RecoveryUsedAt == nil
if !validCode && !autoRecovery {
return RegisterResponse{}, &ServiceError{Code: CodeInstallIDConflict, Message: "installId 已注册,需要该设备的有效 Token", Retryable: false}
}
if existing.RecoveryExpiresAt == nil || !service.Now().Before(*existing.RecoveryExpiresAt) {
return RegisterResponse{}, &ServiceError{Code: CodeRecoveryExpired, Message: "设备身份重置窗口已过期,请在后台重新操作", Retryable: false}
}
}
updates := map[string]any{
"name": request.Name, "manufacturer": request.Manufacturer, "model": request.Model,
"android_version": request.AndroidVersion, "agent_version": request.AgentVersion,
"pdd_version": request.PDDVersion, "last_register_request_id": request.RequestID,
}
if request.Capabilities != nil {
updates["capabilities_json"] = encodeCapabilities(request.Capabilities)
}
if usingRecovery {
newToken, generateErr := service.GenerateToken()
if generateErr != nil {
return RegisterResponse{}, internalError(generateErr)
}
now := service.Now()
updates["token_digest"] = tokenDigest(newToken)
updates["token_issued_at"] = now
updates["recovery_code_digest"] = nil
updates["recovery_expires_at"] = nil
updates["recovery_used_at"] = now
updates["status"] = models.DeviceStatusOnline
response.DeviceToken = newToken
}
result := db.Model(&models.AgentDevice{}).
Where("id = ? AND token_digest = ? AND token_revoked_at IS NULL", existing.ID, existing.TokenDigest).
Updates(updates)
if result.Error != nil {
return RegisterResponse{}, internalError(result.Error)
}
if result.RowsAffected == 0 {
return RegisterResponse{}, &ServiceError{
Code: CodeInstallIDConflict, Message: "设备 Token 已失效", Retryable: false,
}
}
response.DeviceID = existing.ID
return response, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return RegisterResponse{}, internalError(err)
}
token, err := service.GenerateToken()
if err != nil {
return RegisterResponse{}, internalError(err)
}
now := service.Now()
device := models.AgentDevice{
InstallID: request.InstallID, Name: request.Name, Manufacturer: request.Manufacturer,
Model: request.Model, AndroidVersion: request.AndroidVersion, AgentVersion: request.AgentVersion,
PDDVersion: request.PDDVersion, CapabilitiesJSON: encodeCapabilities(request.Capabilities), Status: models.DeviceStatusOnline, TokenDigest: tokenDigest(token),
TokenIssuedAt: now, LastRegisterRequestID: &request.RequestID,
}
if err := db.Create(&device).Error; err != nil {
// A concurrent request may have won the unique install_id race. Query
// outside a failed transaction so this also works on PostgreSQL.
var concurrent models.AgentDevice
if findErr := db.Where("install_id = ?", request.InstallID).First(&concurrent).Error; findErr == nil &&
concurrent.LastRegisterRequestID != nil && *concurrent.LastRegisterRequestID == request.RequestID {
response.DeviceID = concurrent.ID
response.Replayed = true
return response, nil
}
return RegisterResponse{}, internalError(err)
}
response.DeviceID = device.ID
response.DeviceToken = token
response.Created = true
return response, nil
}
// ResetIdentity invalidates the current token and opens a short-lived automatic
// re-registration window for the same installId. It retains the device row and
// task bindings, so no recovery code needs to leave the Admin workflow.
func (service *Service) ResetIdentity(ctx context.Context, deviceID uint64) (ResetIdentityResponse, error) {
if service.DB == nil {
return ResetIdentityResponse{}, internalError(errors.New("database is nil"))
}
now := service.Now()
expiresAt := now.Add(deviceRecoveryLifetime)
response := ResetIdentityResponse{DeviceID: deviceID, ExpiresAt: expiresAt}
err := service.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var device models.AgentDevice
if err := tx.First(&device, deviceID).Error; errors.Is(err, gorm.ErrRecordNotFound) {
return &ServiceError{Code: CodeDeviceNotFound, Message: "设备不存在", Retryable: false}
} else if err != nil {
return internalError(err)
}
if device.Status == models.DeviceStatusDisabled {
return &ServiceError{Code: CodeDeviceDisabled, Message: "设备已停用", Retryable: false}
}
invalidatedToken, generateErr := service.GenerateToken()
if generateErr != nil {
return internalError(generateErr)
}
return tx.Model(&models.AgentDevice{}).Where("id = ?", deviceID).Updates(map[string]any{
"token_digest": tokenDigest(invalidatedToken), "token_issued_at": now,
"recovery_code_digest": nil, "recovery_expires_at": expiresAt,
"recovery_used_at": nil, "status": models.DeviceStatusOffline,
}).Error
})
if err != nil {
return ResetIdentityResponse{}, err
}
return response, nil
}
func (service *Service) Disable(ctx context.Context, deviceID uint64) error {
return service.deactivate(ctx, deviceID, false)
}
func (service *Service) RevokeToken(ctx context.Context, deviceID uint64) error {
return service.deactivate(ctx, deviceID, true)
}
func (service *Service) deactivate(ctx context.Context, deviceID uint64, revoke bool) error {
now := service.Now()
return service.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
updates := map[string]any{"status": models.DeviceStatusDisabled}
if revoke {
updates["token_revoked_at"] = now
}
result := tx.Model(&models.AgentDevice{}).Where("id = ?", deviceID).Updates(updates)
if result.Error != nil {
return internalError(result.Error)
}
if result.RowsAffected == 0 {
var count int64
if err := tx.Model(&models.AgentDevice{}).Where("id = ?", deviceID).Count(&count).Error; err != nil {
return internalError(err)
}
if count == 0 {
return &ServiceError{Code: CodeDeviceNotFound, Message: "设备不存在", Retryable: false}
}
}
errorCode, errorMessage := CodeDeviceDisabled, "设备已由管理员停用"
if err := tx.Session(&gorm.Session{SkipHooks: true}).Model(&models.CollectionTask{}).
Where("device_id = ? AND status = ?", deviceID, models.TaskStatusRunning).
Updates(map[string]any{
"status": models.TaskStatusFailed, "active_slot": gorm.Expr("NULL"),
"device_run_slot": gorm.Expr("NULL"), "lease_expires_at": gorm.Expr("NULL"),
"error_code": errorCode, "error_message": errorMessage, "finished_at": now,
}).Error; err != nil {
return internalError(err)
}
return nil
})
}
func normalizeRegisterRequest(request RegisterRequest) RegisterRequest {
request.RequestID = strings.TrimSpace(request.RequestID)
request.InstallID = strings.ToLower(strings.TrimSpace(request.InstallID))
request.Name = strings.TrimSpace(request.Name)
request.Manufacturer = strings.TrimSpace(request.Manufacturer)
request.Model = strings.TrimSpace(request.Model)
request.AndroidVersion = strings.TrimSpace(request.AndroidVersion)
request.AgentVersion = strings.TrimSpace(request.AgentVersion)
request.PDDVersion = strings.TrimSpace(request.PDDVersion)
if request.Capabilities != nil {
if normalized, err := normalizeCapabilities(request.Capabilities); err == nil {
request.Capabilities = normalized
}
}
return request
}
func validateRegisterRequest(request RegisterRequest) error {
if _, err := uuid.Parse(request.RequestID); err != nil {
return invalidRequest("requestId 必须是 UUID")
}
if _, err := uuid.Parse(request.InstallID); err != nil {
return invalidRequest("installId 必须是 UUID")
}
fields := []struct {
name string
value string
max int
}{
{"name", request.Name, 100}, {"manufacturer", request.Manufacturer, 100},
{"model", request.Model, 100}, {"androidVersion", request.AndroidVersion, 32},
{"agentVersion", request.AgentVersion, 32}, {"pddVersion", request.PDDVersion, 32},
}
for _, field := range fields {
if field.value == "" || len(field.value) > field.max {
return invalidRequest(fmt.Sprintf("%s 必填且长度不能超过 %d", field.name, field.max))
}
}
if _, err := normalizeCapabilities(request.Capabilities); err != nil {
return err
}
return nil
}
func invalidRequest(message string) error {
return &ServiceError{Code: CodeInvalidRequest, Message: message, Retryable: false}
}
func internalError(cause error) error {
return &ServiceError{Code: CodeInternal, Message: "服务端处理失败", Retryable: true, Cause: cause}
}