From a5e9bfb15c31acf189b28b0eed3fabf12646c629 Mon Sep 17 00:00:00 2001 From: Naeel Date: Fri, 10 Apr 2026 19:41:59 +0300 Subject: [PATCH] security: fix critical/high auth, idor, races and persistence --- app/admin/admin.go | 24 +- app/auth/auth_middleware.go | 463 +++++++++++++++++++++++++----- app/gosqs/delete_message_batch.go | 185 ++++++------ app/gosqs/get_queue_url.go | 65 ++--- app/gosqs/list_queues.go | 70 ++--- app/gosqs/receive_message.go | 279 +++++++++--------- app/gosqs/send_message.go | 18 +- app/gosqs/send_message_batch.go | 245 ++++++++-------- app/gosqs/validation.go | 14 +- app/router/router.go | 6 +- app/tenant/tenant_store.go | 4 + 11 files changed, 868 insertions(+), 505 deletions(-) diff --git a/app/admin/admin.go b/app/admin/admin.go index 54a5c96..e394104 100644 --- a/app/admin/admin.go +++ b/app/admin/admin.go @@ -14,6 +14,7 @@ import ( "shared-sqs/app/auth" "shared-sqs/app/models" + "shared-sqs/app/persistence" "shared-sqs/app/tenant" "github.com/google/uuid" @@ -36,9 +37,9 @@ func findQueue(tenantAccessKey, queueName string) (string, *models.Queue) { // Handler — admin API handler, holds TenantStore и admin token type Handler struct { - store *tenant.TenantStore - adminToken string - nubesEndpoint string // URL nubes API для валидации JWT (напр. https://deck-api-test.ngcloud.ru/api/v1) + store *tenant.TenantStore + adminToken string + nubesEndpoint string // URL nubes API для валидации JWT (напр. https://deck-api-test.ngcloud.ru/api/v1) } // NewHandler — создаёт admin handler @@ -188,12 +189,18 @@ func (h *Handler) jwtMiddleware(next http.Handler) http.Handler { } // Проверяем что тенант существует (был создан при /ui/api/auth) - _, ok := h.store.GetBySub(claims.Sub) + jwtTenant, ok := h.store.GetBySub(claims.Sub) if !ok { jsonErr(w, http.StatusForbidden, "tenant not found — authenticate first via POST /ui/api/auth") return } + // IDOR защита: для /ui/api/tenants/{id}/... разрешаем доступ только к своему tenant ID. + if pathTenantID, exists := mux.Vars(r)["id"]; exists && pathTenantID != "" && pathTenantID != jwtTenant.ID { + jsonErr(w, http.StatusForbidden, "forbidden tenant access") + return + } + next.ServeHTTP(w, r) }) } @@ -315,13 +322,18 @@ func (h *Handler) deleteTenant(w http.ResponseWriter, r *http.Request) { } // Удаляем все очереди тенанта из SyncQueues prefix := t.AccessKey + ":" + deletedQueueKeys := make([]string, 0) models.SyncQueues.Lock() for key := range models.SyncQueues.Queues { if strings.HasPrefix(key, prefix) { + deletedQueueKeys = append(deletedQueueKeys, key) delete(models.SyncQueues.Queues, key) } } models.SyncQueues.Unlock() + for _, queueKey := range deletedQueueKeys { + persistence.DeleteQueue(queueKey) + } h.store.Delete(id) w.WriteHeader(http.StatusNoContent) @@ -427,6 +439,7 @@ func (h *Handler) createTenantQueue(w http.ResponseWriter, r *http.Request) { Messages: []models.SqsMessage{}, Duplicates: make(map[string]time.Time), } + persistence.SaveQueue(key, models.SyncQueues.Queues[key]) models.SyncQueues.Unlock() log.Infof("admin: created queue %s for tenant %s", req.Name, t.ID) w.Header().Set("Content-Type", "application/json") @@ -453,6 +466,7 @@ func (h *Handler) deleteTenantQueue(w http.ResponseWriter, r *http.Request) { } delete(models.SyncQueues.Queues, key) models.SyncQueues.Unlock() + persistence.DeleteQueue(key) log.Infof("admin: deleted queue %s for tenant %s", queueName, t.ID) w.WriteHeader(http.StatusNoContent) } @@ -552,6 +566,7 @@ func (h *Handler) sendMessageToQueue(w http.ResponseWriter, r *http.Request) { } models.SyncQueues.Lock() models.SyncQueues.Queues[key].Messages = append(models.SyncQueues.Queues[key].Messages, msg) + persistence.SaveQueue(key, models.SyncQueues.Queues[key]) models.SyncQueues.Unlock() w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusCreated) @@ -575,6 +590,7 @@ func (h *Handler) purgeQueue(w http.ResponseWriter, r *http.Request) { } models.SyncQueues.Lock() models.SyncQueues.Queues[key].Messages = models.SyncQueues.Queues[key].Messages[:0] + persistence.SaveQueue(key, models.SyncQueues.Queues[key]) models.SyncQueues.Unlock() log.Infof("admin: purged queue %s for tenant %s", queueName, t.ID) w.WriteHeader(http.StatusNoContent) diff --git a/app/auth/auth_middleware.go b/app/auth/auth_middleware.go index 997363d..90a0029 100644 --- a/app/auth/auth_middleware.go +++ b/app/auth/auth_middleware.go @@ -4,12 +4,24 @@ package auth import ( -"context" -"encoding/xml" -"net/http" -"strings" + "bytes" + "context" + "crypto/hmac" + "crypto/sha256" + "crypto/subtle" + "encoding/hex" + "encoding/xml" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "sort" + "strconv" + "strings" + "time" -"shared-sqs/app/tenant" + "shared-sqs/app/tenant" ) // TenantContextKey — ключ для хранения тенанта в request context. @@ -18,39 +30,49 @@ type contextKey string const TenantContextKey contextKey = "tenant" +const ( + sigV4DateFormat = "20060102T150405Z" + maxClockSkew = 15 * time.Minute +) + // AuthMiddleware — middleware: ищет тенанта по AccessKeyId из AWS Authorization header. // Пропускает /health и /admin/** без tenant-аутентификации. // Ловушка #4: не ставим короткий таймаут — ReceiveMessage с long polling держит соединение до 20 сек. func AuthMiddleware(store *tenant.TenantStore) func(http.Handler) http.Handler { -return func(next http.Handler) http.Handler { -return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { -// /health — без auth -if r.URL.Path == "/health" { -next.ServeHTTP(w, r) -return -} -// /admin/** — отдельная auth (bearer token, см. admin_handlers.go) -if strings.HasPrefix(r.URL.Path, "/admin/") { -next.ServeHTTP(w, r) -return -} + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // /health — без auth + if r.URL.Path == "/health" { + next.ServeHTTP(w, r) + return + } + // /admin/** — отдельная auth (bearer token, см. admin_handlers.go) + if strings.HasPrefix(r.URL.Path, "/admin/") { + next.ServeHTTP(w, r) + return + } -accessKeyID := extractAccessKeyID(r) -if accessKeyID == "" { -writeSQSAuthError(w, "MissingAuthenticationToken", "Request must contain either AccessKeyId or X-Amz-Credential") -return -} + accessKeyID := extractAccessKeyID(r) + if accessKeyID == "" { + writeSQSAuthError(w, "MissingAuthenticationToken", "Request must contain either AccessKeyId or X-Amz-Credential") + return + } -t, ok := store.GetByAccessKey(accessKeyID) -if !ok || !t.Active { -writeSQSAuthError(w, "InvalidClientTokenId", "The security token included in the request is invalid") -return -} + t, ok := store.GetByAccessKey(accessKeyID) + if !ok || !t.Active { + writeSQSAuthError(w, "InvalidClientTokenId", "The security token included in the request is invalid") + return + } -ctx := context.WithValue(r.Context(), TenantContextKey, t) -next.ServeHTTP(w, r.WithContext(ctx)) -}) -} + if err := verifySigV4(r, t); err != nil { + writeSQSAuthError(w, "SignatureDoesNotMatch", err.Error()) + return + } + + ctx := context.WithValue(r.Context(), TenantContextKey, t) + next.ServeHTTP(w, r.WithContext(ctx)) + }) + } } // extractAccessKeyID — извлекает AWS AccessKeyId из запроса. @@ -58,57 +80,360 @@ next.ServeHTTP(w, r.WithContext(ctx)) // Ловушка #3: AWS CLI ВСЕГДА отправляет Signature V4 — нужно парсить, даже не проверяя подпись. // Ловушка #5: X-Amz-Security-Token (STS) — игнорируем. func extractAccessKeyID(r *http.Request) string { -// Вариант 1: Authorization header -// Формат: "AWS4-HMAC-SHA256 Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request, ..." -auth := r.Header.Get("Authorization") -if strings.HasPrefix(auth, "AWS4-HMAC-SHA256") { -idx := strings.Index(auth, "Credential=") -if idx >= 0 { -rest := auth[idx+len("Credential="):] -slashIdx := strings.Index(rest, "/") -if slashIdx > 0 { -return rest[:slashIdx] -} -} + // Вариант 1: Authorization header + // Формат: "AWS4-HMAC-SHA256 Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request, ..." + auth := r.Header.Get("Authorization") + if strings.HasPrefix(auth, "AWS4-HMAC-SHA256") { + idx := strings.Index(auth, "Credential=") + if idx >= 0 { + rest := auth[idx+len("Credential="):] + slashIdx := strings.Index(rest, "/") + if slashIdx > 0 { + return rest[:slashIdx] + } + } + } + + // Вариант 2: Query parameter (presigned URLs) + // Формат: X-Amz-Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request + if cred := r.URL.Query().Get("X-Amz-Credential"); cred != "" { + parts := strings.SplitN(cred, "/", 2) + if len(parts) > 0 && parts[0] != "" { + return parts[0] + } + } + + return "" } -// Вариант 2: Query parameter (presigned URLs) -// Формат: X-Amz-Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request -if cred := r.URL.Query().Get("X-Amz-Credential"); cred != "" { -parts := strings.SplitN(cred, "/", 2) -if len(parts) > 0 && parts[0] != "" { -return parts[0] -} +// verifySigV4 — полноценная проверка подписи AWS SigV4 для header и presigned запросов. +func verifySigV4(r *http.Request, t *tenant.Tenant) error { + authHeader := r.Header.Get("Authorization") + if strings.HasPrefix(authHeader, "AWS4-HMAC-SHA256") { + return verifyHeaderSigV4(r, t, authHeader) + } + if r.URL.Query().Get("X-Amz-Credential") != "" { + return verifyPresignedSigV4(r, t) + } + return errors.New("missing SigV4 authorization") } -return "" +func verifyHeaderSigV4(r *http.Request, t *tenant.Tenant, authHeader string) error { + credential, signedHeaders, providedSignature, err := parseAuthorizationHeader(authHeader) + if err != nil { + return err + } + accessKey, shortDate, region, service, terminal, err := parseCredentialScope(credential) + if err != nil { + return err + } + if accessKey != t.AccessKey { + return errors.New("access key mismatch") + } + if service != "sqs" || terminal != "aws4_request" { + return errors.New("invalid credential scope") + } + + amzDate := r.Header.Get("X-Amz-Date") + if amzDate == "" { + return errors.New("missing X-Amz-Date header") + } + requestTime, err := time.Parse(sigV4DateFormat, amzDate) + if err != nil { + return fmt.Errorf("invalid X-Amz-Date: %w", err) + } + if absDuration(time.Now().UTC().Sub(requestTime)) > maxClockSkew { + return errors.New("request time skew is too large") + } + + body, err := readAndRestoreBody(r) + if err != nil { + return err + } + payloadHash := r.Header.Get("X-Amz-Content-Sha256") + if payloadHash == "" { + payloadHash = sha256Hex(body) + } + + canonicalRequest, err := buildCanonicalRequest(r, signedHeaders, payloadHash, true) + if err != nil { + return err + } + credentialScope := strings.Join([]string{shortDate, region, service, terminal}, "/") + stringToSign := buildStringToSign(amzDate, credentialScope, canonicalRequest) + expectedSignature := calculateSigV4Signature(t.SecretKey, shortDate, region, service, stringToSign) + + if subtle.ConstantTimeCompare([]byte(strings.ToLower(expectedSignature)), []byte(strings.ToLower(providedSignature))) != 1 { + return errors.New("signature mismatch") + } + return nil +} + +func verifyPresignedSigV4(r *http.Request, t *tenant.Tenant) error { + query := r.URL.Query() + if query.Get("X-Amz-Algorithm") != "AWS4-HMAC-SHA256" { + return errors.New("unsupported X-Amz-Algorithm") + } + credential := query.Get("X-Amz-Credential") + if credential == "" { + return errors.New("missing X-Amz-Credential") + } + providedSignature := query.Get("X-Amz-Signature") + if providedSignature == "" { + return errors.New("missing X-Amz-Signature") + } + signedHeaders := query.Get("X-Amz-SignedHeaders") + if signedHeaders == "" { + return errors.New("missing X-Amz-SignedHeaders") + } + amzDate := query.Get("X-Amz-Date") + if amzDate == "" { + return errors.New("missing X-Amz-Date") + } + expiresSeconds, err := strconv.Atoi(query.Get("X-Amz-Expires")) + if err != nil || expiresSeconds < 1 || expiresSeconds > 604800 { + return errors.New("invalid X-Amz-Expires") + } + + requestTime, err := time.Parse(sigV4DateFormat, amzDate) + if err != nil { + return fmt.Errorf("invalid X-Amz-Date: %w", err) + } + if time.Now().UTC().After(requestTime.Add(time.Duration(expiresSeconds) * time.Second)) { + return errors.New("presigned request has expired") + } + + accessKey, shortDate, region, service, terminal, err := parseCredentialScope(credential) + if err != nil { + return err + } + if accessKey != t.AccessKey { + return errors.New("access key mismatch") + } + if service != "sqs" || terminal != "aws4_request" { + return errors.New("invalid credential scope") + } + + body, err := readAndRestoreBody(r) + if err != nil { + return err + } + payloadHash := query.Get("X-Amz-Content-Sha256") + if payloadHash == "" { + payloadHash = r.Header.Get("X-Amz-Content-Sha256") + } + if payloadHash == "" { + payloadHash = "UNSIGNED-PAYLOAD" + if len(body) > 0 { + payloadHash = sha256Hex(body) + } + } + + canonicalRequest, err := buildCanonicalRequest(r, signedHeaders, payloadHash, false) + if err != nil { + return err + } + credentialScope := strings.Join([]string{shortDate, region, service, terminal}, "/") + stringToSign := buildStringToSign(amzDate, credentialScope, canonicalRequest) + expectedSignature := calculateSigV4Signature(t.SecretKey, shortDate, region, service, stringToSign) + + if subtle.ConstantTimeCompare([]byte(strings.ToLower(expectedSignature)), []byte(strings.ToLower(providedSignature))) != 1 { + return errors.New("signature mismatch") + } + return nil +} + +func parseAuthorizationHeader(authHeader string) (credential string, signedHeaders string, signature string, err error) { + if !strings.HasPrefix(authHeader, "AWS4-HMAC-SHA256") { + return "", "", "", errors.New("unsupported authorization algorithm") + } + fields := strings.Split(strings.TrimSpace(strings.TrimPrefix(authHeader, "AWS4-HMAC-SHA256")), ",") + parts := map[string]string{} + for _, field := range fields { + kv := strings.SplitN(strings.TrimSpace(field), "=", 2) + if len(kv) != 2 { + continue + } + parts[kv[0]] = kv[1] + } + credential = parts["Credential"] + signedHeaders = strings.ToLower(parts["SignedHeaders"]) + signature = parts["Signature"] + if credential == "" || signedHeaders == "" || signature == "" { + return "", "", "", errors.New("malformed Authorization header") + } + return credential, signedHeaders, signature, nil +} + +func parseCredentialScope(credential string) (accessKey, shortDate, region, service, terminal string, err error) { + parts := strings.Split(credential, "/") + if len(parts) != 5 { + return "", "", "", "", "", errors.New("invalid Credential scope") + } + return parts[0], parts[1], parts[2], parts[3], parts[4], nil +} + +func buildCanonicalRequest(r *http.Request, signedHeaders, payloadHash string, includeSignatureQuery bool) (string, error) { + canonicalURI := r.URL.EscapedPath() + if canonicalURI == "" { + canonicalURI = "/" + } + canonicalQuery := canonicalizeQueryString(r.URL.Query(), includeSignatureQuery) + + headers := strings.Split(strings.ToLower(signedHeaders), ";") + canonicalHeaders := strings.Builder{} + cleanSignedHeaders := make([]string, 0, len(headers)) + for _, h := range headers { + name := strings.TrimSpace(h) + if name == "" { + continue + } + value, ok := canonicalHeaderValue(r, name) + if !ok { + return "", fmt.Errorf("missing signed header: %s", name) + } + canonicalHeaders.WriteString(name) + canonicalHeaders.WriteString(":") + canonicalHeaders.WriteString(normalizeHeaderSpace(value)) + canonicalHeaders.WriteString("\n") + cleanSignedHeaders = append(cleanSignedHeaders, name) + } + if len(cleanSignedHeaders) == 0 { + return "", errors.New("empty signed headers") + } + + return strings.Join([]string{ + r.Method, + canonicalURI, + canonicalQuery, + canonicalHeaders.String(), + strings.Join(cleanSignedHeaders, ";"), + payloadHash, + }, "\n"), nil +} + +func canonicalHeaderValue(r *http.Request, name string) (string, bool) { + if name == "host" { + h := r.Host + if h == "" { + h = r.URL.Host + } + return h, h != "" + } + key := http.CanonicalHeaderKey(name) + values, ok := r.Header[key] + if !ok || len(values) == 0 { + return "", false + } + return strings.Join(values, ","), true +} + +func canonicalizeQueryString(query url.Values, includeSignature bool) string { + pairs := make([]string, 0) + for key, values := range query { + if !includeSignature && strings.EqualFold(key, "X-Amz-Signature") { + continue + } + sortedValues := append([]string(nil), values...) + sort.Strings(sortedValues) + if len(sortedValues) == 0 { + pairs = append(pairs, awsPercentEncode(key)+"=") + continue + } + for _, value := range sortedValues { + pairs = append(pairs, awsPercentEncode(key)+"="+awsPercentEncode(value)) + } + } + sort.Strings(pairs) + return strings.Join(pairs, "&") +} + +func awsPercentEncode(s string) string { + encoded := url.QueryEscape(s) + encoded = strings.ReplaceAll(encoded, "+", "%20") + encoded = strings.ReplaceAll(encoded, "*", "%2A") + encoded = strings.ReplaceAll(encoded, "%7E", "~") + return encoded +} + +func normalizeHeaderSpace(v string) string { + return strings.Join(strings.Fields(strings.TrimSpace(v)), " ") +} + +func buildStringToSign(amzDate, credentialScope, canonicalRequest string) string { + canonicalHash := sha256.Sum256([]byte(canonicalRequest)) + return strings.Join([]string{ + "AWS4-HMAC-SHA256", + amzDate, + credentialScope, + hex.EncodeToString(canonicalHash[:]), + }, "\n") +} + +func calculateSigV4Signature(secretKey, shortDate, region, service, stringToSign string) string { + kDate := hmacSHA256([]byte("AWS4"+secretKey), shortDate) + kRegion := hmacSHA256(kDate, region) + kService := hmacSHA256(kRegion, service) + kSigning := hmacSHA256(kService, "aws4_request") + sig := hmacSHA256(kSigning, stringToSign) + return hex.EncodeToString(sig) +} + +func hmacSHA256(key []byte, data string) []byte { + m := hmac.New(sha256.New, key) + _, _ = m.Write([]byte(data)) + return m.Sum(nil) +} + +func readAndRestoreBody(r *http.Request) ([]byte, error) { + if r.Body == nil { + return nil, nil + } + body, err := io.ReadAll(r.Body) + if err != nil { + return nil, fmt.Errorf("read request body: %w", err) + } + r.Body = io.NopCloser(bytes.NewReader(body)) + return body, nil +} + +func sha256Hex(data []byte) string { + sum := sha256.Sum256(data) + return hex.EncodeToString(sum[:]) +} + +func absDuration(d time.Duration) time.Duration { + if d < 0 { + return -d + } + return d } // sqsAuthError — AWS-совместимый XML ответ об ошибке аутентификации. type sqsAuthError struct { -XMLName xml.Name `xml:"ErrorResponse"` -Error sqsErrorBody `xml:"Error"` -RequestID string `xml:"RequestId"` + XMLName xml.Name `xml:"ErrorResponse"` + Error sqsErrorBody `xml:"Error"` + RequestID string `xml:"RequestId"` } type sqsErrorBody struct { -Type string `xml:"Type"` -Code string `xml:"Code"` -Message string `xml:"Message"` + Type string `xml:"Type"` + Code string `xml:"Code"` + Message string `xml:"Message"` } // writeSQSAuthError — отвечает AWS-совместимым XML с кодом 403. func writeSQSAuthError(w http.ResponseWriter, code, message string) { -w.Header().Set("Content-Type", "application/xml") -w.WriteHeader(http.StatusForbidden) -resp := sqsAuthError{ -Error: sqsErrorBody{ -Type: "Sender", -Code: code, -Message: message, -}, -RequestID: "00000000-0000-0000-0000-000000000000", -} -data, _ := xml.Marshal(resp) -w.Write(data) + w.Header().Set("Content-Type", "application/xml") + w.WriteHeader(http.StatusForbidden) + resp := sqsAuthError{ + Error: sqsErrorBody{ + Type: "Sender", + Code: code, + Message: message, + }, + RequestID: "00000000-0000-0000-0000-000000000000", + } + data, _ := xml.Marshal(resp) + w.Write(data) } diff --git a/app/gosqs/delete_message_batch.go b/app/gosqs/delete_message_batch.go index ebafc99..df17c68 100644 --- a/app/gosqs/delete_message_batch.go +++ b/app/gosqs/delete_message_batch.go @@ -3,117 +3,120 @@ package gosqs import ( -"net/http" -"strings" + "net/http" + "strings" -"shared-sqs/app/interfaces" -"shared-sqs/app/models" -"shared-sqs/app/utils" -"github.com/gorilla/mux" -log "github.com/sirupsen/logrus" + "github.com/gorilla/mux" + log "github.com/sirupsen/logrus" + "shared-sqs/app/interfaces" + "shared-sqs/app/models" + "shared-sqs/app/persistence" + "shared-sqs/app/utils" ) func DeleteMessageBatchV1(req *http.Request) (int, interfaces.AbstractResponseBody) { -requestBody := models.NewDeleteMessageBatchRequest() -ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) -if !ok { -log.Error("Invalid Request - DeleteMessageBatchV1") -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} + requestBody := models.NewDeleteMessageBatchRequest() + ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) + if !ok { + log.Error("Invalid Request - DeleteMessageBatchV1") + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } -t := getTenantFromContext(req) -if t == nil { -return utils.CreateErrorResponseV1("InvalidClientTokenId", true) -} + t := getTenantFromContext(req) + if t == nil { + return utils.CreateErrorResponseV1("InvalidClientTokenId", true) + } -queueUrl := requestBody.QueueUrl -queueName := "" -if queueUrl == "" { -vars := mux.Vars(req) -queueName = vars["queueName"] -} else { -uriSegments := strings.Split(queueUrl, "/") -queueName = uriSegments[len(uriSegments)-1] -} + queueUrl := requestBody.QueueUrl + queueName := "" + if queueUrl == "" { + vars := mux.Vars(req) + queueName = vars["queueName"] + } else { + uriSegments := strings.Split(queueUrl, "/") + queueName = uriSegments[len(uriSegments)-1] + } -key := tenantQueueKey(t.AccessKey, queueName) + key := tenantQueueKey(t.AccessKey, queueName) -if _, ok := models.SyncQueues.Queues[key]; !ok { -return utils.CreateErrorResponseV1("QueueNotFound", true) -} + if _, ok := models.SyncQueues.Queues[key]; !ok { + return utils.CreateErrorResponseV1("QueueNotFound", true) + } -if len(requestBody.Entries) == 0 { -return utils.CreateErrorResponseV1("EmptyBatchRequest", true) -} + if len(requestBody.Entries) == 0 { + return utils.CreateErrorResponseV1("EmptyBatchRequest", true) + } -if len(requestBody.Entries) > 10 { -return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true) -} + if len(requestBody.Entries) > 10 { + return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true) + } -ids := map[string]bool{} -for _, v := range requestBody.Entries { -if _, found := ids[v.Id]; found { -return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true) -} -ids[v.Id] = true -} + ids := map[string]bool{} + for _, v := range requestBody.Entries { + if _, found := ids[v.Id]; found { + return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true) + } + ids[v.Id] = true + } -models.SyncQueues.Lock() -defer models.SyncQueues.Unlock() + models.SyncQueues.Lock() + defer models.SyncQueues.Unlock() -deleteMessageMap := make(map[string]*deleteEntry) -for _, entry := range requestBody.Entries { -deleteMessageMap[entry.ReceiptHandle] = &deleteEntry{ -Id: entry.Id, -ReceiptHandle: entry.ReceiptHandle, -Deleted: false, -} -} + deleteMessageMap := make(map[string]*deleteEntry) + for _, entry := range requestBody.Entries { + deleteMessageMap[entry.ReceiptHandle] = &deleteEntry{ + Id: entry.Id, + ReceiptHandle: entry.ReceiptHandle, + Deleted: false, + } + } -deletedEntries := make([]models.DeleteMessageBatchResultEntry, 0) -remainingMessages := make([]models.SqsMessage, 0, len(models.SyncQueues.Queues[key].Messages)) + deletedEntries := make([]models.DeleteMessageBatchResultEntry, 0) + remainingMessages := make([]models.SqsMessage, 0, len(models.SyncQueues.Queues[key].Messages)) -for _, message := range models.SyncQueues.Queues[key].Messages { -if de, found := deleteMessageMap[message.ReceiptHandle]; found { -log.Debugf("FIFO Queue %s unlocking group %s:", queueName, message.GroupID) -models.SyncQueues.Queues[key].UnlockGroup(message.GroupID) -delete(models.SyncQueues.Queues[key].Duplicates, message.DeduplicationID) -de.Deleted = true -deletedEntries = append(deletedEntries, models.DeleteMessageBatchResultEntry{Id: de.Id}) -} else { -remainingMessages = append(remainingMessages, message) -} -} + for _, message := range models.SyncQueues.Queues[key].Messages { + if de, found := deleteMessageMap[message.ReceiptHandle]; found { + log.Debugf("FIFO Queue %s unlocking group %s:", queueName, message.GroupID) + models.SyncQueues.Queues[key].UnlockGroup(message.GroupID) + delete(models.SyncQueues.Queues[key].Duplicates, message.DeduplicationID) + de.Deleted = true + deletedEntries = append(deletedEntries, models.DeleteMessageBatchResultEntry{Id: de.Id}) + } else { + remainingMessages = append(remainingMessages, message) + } + } -models.SyncQueues.Queues[key].Messages = remainingMessages + models.SyncQueues.Queues[key].Messages = remainingMessages + // Персистим обновлённое состояние очереди, чтобы не терять batch-delete после рестарта. + persistence.SaveQueue(key, models.SyncQueues.Queues[key]) -notFoundEntries := make([]models.BatchResultErrorEntry, 0) -for _, de := range deleteMessageMap { -if !de.Deleted { -notFoundEntries = append(notFoundEntries, models.BatchResultErrorEntry{ -Code: "1", -Id: de.Id, -Message: "Message not found", -SenderFault: true, -}) -} -} + notFoundEntries := make([]models.BatchResultErrorEntry, 0) + for _, de := range deleteMessageMap { + if !de.Deleted { + notFoundEntries = append(notFoundEntries, models.BatchResultErrorEntry{ + Code: "1", + Id: de.Id, + Message: "Message not found", + SenderFault: true, + }) + } + } -respStruct := models.DeleteMessageBatchResponse{ -Xmlns: models.BaseXmlns, -Result: models.DeleteMessageBatchResult{ -Successful: deletedEntries, -Failed: notFoundEntries, -}, -Metadata: models.BaseResponseMetadata, -} + respStruct := models.DeleteMessageBatchResponse{ + Xmlns: models.BaseXmlns, + Result: models.DeleteMessageBatchResult{ + Successful: deletedEntries, + Failed: notFoundEntries, + }, + Metadata: models.BaseResponseMetadata, + } -return http.StatusOK, respStruct + return http.StatusOK, respStruct } type deleteEntry struct { -Id string -ReceiptHandle string -Error string -Deleted bool + Id string + ReceiptHandle string + Error string + Deleted bool } diff --git a/app/gosqs/get_queue_url.go b/app/gosqs/get_queue_url.go index e932ec0..090bfbf 100644 --- a/app/gosqs/get_queue_url.go +++ b/app/gosqs/get_queue_url.go @@ -3,44 +3,45 @@ package gosqs import ( -"net/http" + "net/http" -"shared-sqs/app/interfaces" -"shared-sqs/app/models" -"shared-sqs/app/utils" -log "github.com/sirupsen/logrus" + "shared-sqs/app/interfaces" + "shared-sqs/app/models" + "shared-sqs/app/utils" + + log "github.com/sirupsen/logrus" ) func GetQueueUrlV1(req *http.Request) (int, interfaces.AbstractResponseBody) { -requestBody := models.NewGetQueueUrlRequest() -ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) -if !ok { -log.Error("Invalid Request - GetQueueUrlV1") -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} + requestBody := models.NewGetQueueUrlRequest() + ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) + if !ok { + log.Error("Invalid Request - GetQueueUrlV1") + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } -t := getTenantFromContext(req) -if t == nil { -return utils.CreateErrorResponseV1("InvalidClientTokenId", true) -} + t := getTenantFromContext(req) + if t == nil { + return utils.CreateErrorResponseV1("InvalidClientTokenId", true) + } -queueName := requestBody.QueueName -key := tenantQueueKey(t.AccessKey, queueName) + queueName := requestBody.QueueName + key := tenantQueueKey(t.AccessKey, queueName) -// Fix #10: RLock перед чтением SyncQueues — иначе data race -models.SyncQueues.RLock() -queue, ok := models.SyncQueues.Queues[key] -models.SyncQueues.RUnlock() -if !ok { -log.Errorf("Get Queue URL: %s, queue does not exist for tenant %s", queueName, t.ID) -return utils.CreateErrorResponseV1("QueueNotFound", true) -} -log.Debug("Get Queue URL:", queue.Name) + // Fix #10: RLock перед чтением SyncQueues — иначе data race + models.SyncQueues.RLock() + queue, ok := models.SyncQueues.Queues[key] + models.SyncQueues.RUnlock() + if !ok { + log.Errorf("Get Queue URL: %s, queue does not exist for tenant %s", queueName, t.ID) + return utils.CreateErrorResponseV1("QueueNotFound", true) + } + log.Debug("Get Queue URL:", queue.Name) -respStruct := models.GetQueueUrlResponse{ -Xmlns: models.BaseXmlns, -Result: models.GetQueueUrlResult{QueueUrl: queue.URL}, -Metadata: models.BaseResponseMetadata, -} -return http.StatusOK, respStruct + respStruct := models.GetQueueUrlResponse{ + Xmlns: models.BaseXmlns, + Result: models.GetQueueUrlResult{QueueUrl: queue.URL}, + Metadata: models.BaseResponseMetadata, + } + return http.StatusOK, respStruct } diff --git a/app/gosqs/list_queues.go b/app/gosqs/list_queues.go index 8bb098f..b6cfa1d 100644 --- a/app/gosqs/list_queues.go +++ b/app/gosqs/list_queues.go @@ -4,48 +4,48 @@ package gosqs import ( -"net/http" -"strings" + "net/http" + "strings" -"shared-sqs/app/interfaces" -"shared-sqs/app/models" -"shared-sqs/app/utils" -log "github.com/sirupsen/logrus" + log "github.com/sirupsen/logrus" + "shared-sqs/app/interfaces" + "shared-sqs/app/models" + "shared-sqs/app/utils" ) func ListQueuesV1(req *http.Request) (int, interfaces.AbstractResponseBody) { -requestBody := models.NewListQueuesRequest() -ok := utils.REQUEST_TRANSFORMER(requestBody, req, true) -if !ok { -log.Error("Invalid Request - ListQueuesV1") -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} + requestBody := models.NewListQueuesRequest() + ok := utils.REQUEST_TRANSFORMER(requestBody, req, true) + if !ok { + log.Error("Invalid Request - ListQueuesV1") + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } -t := getTenantFromContext(req) -if t == nil { -return utils.CreateErrorResponseV1("InvalidClientTokenId", true) -} + t := getTenantFromContext(req) + if t == nil { + return utils.CreateErrorResponseV1("InvalidClientTokenId", true) + } -log.Infof("Listing Queues for tenant: %s", t.ID) -queueUrls := make([]string, 0) -prefix := t.AccessKey + ":" + log.Infof("Listing Queues for tenant: %s", t.ID) + queueUrls := make([]string, 0) + prefix := t.AccessKey + ":" -models.SyncQueues.Lock() -for key, queue := range models.SyncQueues.Queues { -// Показываем только очереди этого тенанта -if strings.HasPrefix(key, prefix) { -if strings.HasPrefix(queue.Name, requestBody.QueueNamePrefix) { -queueUrls = append(queueUrls, queue.URL) -} -} -} -models.SyncQueues.Unlock() + models.SyncQueues.RLock() + for key, queue := range models.SyncQueues.Queues { + // Показываем только очереди этого тенанта + if strings.HasPrefix(key, prefix) { + if strings.HasPrefix(queue.Name, requestBody.QueueNamePrefix) { + queueUrls = append(queueUrls, queue.URL) + } + } + } + models.SyncQueues.RUnlock() -respStruct := models.ListQueuesResponse{ -Xmlns: models.BaseXmlns, -Metadata: models.BaseResponseMetadata, -Result: models.ListQueuesResult{QueueUrls: queueUrls}, -} + respStruct := models.ListQueuesResponse{ + Xmlns: models.BaseXmlns, + Metadata: models.BaseResponseMetadata, + Result: models.ListQueuesResult{QueueUrls: queueUrls}, + } -return http.StatusOK, respStruct + return http.StatusOK, respStruct } diff --git a/app/gosqs/receive_message.go b/app/gosqs/receive_message.go index d25c9b4..f083d12 100644 --- a/app/gosqs/receive_message.go +++ b/app/gosqs/receive_message.go @@ -4,169 +4,170 @@ package gosqs import ( -"fmt" -"net/http" -"strings" -"time" + "fmt" + "net/http" + "strings" + "time" -"github.com/google/uuid" + "github.com/google/uuid" -"shared-sqs/app/interfaces" -"shared-sqs/app/models" -"shared-sqs/app/utils" -"github.com/gorilla/mux" -log "github.com/sirupsen/logrus" + "shared-sqs/app/interfaces" + "shared-sqs/app/models" + "shared-sqs/app/utils" + + "github.com/gorilla/mux" + log "github.com/sirupsen/logrus" ) func ReceiveMessageV1(req *http.Request) (int, interfaces.AbstractResponseBody) { -requestBody := models.NewReceiveMessageRequest() -ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) -if !ok { -log.Error("Invalid Request - ReceiveMessageV1") -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} + requestBody := models.NewReceiveMessageRequest() + ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) + if !ok { + log.Error("Invalid Request - ReceiveMessageV1") + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } -t := getTenantFromContext(req) -if t == nil { -return utils.CreateErrorResponseV1("InvalidClientTokenId", true) -} + t := getTenantFromContext(req) + if t == nil { + return utils.CreateErrorResponseV1("InvalidClientTokenId", true) + } -maxNumberOfMessages := requestBody.MaxNumberOfMessages -if maxNumberOfMessages == 0 { -maxNumberOfMessages = 1 -} -// Fix #8: clamp MaxNumberOfMessages к AWS лимиту 1–10 -maxNumberOfMessages = ClampInt(maxNumberOfMessages, MinNumberOfMessagesLimit, MaxNumberOfMessagesLimit) + maxNumberOfMessages := requestBody.MaxNumberOfMessages + if maxNumberOfMessages == 0 { + maxNumberOfMessages = 1 + } + // Fix #8: clamp MaxNumberOfMessages к AWS лимиту 1–10 + maxNumberOfMessages = ClampInt(maxNumberOfMessages, MinNumberOfMessagesLimit, MaxNumberOfMessagesLimit) -queueName := "" -if requestBody.QueueUrl == "" { -vars := mux.Vars(req) -queueName = vars["queueName"] -} else { -uriSegments := strings.Split(requestBody.QueueUrl, "/") -queueName = uriSegments[len(uriSegments)-1] -} + queueName := "" + if requestBody.QueueUrl == "" { + vars := mux.Vars(req) + queueName = vars["queueName"] + } else { + uriSegments := strings.Split(requestBody.QueueUrl, "/") + queueName = uriSegments[len(uriSegments)-1] + } -key := tenantQueueKey(t.AccessKey, queueName) + key := tenantQueueKey(t.AccessKey, queueName) -if _, ok := models.SyncQueues.Queues[key]; !ok { -return utils.CreateErrorResponseV1("QueueNotFound", true) -} + if _, ok := models.SyncQueues.Queues[key]; !ok { + return utils.CreateErrorResponseV1("QueueNotFound", true) + } -var messages []*models.ResultMessage -respStruct := models.ReceiveMessageResponse{} + var messages []*models.ResultMessage + respStruct := models.ReceiveMessageResponse{} -waitTimeSeconds := requestBody.WaitTimeSeconds -if waitTimeSeconds == 0 { -models.SyncQueues.RLock() -waitTimeSeconds = models.SyncQueues.Queues[key].ReceiveMessageWaitTimeSeconds -models.SyncQueues.RUnlock() -} -// Fix #4: clamp WaitTimeSeconds к AWS лимиту 0–20 -waitTimeSeconds = ClampInt(waitTimeSeconds, 0, MaxReceiveMessageWaitTimeSeconds) + waitTimeSeconds := requestBody.WaitTimeSeconds + if waitTimeSeconds == 0 { + models.SyncQueues.RLock() + waitTimeSeconds = models.SyncQueues.Queues[key].ReceiveMessageWaitTimeSeconds + models.SyncQueues.RUnlock() + } + // Fix #4: clamp WaitTimeSeconds к AWS лимиту 0–20 + waitTimeSeconds = ClampInt(waitTimeSeconds, 0, MaxReceiveMessageWaitTimeSeconds) -// Long polling: ждём появления сообщения до waitTimeSeconds*10 итераций по 100ms -loops := waitTimeSeconds * 10 -for loops > 0 { -models.SyncQueues.RLock() -_, queueFound := models.SyncQueues.Queues[key] -if !queueFound { -models.SyncQueues.RUnlock() -return utils.CreateErrorResponseV1("QueueNotFound", true) -} -messageFound := len(models.SyncQueues.Queues[key].Messages)-numberOfHiddenMessagesInQueue(*models.SyncQueues.Queues[key]) != 0 -models.SyncQueues.RUnlock() -if !messageFound { -continueTimer := time.NewTimer(100 * time.Millisecond) -select { -case <-req.Context().Done(): -continueTimer.Stop() -return http.StatusOK, models.ReceiveMessageResponse{ -Xmlns: models.BaseXmlns, -Result: models.ReceiveMessageResult{}, -Metadata: models.BaseResponseMetadata, -} -case <-continueTimer.C: -continueTimer.Stop() -} -loops-- -} else { -break -} -} -log.Debugf("Getting Message from Queue:%s (tenant: %s)", queueName, t.ID) + // Long polling: ждём появления сообщения до waitTimeSeconds*10 итераций по 100ms + loops := waitTimeSeconds * 10 + for loops > 0 { + models.SyncQueues.RLock() + _, queueFound := models.SyncQueues.Queues[key] + if !queueFound { + models.SyncQueues.RUnlock() + return utils.CreateErrorResponseV1("QueueNotFound", true) + } + messageFound := len(models.SyncQueues.Queues[key].Messages)-numberOfHiddenMessagesInQueue(*models.SyncQueues.Queues[key]) != 0 + models.SyncQueues.RUnlock() + if !messageFound { + continueTimer := time.NewTimer(100 * time.Millisecond) + select { + case <-req.Context().Done(): + continueTimer.Stop() + return http.StatusOK, models.ReceiveMessageResponse{ + Xmlns: models.BaseXmlns, + Result: models.ReceiveMessageResult{}, + Metadata: models.BaseResponseMetadata, + } + case <-continueTimer.C: + continueTimer.Stop() + } + loops-- + } else { + break + } + } + log.Debugf("Getting Message from Queue:%s (tenant: %s)", queueName, t.ID) -models.SyncQueues.Lock() -defer models.SyncQueues.Unlock() + models.SyncQueues.Lock() + defer models.SyncQueues.Unlock() -if len(models.SyncQueues.Queues[key].Messages) > 0 { -numMsg := 0 -messages = make([]*models.ResultMessage, 0) -for i := range models.SyncQueues.Queues[key].Messages { -if numMsg >= maxNumberOfMessages { -break -} + if len(models.SyncQueues.Queues[key].Messages) > 0 { + numMsg := 0 + messages = make([]*models.ResultMessage, 0) + for i := range models.SyncQueues.Queues[key].Messages { + if numMsg >= maxNumberOfMessages { + break + } -if models.SyncQueues.Queues[key].Messages[i].ReceiptHandle != "" { -continue -} + if models.SyncQueues.Queues[key].Messages[i].ReceiptHandle != "" { + continue + } -msg := &models.SyncQueues.Queues[key].Messages[i] -if !msg.IsReadyForReceipt() { -continue -} + msg := &models.SyncQueues.Queues[key].Messages[i] + if !msg.IsReadyForReceipt() { + continue + } -if models.SyncQueues.Queues[key].IsFIFO { -if models.SyncQueues.Queues[key].IsLocked(msg.GroupID) { -continue -} -models.SyncQueues.Queues[key].LockGroup(msg.GroupID) -} + if models.SyncQueues.Queues[key].IsFIFO { + if models.SyncQueues.Queues[key].IsLocked(msg.GroupID) { + continue + } + models.SyncQueues.Queues[key].LockGroup(msg.GroupID) + } -randomId := uuid.NewString() -msg.ReceiptHandle = msg.Uuid + "#" + randomId -msg.ReceiptTime = time.Now().UTC() + randomId := uuid.NewString() + msg.ReceiptHandle = msg.Uuid + "#" + randomId + msg.ReceiptTime = time.Now().UTC() -if requestBody.VisibilityTimeout != 0 { -msg.VisibilityTimeout = time.Now().Add(time.Duration(requestBody.VisibilityTimeout) * time.Second) -} else { -msg.VisibilityTimeout = time.Now().Add(time.Duration(models.SyncQueues.Queues[key].VisibilityTimeout) * time.Second) -} + if requestBody.VisibilityTimeout != 0 { + msg.VisibilityTimeout = time.Now().Add(time.Duration(requestBody.VisibilityTimeout) * time.Second) + } else { + msg.VisibilityTimeout = time.Now().Add(time.Duration(models.SyncQueues.Queues[key].VisibilityTimeout) * time.Second) + } -messages = append(messages, buildResultMessage(msg)) -numMsg++ -} + messages = append(messages, buildResultMessage(msg)) + numMsg++ + } -respStruct = models.ReceiveMessageResponse{ -"http://queue.amazonaws.com/doc/2012-11-05/", -models.ReceiveMessageResult{Messages: messages}, -models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"}, -} -} else { -log.Warning("No messages in Queue:", queueName) -respStruct = models.ReceiveMessageResponse{ -Xmlns: "http://queue.amazonaws.com/doc/2012-11-05/", -Result: models.ReceiveMessageResult{}, -Metadata: models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"}, -} -} + respStruct = models.ReceiveMessageResponse{ + "http://queue.amazonaws.com/doc/2012-11-05/", + models.ReceiveMessageResult{Messages: messages}, + models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"}, + } + } else { + log.Warning("No messages in Queue:", queueName) + respStruct = models.ReceiveMessageResponse{ + Xmlns: "http://queue.amazonaws.com/doc/2012-11-05/", + Result: models.ReceiveMessageResult{}, + Metadata: models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"}, + } + } -return http.StatusOK, respStruct + return http.StatusOK, respStruct } func buildResultMessage(m *models.SqsMessage) *models.ResultMessage { -return &models.ResultMessage{ -MessageId: m.Uuid, -Body: m.MessageBody, -ReceiptHandle: m.ReceiptHandle, -MD5OfBody: utils.GetMD5Hash(m.MessageBody), -MD5OfMessageAttributes: m.MD5OfMessageAttributes, -MessageAttributes: m.MessageAttributes, -Attributes: map[string]string{ -"ApproximateFirstReceiveTimestamp": fmt.Sprintf("%d", m.ReceiptTime.UnixNano()/int64(time.Millisecond)), -"SenderId": models.CurrentEnvironment.AccountID, -"ApproximateReceiveCount": fmt.Sprintf("%d", m.NumberOfReceives+1), -"SentTimestamp": fmt.Sprintf("%d", time.Now().UTC().UnixNano()/int64(time.Millisecond)), -}, -} + return &models.ResultMessage{ + MessageId: m.Uuid, + Body: m.MessageBody, + ReceiptHandle: m.ReceiptHandle, + MD5OfBody: utils.GetMD5Hash(m.MessageBody), + MD5OfMessageAttributes: m.MD5OfMessageAttributes, + MessageAttributes: m.MessageAttributes, + Attributes: map[string]string{ + "ApproximateFirstReceiveTimestamp": fmt.Sprintf("%d", m.ReceiptTime.UnixNano()/int64(time.Millisecond)), + "SenderId": models.CurrentEnvironment.AccountID, + "ApproximateReceiveCount": fmt.Sprintf("%d", m.NumberOfReceives+1), + "SentTimestamp": fmt.Sprintf("%d", time.Now().UTC().UnixNano()/int64(time.Millisecond)), + }, + } } diff --git a/app/gosqs/send_message.go b/app/gosqs/send_message.go index cee626b..196c794 100644 --- a/app/gosqs/send_message.go +++ b/app/gosqs/send_message.go @@ -63,25 +63,27 @@ func SendMessageV1(req *http.Request) (int, interfaces.AbstractResponseBody) { key := tenantQueueKey(t.AccessKey, queueName) - if _, ok := models.SyncQueues.Queues[key]; !ok { + models.SyncQueues.RLock() + queue, exists := models.SyncQueues.Queues[key] + if !exists { + models.SyncQueues.RUnlock() return utils.CreateErrorResponseV1("QueueNotFound", true) } + maxMessageSize := queue.MaximumMessageSize + currentMsgCount := len(queue.Messages) + queueIsFIFO := queue.IsFIFO + delaySecs := queue.DelaySeconds + models.SyncQueues.RUnlock() - if models.SyncQueues.Queues[key].MaximumMessageSize > 0 && - len(messageBody) > models.SyncQueues.Queues[key].MaximumMessageSize { + if maxMessageSize > 0 && len(messageBody) > maxMessageSize { return utils.CreateErrorResponseV1("MessageTooBig", true) } // Fix #11: лимит сообщений в очереди — защита от OOM - models.SyncQueues.RLock() - currentMsgCount := len(models.SyncQueues.Queues[key].Messages) - queueIsFIFO := models.SyncQueues.Queues[key].IsFIFO - models.SyncQueues.RUnlock() if currentMsgCount >= MaxMessagesForQueue(queueIsFIFO) { return utils.CreateErrorResponseV1("OverLimit", true) } - delaySecs := models.SyncQueues.Queues[key].DelaySeconds if requestBody.DelaySeconds != 0 { delaySecs = requestBody.DelaySeconds } diff --git a/app/gosqs/send_message_batch.go b/app/gosqs/send_message_batch.go index f7953aa..4c31702 100644 --- a/app/gosqs/send_message_batch.go +++ b/app/gosqs/send_message_batch.go @@ -3,143 +3,150 @@ package gosqs import ( -"net/http" -"strings" -"time" + "net/http" + "strings" + "time" -"github.com/google/uuid" + "github.com/google/uuid" -"shared-sqs/app/interfaces" -"shared-sqs/app/models" -"shared-sqs/app/utils" -"github.com/gorilla/mux" -log "github.com/sirupsen/logrus" + "shared-sqs/app/interfaces" + "shared-sqs/app/models" + "shared-sqs/app/persistence" + "shared-sqs/app/utils" + + "github.com/gorilla/mux" + log "github.com/sirupsen/logrus" ) func SendMessageBatchV1(req *http.Request) (int, interfaces.AbstractResponseBody) { -requestBody := models.NewSendMessageBatchRequest() -ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) -if !ok { -log.Error("Invalid Request - SendMessageBatchV1") -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} + requestBody := models.NewSendMessageBatchRequest() + ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) + if !ok { + log.Error("Invalid Request - SendMessageBatchV1") + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } -t := getTenantFromContext(req) -if t == nil { -return utils.CreateErrorResponseV1("InvalidClientTokenId", true) -} + t := getTenantFromContext(req) + if t == nil { + return utils.CreateErrorResponseV1("InvalidClientTokenId", true) + } -queueUrl := requestBody.QueueUrl -queueName := "" -if queueUrl == "" { -vars := mux.Vars(req) -queueName = vars["queueName"] -} else { -uriSegments := strings.Split(queueUrl, "/") -queueName = uriSegments[len(uriSegments)-1] -} + queueUrl := requestBody.QueueUrl + queueName := "" + if queueUrl == "" { + vars := mux.Vars(req) + queueName = vars["queueName"] + } else { + uriSegments := strings.Split(queueUrl, "/") + queueName = uriSegments[len(uriSegments)-1] + } -key := tenantQueueKey(t.AccessKey, queueName) + key := tenantQueueKey(t.AccessKey, queueName) -if _, ok := models.SyncQueues.Queues[key]; !ok { -return utils.CreateErrorResponseV1("QueueNotFound", true) -} + models.SyncQueues.RLock() + queue, exists := models.SyncQueues.Queues[key] + if !exists { + models.SyncQueues.RUnlock() + return utils.CreateErrorResponseV1("QueueNotFound", true) + } + maxMsgSize := queue.MaximumMessageSize + currentMsgCount := len(queue.Messages) + queueIsFIFO := queue.IsFIFO + models.SyncQueues.RUnlock() -sendEntries := requestBody.Entries + sendEntries := requestBody.Entries -if len(sendEntries) == 0 { -return utils.CreateErrorResponseV1("EmptyBatchRequest", true) -} + if len(sendEntries) == 0 { + return utils.CreateErrorResponseV1("EmptyBatchRequest", true) + } -if len(sendEntries) > 10 { -return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true) -} -ids := map[string]struct{}{} -for _, v := range sendEntries { -if _, ok := ids[v.Id]; ok { -return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true) -} -// Валидация длины BatchEntryId (макс 80 chars) -if err := ValidateBatchEntryID(v.Id); err != nil { -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} -// Валидация DeduplicationID и GroupID (макс 128 chars) -if err := ValidateDeduplicationID(v.MessageDeduplicationId); err != nil { -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} -if err := ValidateGroupID(v.MessageGroupId); err != nil { -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} -ids[v.Id] = struct{}{} -} + if len(sendEntries) > 10 { + return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true) + } + ids := map[string]struct{}{} + for _, v := range sendEntries { + if _, ok := ids[v.Id]; ok { + return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true) + } + // Валидация длины BatchEntryId (макс 80 chars) + if err := ValidateBatchEntryID(v.Id); err != nil { + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } + // Валидация DeduplicationID и GroupID (макс 128 chars) + if err := ValidateDeduplicationID(v.MessageDeduplicationId); err != nil { + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } + if err := ValidateGroupID(v.MessageGroupId); err != nil { + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } + ids[v.Id] = struct{}{} + } -sentEntries := make([]models.SendMessageBatchResultEntry, 0) -log.Debugf("Batch sending to Queue: %s (tenant: %s)", queueName, t.ID) + sentEntries := make([]models.SendMessageBatchResultEntry, 0) + log.Debugf("Batch sending to Queue: %s (tenant: %s)", queueName, t.ID) -// Fix #2: проверяем размер каждого сообщения в batch (Critical — batch size bypass) -maxMsgSize := models.SyncQueues.Queues[key].MaximumMessageSize -if maxMsgSize <= 0 { -maxMsgSize = MaxMessageSizeDefault -} -for _, entry := range sendEntries { -if len(entry.MessageBody) > maxMsgSize { -return utils.CreateErrorResponseV1("MessageTooBig", true) -} -// Валидация количества message attributes (макс 10 по AWS) -if len(entry.MessageAttributes) > MaxMessageAttributes { -return utils.CreateErrorResponseV1("InvalidParameterValue", true) -} -} + // Fix #2: проверяем размер каждого сообщения в batch (Critical — batch size bypass) + if maxMsgSize <= 0 { + maxMsgSize = MaxMessageSizeDefault + } + for _, entry := range sendEntries { + if len(entry.MessageBody) > maxMsgSize { + return utils.CreateErrorResponseV1("MessageTooBig", true) + } + // Валидация количества message attributes (макс 10 по AWS) + if len(entry.MessageAttributes) > MaxMessageAttributes { + return utils.CreateErrorResponseV1("InvalidParameterValue", true) + } + } -// Fix #11: лимит сообщений в очереди — защита от OOM -models.SyncQueues.RLock() -currentMsgCount := len(models.SyncQueues.Queues[key].Messages) -queueIsFIFO := models.SyncQueues.Queues[key].IsFIFO -models.SyncQueues.RUnlock() -if currentMsgCount+len(sendEntries) > MaxMessagesForQueue(queueIsFIFO) { -return utils.CreateErrorResponseV1("OverLimit", true) -} + // Fix #11: лимит сообщений в очереди — защита от OOM + if currentMsgCount+len(sendEntries) > MaxMessagesForQueue(queueIsFIFO) { + return utils.CreateErrorResponseV1("OverLimit", true) + } -for _, sendEntry := range sendEntries { -msg := models.SqsMessage{MessageBody: sendEntry.MessageBody} -if len(sendEntry.MessageAttributes) > 0 { -msg.MessageAttributes = sendEntry.MessageAttributes -msg.MD5OfMessageAttributes = utils.HashAttributes(sendEntry.MessageAttributes) -} -msg.MD5OfMessageBody = utils.GetMD5Hash(sendEntry.MessageBody) -msg.GroupID = sendEntry.MessageGroupId -msg.DeduplicationID = sendEntry.MessageDeduplicationId -msg.Uuid = uuid.NewString() -msg.SentTime = time.Now() + models.SyncQueues.Lock() + queue = models.SyncQueues.Queues[key] + for _, sendEntry := range sendEntries { + msg := models.SqsMessage{MessageBody: sendEntry.MessageBody} + if len(sendEntry.MessageAttributes) > 0 { + msg.MessageAttributes = sendEntry.MessageAttributes + msg.MD5OfMessageAttributes = utils.HashAttributes(sendEntry.MessageAttributes) + } + msg.MD5OfMessageBody = utils.GetMD5Hash(sendEntry.MessageBody) + msg.GroupID = sendEntry.MessageGroupId + msg.DeduplicationID = sendEntry.MessageDeduplicationId + msg.Uuid = uuid.NewString() + msg.SentTime = time.Now() -models.SyncQueues.Lock() -fifoSeqNumber := "" -if models.SyncQueues.Queues[key].IsFIFO { -fifoSeqNumber = models.SyncQueues.Queues[key].NextSequenceNumber(sendEntry.MessageGroupId) -} -if !models.SyncQueues.Queues[key].IsDuplicate(sendEntry.MessageDeduplicationId) { -models.SyncQueues.Queues[key].Messages = append(models.SyncQueues.Queues[key].Messages, msg) -} else { -log.Debugf("Duplicate deduplicationId [%s] in queue [%s]", sendEntry.MessageDeduplicationId, queueName) -} -models.SyncQueues.Queues[key].InitDuplicatation(sendEntry.MessageDeduplicationId) -models.SyncQueues.Unlock() + fifoSeqNumber := "" + if queue.IsFIFO { + fifoSeqNumber = queue.NextSequenceNumber(sendEntry.MessageGroupId) + } + if !queue.IsDuplicate(sendEntry.MessageDeduplicationId) { + queue.Messages = append(queue.Messages, msg) + } else { + log.Debugf("Duplicate deduplicationId [%s] in queue [%s]", sendEntry.MessageDeduplicationId, queueName) + } + queue.InitDuplicatation(sendEntry.MessageDeduplicationId) -sentEntries = append(sentEntries, models.SendMessageBatchResultEntry{ -Id: sendEntry.Id, -MessageId: msg.Uuid, -MD5OfMessageBody: msg.MD5OfMessageBody, -MD5OfMessageAttributes: msg.MD5OfMessageAttributes, -SequenceNumber: fifoSeqNumber, -}) -log.Infof("%s: Queue: %s, Message: %s", time.Now().Format("2006-01-02 15:04:05"), queueName, msg.MessageBody) -} + sentEntries = append(sentEntries, models.SendMessageBatchResultEntry{ + Id: sendEntry.Id, + MessageId: msg.Uuid, + MD5OfMessageBody: msg.MD5OfMessageBody, + MD5OfMessageAttributes: msg.MD5OfMessageAttributes, + SequenceNumber: fifoSeqNumber, + }) + log.Infof("%s: Queue: %s, Message: %s", time.Now().Format("2006-01-02 15:04:05"), queueName, msg.MessageBody) + } + // Персистим батч-изменение одним снапшотом под lock. + persistence.SaveQueue(key, queue) + models.SyncQueues.Unlock() -respStruct := models.SendMessageBatchResponse{ -Xmlns: models.BaseXmlns, -Result: models.SendMessageBatchResult{Entry: sentEntries}, -Metadata: models.BaseResponseMetadata, -} + respStruct := models.SendMessageBatchResponse{ + Xmlns: models.BaseXmlns, + Result: models.SendMessageBatchResult{Entry: sentEntries}, + Metadata: models.BaseResponseMetadata, + } -return http.StatusOK, respStruct + return http.StatusOK, respStruct } diff --git a/app/gosqs/validation.go b/app/gosqs/validation.go index 0c07075..6e93526 100644 --- a/app/gosqs/validation.go +++ b/app/gosqs/validation.go @@ -12,16 +12,16 @@ import ( // ── AWS SQS лимиты ──────────────────────────────────────────────────────────── const ( // Очередь - MaxQueueNameLength = 80 + MaxQueueNameLength = 80 MaxMessageSizeDefault = 262144 // 256 KB MaxMessageSizeLimit = 262144 - MinMessageSizeLimit = 1024 // 1 KB + MinMessageSizeLimit = 1024 // 1 KB // Атрибуты очереди - MaxDelaySeconds = 900 // 15 min - MaxVisibilityTimeout = 43200 // 12 hours - MaxReceiveMessageWaitTimeSeconds = 20 // long polling cap - MinMessageRetentionPeriod = 60 // 1 min + MaxDelaySeconds = 900 // 15 min + MaxVisibilityTimeout = 43200 // 12 hours + MaxReceiveMessageWaitTimeSeconds = 20 // long polling cap + MinMessageRetentionPeriod = 60 // 1 min MaxMessageRetentionPeriod = 1209600 // 14 days // Receive @@ -29,7 +29,7 @@ const ( MinNumberOfMessagesLimit = 1 // Message attributes - MaxMessageAttributes = 10 + MaxMessageAttributes = 10 MaxMessageAttributeSize = 262144 // 256 KB суммарно (тело + атрибуты) // Deduplication / GroupID diff --git a/app/router/router.go b/app/router/router.go index aa2c636..48f17d4 100644 --- a/app/router/router.go +++ b/app/router/router.go @@ -147,7 +147,11 @@ func extractAction(req *http.Request) string { switch protocol { case AwsJsonProtocol: action := req.Header.Get("X-Amz-Target") - return strings.Split(action, ".")[1] + parts := strings.SplitN(action, ".", 2) + if len(parts) != 2 || parts[1] == "" { + return "" + } + return parts[1] case AwsQueryProtocol: return req.FormValue("Action") } diff --git a/app/tenant/tenant_store.go b/app/tenant/tenant_store.go index d87458f..96635a2 100644 --- a/app/tenant/tenant_store.go +++ b/app/tenant/tenant_store.go @@ -202,6 +202,10 @@ func (s *TenantStore) CreateFromJWT(tenantID, sub, email string, maxQueues int) s.mu.Unlock() return existing, nil } + if len(s.byID) >= MaxTenantsGlobal { + s.mu.Unlock() + return nil, fmt.Errorf("global tenant limit reached (%d)", MaxTenantsGlobal) + } accessKey, err := generateAccessKey() if err != nil {