391 lines
14 KiB
Go
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}
|
|
}
|