Files
cmbuyer/admin/internal/server/task_claims.go
T

165 lines
4.9 KiB
Go

package server
import (
"bytes"
"encoding/json"
"errors"
"io"
"net/http"
"unicode/utf8"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/taskclaim"
"github.com/gin-gonic/gin"
)
const maxClaimJSONBytes = 4096
func claimNext(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
principal, ok := authenticateDevice(context, options)
if !ok {
return
}
var command taskclaim.ClaimCommand
if !decodeClaimJSON(context, &command) {
return
}
response, found, err := options.TaskClaims.ClaimNext(context.Request.Context(), principal.ID, command)
if err != nil {
writeTaskClaimError(context, err)
return
}
if !found {
context.Status(http.StatusNoContent)
return
}
context.JSON(http.StatusOK, response)
}
}
func renewLease(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
principal, ok := authenticateDevice(context, options)
if !ok {
return
}
var command taskclaim.RenewCommand
if !decodeClaimJSON(context, &command) {
return
}
command.TaskID = context.Param("id")
response, err := options.TaskClaims.Renew(context.Request.Context(), principal.ID, command)
if err != nil {
writeTaskClaimError(context, err)
return
}
context.JSON(http.StatusOK, response)
}
}
// Authentication precedes path interpretation, Content-Type parsing and every body read. This
// keeps rejected devices from using parsing differences as an oracle or making the server buffer data.
func authenticateDevice(context *gin.Context, options Options) (deviceauth.Principal, bool) {
principal, err := options.DeviceAuthenticator.Authenticate(context.Request)
if errors.Is(err, deviceauth.ErrUnauthenticated) {
context.Header("WWW-Authenticate", "Bearer")
context.Status(http.StatusUnauthorized)
return deviceauth.Principal{}, false
}
if err != nil || !deviceauth.ValidDeviceID(principal.ID) {
context.Status(http.StatusServiceUnavailable)
return deviceauth.Principal{}, false
}
return principal, true
}
func decodeClaimJSON(context *gin.Context, target any) bool {
if !isJSONContentType(context.GetHeader("Content-Type")) {
writeFixedError(context, http.StatusUnsupportedMediaType, "unsupported_media_type")
return false
}
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, maxClaimJSONBytes)
raw, err := io.ReadAll(context.Request.Body)
if err != nil {
var tooLarge *http.MaxBytesError
if errors.As(err, &tooLarge) {
writeFixedError(context, http.StatusRequestEntityTooLarge, "request_too_large")
} else {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
}
return false
}
if len(raw) == 0 || !utf8.Valid(raw) {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
if !hasUniqueTopLevelJSONFields(raw) {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
var extra any
if err := decoder.Decode(&extra); err != io.EOF {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
return true
}
func hasUniqueTopLevelJSONFields(raw []byte) bool {
decoder := json.NewDecoder(bytes.NewReader(raw))
first, err := decoder.Token()
if err != nil || first != json.Delim('{') {
return false
}
seen := make(map[string]struct{})
for decoder.More() {
key, err := decoder.Token()
name, ok := key.(string)
if err != nil || !ok {
return false
}
if _, duplicate := seen[name]; duplicate {
return false
}
seen[name] = struct{}{}
var value json.RawMessage
if err := decoder.Decode(&value); err != nil {
return false
}
}
last, err := decoder.Token()
return err == nil && last == json.Delim('}')
}
func writeTaskClaimError(context *gin.Context, err error) {
switch {
case errors.Is(err, taskclaim.ErrInvalid):
writeFixedError(context, http.StatusBadRequest, "invalid_request")
case errors.Is(err, taskclaim.ErrIdempotencyConflict):
writeFixedError(context, http.StatusConflict, "idempotency_conflict")
case errors.Is(err, taskclaim.ErrRequiresManual):
writeFixedError(context, http.StatusConflict, "claim_requires_manual")
case errors.Is(err, taskclaim.ErrNotCurrent):
writeFixedError(context, http.StatusConflict, "claim_not_current")
case errors.Is(err, taskclaim.ErrDeviceInactive):
context.Header("WWW-Authenticate", "Bearer")
context.Status(http.StatusUnauthorized)
default:
// Storage and transaction failures are intentionally bodyless: SQL, paths and candidate
// details are server-only and must not become a device-facing diagnostic oracle.
context.Status(http.StatusServiceUnavailable)
}
}
func writeFixedError(context *gin.Context, status int, code string) {
context.JSON(status, gin.H{"error": code})
}