security: fix critical/high auth, idor, races and persistence

This commit is contained in:
Naeel
2026-04-10 19:41:59 +03:00
parent e9a26f7975
commit a5e9bfb15c
11 changed files with 868 additions and 505 deletions
+20 -4
View File
@@ -14,6 +14,7 @@ import (
"shared-sqs/app/auth" "shared-sqs/app/auth"
"shared-sqs/app/models" "shared-sqs/app/models"
"shared-sqs/app/persistence"
"shared-sqs/app/tenant" "shared-sqs/app/tenant"
"github.com/google/uuid" "github.com/google/uuid"
@@ -36,9 +37,9 @@ func findQueue(tenantAccessKey, queueName string) (string, *models.Queue) {
// Handler — admin API handler, holds TenantStore и admin token // Handler — admin API handler, holds TenantStore и admin token
type Handler struct { type Handler struct {
store *tenant.TenantStore store *tenant.TenantStore
adminToken string adminToken string
nubesEndpoint string // URL nubes API для валидации JWT (напр. https://deck-api-test.ngcloud.ru/api/v1) nubesEndpoint string // URL nubes API для валидации JWT (напр. https://deck-api-test.ngcloud.ru/api/v1)
} }
// NewHandler — создаёт admin handler // NewHandler — создаёт admin handler
@@ -188,12 +189,18 @@ func (h *Handler) jwtMiddleware(next http.Handler) http.Handler {
} }
// Проверяем что тенант существует (был создан при /ui/api/auth) // Проверяем что тенант существует (был создан при /ui/api/auth)
_, ok := h.store.GetBySub(claims.Sub) jwtTenant, ok := h.store.GetBySub(claims.Sub)
if !ok { if !ok {
jsonErr(w, http.StatusForbidden, "tenant not found — authenticate first via POST /ui/api/auth") jsonErr(w, http.StatusForbidden, "tenant not found — authenticate first via POST /ui/api/auth")
return 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) next.ServeHTTP(w, r)
}) })
} }
@@ -315,13 +322,18 @@ func (h *Handler) deleteTenant(w http.ResponseWriter, r *http.Request) {
} }
// Удаляем все очереди тенанта из SyncQueues // Удаляем все очереди тенанта из SyncQueues
prefix := t.AccessKey + ":" prefix := t.AccessKey + ":"
deletedQueueKeys := make([]string, 0)
models.SyncQueues.Lock() models.SyncQueues.Lock()
for key := range models.SyncQueues.Queues { for key := range models.SyncQueues.Queues {
if strings.HasPrefix(key, prefix) { if strings.HasPrefix(key, prefix) {
deletedQueueKeys = append(deletedQueueKeys, key)
delete(models.SyncQueues.Queues, key) delete(models.SyncQueues.Queues, key)
} }
} }
models.SyncQueues.Unlock() models.SyncQueues.Unlock()
for _, queueKey := range deletedQueueKeys {
persistence.DeleteQueue(queueKey)
}
h.store.Delete(id) h.store.Delete(id)
w.WriteHeader(http.StatusNoContent) w.WriteHeader(http.StatusNoContent)
@@ -427,6 +439,7 @@ func (h *Handler) createTenantQueue(w http.ResponseWriter, r *http.Request) {
Messages: []models.SqsMessage{}, Messages: []models.SqsMessage{},
Duplicates: make(map[string]time.Time), Duplicates: make(map[string]time.Time),
} }
persistence.SaveQueue(key, models.SyncQueues.Queues[key])
models.SyncQueues.Unlock() models.SyncQueues.Unlock()
log.Infof("admin: created queue %s for tenant %s", req.Name, t.ID) log.Infof("admin: created queue %s for tenant %s", req.Name, t.ID)
w.Header().Set("Content-Type", "application/json") 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) delete(models.SyncQueues.Queues, key)
models.SyncQueues.Unlock() models.SyncQueues.Unlock()
persistence.DeleteQueue(key)
log.Infof("admin: deleted queue %s for tenant %s", queueName, t.ID) log.Infof("admin: deleted queue %s for tenant %s", queueName, t.ID)
w.WriteHeader(http.StatusNoContent) w.WriteHeader(http.StatusNoContent)
} }
@@ -552,6 +566,7 @@ func (h *Handler) sendMessageToQueue(w http.ResponseWriter, r *http.Request) {
} }
models.SyncQueues.Lock() models.SyncQueues.Lock()
models.SyncQueues.Queues[key].Messages = append(models.SyncQueues.Queues[key].Messages, msg) models.SyncQueues.Queues[key].Messages = append(models.SyncQueues.Queues[key].Messages, msg)
persistence.SaveQueue(key, models.SyncQueues.Queues[key])
models.SyncQueues.Unlock() models.SyncQueues.Unlock()
w.Header().Set("Content-Type", "application/json") w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusCreated) w.WriteHeader(http.StatusCreated)
@@ -575,6 +590,7 @@ func (h *Handler) purgeQueue(w http.ResponseWriter, r *http.Request) {
} }
models.SyncQueues.Lock() models.SyncQueues.Lock()
models.SyncQueues.Queues[key].Messages = models.SyncQueues.Queues[key].Messages[:0] models.SyncQueues.Queues[key].Messages = models.SyncQueues.Queues[key].Messages[:0]
persistence.SaveQueue(key, models.SyncQueues.Queues[key])
models.SyncQueues.Unlock() models.SyncQueues.Unlock()
log.Infof("admin: purged queue %s for tenant %s", queueName, t.ID) log.Infof("admin: purged queue %s for tenant %s", queueName, t.ID)
w.WriteHeader(http.StatusNoContent) w.WriteHeader(http.StatusNoContent)
+394 -69
View File
@@ -4,12 +4,24 @@
package auth package auth
import ( import (
"context" "bytes"
"encoding/xml" "context"
"net/http" "crypto/hmac"
"strings" "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. // TenantContextKey — ключ для хранения тенанта в request context.
@@ -18,39 +30,49 @@ type contextKey string
const TenantContextKey contextKey = "tenant" const TenantContextKey contextKey = "tenant"
const (
sigV4DateFormat = "20060102T150405Z"
maxClockSkew = 15 * time.Minute
)
// AuthMiddleware — middleware: ищет тенанта по AccessKeyId из AWS Authorization header. // AuthMiddleware — middleware: ищет тенанта по AccessKeyId из AWS Authorization header.
// Пропускает /health и /admin/** без tenant-аутентификации. // Пропускает /health и /admin/** без tenant-аутентификации.
// Ловушка #4: не ставим короткий таймаут — ReceiveMessage с long polling держит соединение до 20 сек. // Ловушка #4: не ставим короткий таймаут — ReceiveMessage с long polling держит соединение до 20 сек.
func AuthMiddleware(store *tenant.TenantStore) func(http.Handler) http.Handler { func AuthMiddleware(store *tenant.TenantStore) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler { return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// /health — без auth // /health — без auth
if r.URL.Path == "/health" { if r.URL.Path == "/health" {
next.ServeHTTP(w, r) next.ServeHTTP(w, r)
return return
} }
// /admin/** — отдельная auth (bearer token, см. admin_handlers.go) // /admin/** — отдельная auth (bearer token, см. admin_handlers.go)
if strings.HasPrefix(r.URL.Path, "/admin/") { if strings.HasPrefix(r.URL.Path, "/admin/") {
next.ServeHTTP(w, r) next.ServeHTTP(w, r)
return return
} }
accessKeyID := extractAccessKeyID(r) accessKeyID := extractAccessKeyID(r)
if accessKeyID == "" { if accessKeyID == "" {
writeSQSAuthError(w, "MissingAuthenticationToken", "Request must contain either AccessKeyId or X-Amz-Credential") writeSQSAuthError(w, "MissingAuthenticationToken", "Request must contain either AccessKeyId or X-Amz-Credential")
return return
} }
t, ok := store.GetByAccessKey(accessKeyID) t, ok := store.GetByAccessKey(accessKeyID)
if !ok || !t.Active { if !ok || !t.Active {
writeSQSAuthError(w, "InvalidClientTokenId", "The security token included in the request is invalid") writeSQSAuthError(w, "InvalidClientTokenId", "The security token included in the request is invalid")
return return
} }
ctx := context.WithValue(r.Context(), TenantContextKey, t) if err := verifySigV4(r, t); err != nil {
next.ServeHTTP(w, r.WithContext(ctx)) writeSQSAuthError(w, "SignatureDoesNotMatch", err.Error())
}) return
} }
ctx := context.WithValue(r.Context(), TenantContextKey, t)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
} }
// extractAccessKeyID — извлекает AWS AccessKeyId из запроса. // extractAccessKeyID — извлекает AWS AccessKeyId из запроса.
@@ -58,57 +80,360 @@ next.ServeHTTP(w, r.WithContext(ctx))
// Ловушка #3: AWS CLI ВСЕГДА отправляет Signature V4 — нужно парсить, даже не проверяя подпись. // Ловушка #3: AWS CLI ВСЕГДА отправляет Signature V4 — нужно парсить, даже не проверяя подпись.
// Ловушка #5: X-Amz-Security-Token (STS) — игнорируем. // Ловушка #5: X-Amz-Security-Token (STS) — игнорируем.
func extractAccessKeyID(r *http.Request) string { func extractAccessKeyID(r *http.Request) string {
// Вариант 1: Authorization header // Вариант 1: Authorization header
// Формат: "AWS4-HMAC-SHA256 Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request, ..." // Формат: "AWS4-HMAC-SHA256 Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request, ..."
auth := r.Header.Get("Authorization") auth := r.Header.Get("Authorization")
if strings.HasPrefix(auth, "AWS4-HMAC-SHA256") { if strings.HasPrefix(auth, "AWS4-HMAC-SHA256") {
idx := strings.Index(auth, "Credential=") idx := strings.Index(auth, "Credential=")
if idx >= 0 { if idx >= 0 {
rest := auth[idx+len("Credential="):] rest := auth[idx+len("Credential="):]
slashIdx := strings.Index(rest, "/") slashIdx := strings.Index(rest, "/")
if slashIdx > 0 { if slashIdx > 0 {
return rest[:slashIdx] 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) // verifySigV4 — полноценная проверка подписи AWS SigV4 для header и presigned запросов.
// Формат: X-Amz-Credential={AccessKeyId}/{date}/{region}/sqs/aws4_request func verifySigV4(r *http.Request, t *tenant.Tenant) error {
if cred := r.URL.Query().Get("X-Amz-Credential"); cred != "" { authHeader := r.Header.Get("Authorization")
parts := strings.SplitN(cred, "/", 2) if strings.HasPrefix(authHeader, "AWS4-HMAC-SHA256") {
if len(parts) > 0 && parts[0] != "" { return verifyHeaderSigV4(r, t, authHeader)
return parts[0] }
} 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 ответ об ошибке аутентификации. // sqsAuthError — AWS-совместимый XML ответ об ошибке аутентификации.
type sqsAuthError struct { type sqsAuthError struct {
XMLName xml.Name `xml:"ErrorResponse"` XMLName xml.Name `xml:"ErrorResponse"`
Error sqsErrorBody `xml:"Error"` Error sqsErrorBody `xml:"Error"`
RequestID string `xml:"RequestId"` RequestID string `xml:"RequestId"`
} }
type sqsErrorBody struct { type sqsErrorBody struct {
Type string `xml:"Type"` Type string `xml:"Type"`
Code string `xml:"Code"` Code string `xml:"Code"`
Message string `xml:"Message"` Message string `xml:"Message"`
} }
// writeSQSAuthError — отвечает AWS-совместимым XML с кодом 403. // writeSQSAuthError — отвечает AWS-совместимым XML с кодом 403.
func writeSQSAuthError(w http.ResponseWriter, code, message string) { func writeSQSAuthError(w http.ResponseWriter, code, message string) {
w.Header().Set("Content-Type", "application/xml") w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(http.StatusForbidden) w.WriteHeader(http.StatusForbidden)
resp := sqsAuthError{ resp := sqsAuthError{
Error: sqsErrorBody{ Error: sqsErrorBody{
Type: "Sender", Type: "Sender",
Code: code, Code: code,
Message: message, Message: message,
}, },
RequestID: "00000000-0000-0000-0000-000000000000", RequestID: "00000000-0000-0000-0000-000000000000",
} }
data, _ := xml.Marshal(resp) data, _ := xml.Marshal(resp)
w.Write(data) w.Write(data)
} }
+94 -91
View File
@@ -3,117 +3,120 @@
package gosqs package gosqs
import ( import (
"net/http" "net/http"
"strings" "strings"
"shared-sqs/app/interfaces" "github.com/gorilla/mux"
"shared-sqs/app/models" log "github.com/sirupsen/logrus"
"shared-sqs/app/utils" "shared-sqs/app/interfaces"
"github.com/gorilla/mux" "shared-sqs/app/models"
log "github.com/sirupsen/logrus" "shared-sqs/app/persistence"
"shared-sqs/app/utils"
) )
func DeleteMessageBatchV1(req *http.Request) (int, interfaces.AbstractResponseBody) { func DeleteMessageBatchV1(req *http.Request) (int, interfaces.AbstractResponseBody) {
requestBody := models.NewDeleteMessageBatchRequest() requestBody := models.NewDeleteMessageBatchRequest()
ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) ok := utils.REQUEST_TRANSFORMER(requestBody, req, false)
if !ok { if !ok {
log.Error("Invalid Request - DeleteMessageBatchV1") log.Error("Invalid Request - DeleteMessageBatchV1")
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
t := getTenantFromContext(req) t := getTenantFromContext(req)
if t == nil { if t == nil {
return utils.CreateErrorResponseV1("InvalidClientTokenId", true) return utils.CreateErrorResponseV1("InvalidClientTokenId", true)
} }
queueUrl := requestBody.QueueUrl queueUrl := requestBody.QueueUrl
queueName := "" queueName := ""
if queueUrl == "" { if queueUrl == "" {
vars := mux.Vars(req) vars := mux.Vars(req)
queueName = vars["queueName"] queueName = vars["queueName"]
} else { } else {
uriSegments := strings.Split(queueUrl, "/") uriSegments := strings.Split(queueUrl, "/")
queueName = uriSegments[len(uriSegments)-1] queueName = uriSegments[len(uriSegments)-1]
} }
key := tenantQueueKey(t.AccessKey, queueName) key := tenantQueueKey(t.AccessKey, queueName)
if _, ok := models.SyncQueues.Queues[key]; !ok { if _, ok := models.SyncQueues.Queues[key]; !ok {
return utils.CreateErrorResponseV1("QueueNotFound", true) return utils.CreateErrorResponseV1("QueueNotFound", true)
} }
if len(requestBody.Entries) == 0 { if len(requestBody.Entries) == 0 {
return utils.CreateErrorResponseV1("EmptyBatchRequest", true) return utils.CreateErrorResponseV1("EmptyBatchRequest", true)
} }
if len(requestBody.Entries) > 10 { if len(requestBody.Entries) > 10 {
return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true) return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true)
} }
ids := map[string]bool{} ids := map[string]bool{}
for _, v := range requestBody.Entries { for _, v := range requestBody.Entries {
if _, found := ids[v.Id]; found { if _, found := ids[v.Id]; found {
return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true) return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true)
} }
ids[v.Id] = true ids[v.Id] = true
} }
models.SyncQueues.Lock() models.SyncQueues.Lock()
defer models.SyncQueues.Unlock() defer models.SyncQueues.Unlock()
deleteMessageMap := make(map[string]*deleteEntry) deleteMessageMap := make(map[string]*deleteEntry)
for _, entry := range requestBody.Entries { for _, entry := range requestBody.Entries {
deleteMessageMap[entry.ReceiptHandle] = &deleteEntry{ deleteMessageMap[entry.ReceiptHandle] = &deleteEntry{
Id: entry.Id, Id: entry.Id,
ReceiptHandle: entry.ReceiptHandle, ReceiptHandle: entry.ReceiptHandle,
Deleted: false, Deleted: false,
} }
} }
deletedEntries := make([]models.DeleteMessageBatchResultEntry, 0) deletedEntries := make([]models.DeleteMessageBatchResultEntry, 0)
remainingMessages := make([]models.SqsMessage, 0, len(models.SyncQueues.Queues[key].Messages)) remainingMessages := make([]models.SqsMessage, 0, len(models.SyncQueues.Queues[key].Messages))
for _, message := range models.SyncQueues.Queues[key].Messages { for _, message := range models.SyncQueues.Queues[key].Messages {
if de, found := deleteMessageMap[message.ReceiptHandle]; found { if de, found := deleteMessageMap[message.ReceiptHandle]; found {
log.Debugf("FIFO Queue %s unlocking group %s:", queueName, message.GroupID) log.Debugf("FIFO Queue %s unlocking group %s:", queueName, message.GroupID)
models.SyncQueues.Queues[key].UnlockGroup(message.GroupID) models.SyncQueues.Queues[key].UnlockGroup(message.GroupID)
delete(models.SyncQueues.Queues[key].Duplicates, message.DeduplicationID) delete(models.SyncQueues.Queues[key].Duplicates, message.DeduplicationID)
de.Deleted = true de.Deleted = true
deletedEntries = append(deletedEntries, models.DeleteMessageBatchResultEntry{Id: de.Id}) deletedEntries = append(deletedEntries, models.DeleteMessageBatchResultEntry{Id: de.Id})
} else { } else {
remainingMessages = append(remainingMessages, message) 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) notFoundEntries := make([]models.BatchResultErrorEntry, 0)
for _, de := range deleteMessageMap { for _, de := range deleteMessageMap {
if !de.Deleted { if !de.Deleted {
notFoundEntries = append(notFoundEntries, models.BatchResultErrorEntry{ notFoundEntries = append(notFoundEntries, models.BatchResultErrorEntry{
Code: "1", Code: "1",
Id: de.Id, Id: de.Id,
Message: "Message not found", Message: "Message not found",
SenderFault: true, SenderFault: true,
}) })
} }
} }
respStruct := models.DeleteMessageBatchResponse{ respStruct := models.DeleteMessageBatchResponse{
Xmlns: models.BaseXmlns, Xmlns: models.BaseXmlns,
Result: models.DeleteMessageBatchResult{ Result: models.DeleteMessageBatchResult{
Successful: deletedEntries, Successful: deletedEntries,
Failed: notFoundEntries, Failed: notFoundEntries,
}, },
Metadata: models.BaseResponseMetadata, Metadata: models.BaseResponseMetadata,
} }
return http.StatusOK, respStruct return http.StatusOK, respStruct
} }
type deleteEntry struct { type deleteEntry struct {
Id string Id string
ReceiptHandle string ReceiptHandle string
Error string Error string
Deleted bool Deleted bool
} }
+33 -32
View File
@@ -3,44 +3,45 @@
package gosqs package gosqs
import ( import (
"net/http" "net/http"
"shared-sqs/app/interfaces" "shared-sqs/app/interfaces"
"shared-sqs/app/models" "shared-sqs/app/models"
"shared-sqs/app/utils" "shared-sqs/app/utils"
log "github.com/sirupsen/logrus"
log "github.com/sirupsen/logrus"
) )
func GetQueueUrlV1(req *http.Request) (int, interfaces.AbstractResponseBody) { func GetQueueUrlV1(req *http.Request) (int, interfaces.AbstractResponseBody) {
requestBody := models.NewGetQueueUrlRequest() requestBody := models.NewGetQueueUrlRequest()
ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) ok := utils.REQUEST_TRANSFORMER(requestBody, req, false)
if !ok { if !ok {
log.Error("Invalid Request - GetQueueUrlV1") log.Error("Invalid Request - GetQueueUrlV1")
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
t := getTenantFromContext(req) t := getTenantFromContext(req)
if t == nil { if t == nil {
return utils.CreateErrorResponseV1("InvalidClientTokenId", true) return utils.CreateErrorResponseV1("InvalidClientTokenId", true)
} }
queueName := requestBody.QueueName queueName := requestBody.QueueName
key := tenantQueueKey(t.AccessKey, queueName) key := tenantQueueKey(t.AccessKey, queueName)
// Fix #10: RLock перед чтением SyncQueues — иначе data race // Fix #10: RLock перед чтением SyncQueues — иначе data race
models.SyncQueues.RLock() models.SyncQueues.RLock()
queue, ok := models.SyncQueues.Queues[key] queue, ok := models.SyncQueues.Queues[key]
models.SyncQueues.RUnlock() models.SyncQueues.RUnlock()
if !ok { if !ok {
log.Errorf("Get Queue URL: %s, queue does not exist for tenant %s", queueName, t.ID) log.Errorf("Get Queue URL: %s, queue does not exist for tenant %s", queueName, t.ID)
return utils.CreateErrorResponseV1("QueueNotFound", true) return utils.CreateErrorResponseV1("QueueNotFound", true)
} }
log.Debug("Get Queue URL:", queue.Name) log.Debug("Get Queue URL:", queue.Name)
respStruct := models.GetQueueUrlResponse{ respStruct := models.GetQueueUrlResponse{
Xmlns: models.BaseXmlns, Xmlns: models.BaseXmlns,
Result: models.GetQueueUrlResult{QueueUrl: queue.URL}, Result: models.GetQueueUrlResult{QueueUrl: queue.URL},
Metadata: models.BaseResponseMetadata, Metadata: models.BaseResponseMetadata,
} }
return http.StatusOK, respStruct return http.StatusOK, respStruct
} }
+35 -35
View File
@@ -4,48 +4,48 @@
package gosqs package gosqs
import ( import (
"net/http" "net/http"
"strings" "strings"
"shared-sqs/app/interfaces" log "github.com/sirupsen/logrus"
"shared-sqs/app/models" "shared-sqs/app/interfaces"
"shared-sqs/app/utils" "shared-sqs/app/models"
log "github.com/sirupsen/logrus" "shared-sqs/app/utils"
) )
func ListQueuesV1(req *http.Request) (int, interfaces.AbstractResponseBody) { func ListQueuesV1(req *http.Request) (int, interfaces.AbstractResponseBody) {
requestBody := models.NewListQueuesRequest() requestBody := models.NewListQueuesRequest()
ok := utils.REQUEST_TRANSFORMER(requestBody, req, true) ok := utils.REQUEST_TRANSFORMER(requestBody, req, true)
if !ok { if !ok {
log.Error("Invalid Request - ListQueuesV1") log.Error("Invalid Request - ListQueuesV1")
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
t := getTenantFromContext(req) t := getTenantFromContext(req)
if t == nil { if t == nil {
return utils.CreateErrorResponseV1("InvalidClientTokenId", true) return utils.CreateErrorResponseV1("InvalidClientTokenId", true)
} }
log.Infof("Listing Queues for tenant: %s", t.ID) log.Infof("Listing Queues for tenant: %s", t.ID)
queueUrls := make([]string, 0) queueUrls := make([]string, 0)
prefix := t.AccessKey + ":" prefix := t.AccessKey + ":"
models.SyncQueues.Lock() models.SyncQueues.RLock()
for key, queue := range models.SyncQueues.Queues { for key, queue := range models.SyncQueues.Queues {
// Показываем только очереди этого тенанта // Показываем только очереди этого тенанта
if strings.HasPrefix(key, prefix) { if strings.HasPrefix(key, prefix) {
if strings.HasPrefix(queue.Name, requestBody.QueueNamePrefix) { if strings.HasPrefix(queue.Name, requestBody.QueueNamePrefix) {
queueUrls = append(queueUrls, queue.URL) queueUrls = append(queueUrls, queue.URL)
} }
} }
} }
models.SyncQueues.Unlock() models.SyncQueues.RUnlock()
respStruct := models.ListQueuesResponse{ respStruct := models.ListQueuesResponse{
Xmlns: models.BaseXmlns, Xmlns: models.BaseXmlns,
Metadata: models.BaseResponseMetadata, Metadata: models.BaseResponseMetadata,
Result: models.ListQueuesResult{QueueUrls: queueUrls}, Result: models.ListQueuesResult{QueueUrls: queueUrls},
} }
return http.StatusOK, respStruct return http.StatusOK, respStruct
} }
+140 -139
View File
@@ -4,169 +4,170 @@
package gosqs package gosqs
import ( import (
"fmt" "fmt"
"net/http" "net/http"
"strings" "strings"
"time" "time"
"github.com/google/uuid" "github.com/google/uuid"
"shared-sqs/app/interfaces" "shared-sqs/app/interfaces"
"shared-sqs/app/models" "shared-sqs/app/models"
"shared-sqs/app/utils" "shared-sqs/app/utils"
"github.com/gorilla/mux"
log "github.com/sirupsen/logrus" "github.com/gorilla/mux"
log "github.com/sirupsen/logrus"
) )
func ReceiveMessageV1(req *http.Request) (int, interfaces.AbstractResponseBody) { func ReceiveMessageV1(req *http.Request) (int, interfaces.AbstractResponseBody) {
requestBody := models.NewReceiveMessageRequest() requestBody := models.NewReceiveMessageRequest()
ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) ok := utils.REQUEST_TRANSFORMER(requestBody, req, false)
if !ok { if !ok {
log.Error("Invalid Request - ReceiveMessageV1") log.Error("Invalid Request - ReceiveMessageV1")
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
t := getTenantFromContext(req) t := getTenantFromContext(req)
if t == nil { if t == nil {
return utils.CreateErrorResponseV1("InvalidClientTokenId", true) return utils.CreateErrorResponseV1("InvalidClientTokenId", true)
} }
maxNumberOfMessages := requestBody.MaxNumberOfMessages maxNumberOfMessages := requestBody.MaxNumberOfMessages
if maxNumberOfMessages == 0 { if maxNumberOfMessages == 0 {
maxNumberOfMessages = 1 maxNumberOfMessages = 1
} }
// Fix #8: clamp MaxNumberOfMessages к AWS лимиту 1–10 // Fix #8: clamp MaxNumberOfMessages к AWS лимиту 1–10
maxNumberOfMessages = ClampInt(maxNumberOfMessages, MinNumberOfMessagesLimit, MaxNumberOfMessagesLimit) maxNumberOfMessages = ClampInt(maxNumberOfMessages, MinNumberOfMessagesLimit, MaxNumberOfMessagesLimit)
queueName := "" queueName := ""
if requestBody.QueueUrl == "" { if requestBody.QueueUrl == "" {
vars := mux.Vars(req) vars := mux.Vars(req)
queueName = vars["queueName"] queueName = vars["queueName"]
} else { } else {
uriSegments := strings.Split(requestBody.QueueUrl, "/") uriSegments := strings.Split(requestBody.QueueUrl, "/")
queueName = uriSegments[len(uriSegments)-1] queueName = uriSegments[len(uriSegments)-1]
} }
key := tenantQueueKey(t.AccessKey, queueName) key := tenantQueueKey(t.AccessKey, queueName)
if _, ok := models.SyncQueues.Queues[key]; !ok { if _, ok := models.SyncQueues.Queues[key]; !ok {
return utils.CreateErrorResponseV1("QueueNotFound", true) return utils.CreateErrorResponseV1("QueueNotFound", true)
} }
var messages []*models.ResultMessage var messages []*models.ResultMessage
respStruct := models.ReceiveMessageResponse{} respStruct := models.ReceiveMessageResponse{}
waitTimeSeconds := requestBody.WaitTimeSeconds waitTimeSeconds := requestBody.WaitTimeSeconds
if waitTimeSeconds == 0 { if waitTimeSeconds == 0 {
models.SyncQueues.RLock() models.SyncQueues.RLock()
waitTimeSeconds = models.SyncQueues.Queues[key].ReceiveMessageWaitTimeSeconds waitTimeSeconds = models.SyncQueues.Queues[key].ReceiveMessageWaitTimeSeconds
models.SyncQueues.RUnlock() models.SyncQueues.RUnlock()
} }
// Fix #4: clamp WaitTimeSeconds к AWS лимиту 0–20 // Fix #4: clamp WaitTimeSeconds к AWS лимиту 0–20
waitTimeSeconds = ClampInt(waitTimeSeconds, 0, MaxReceiveMessageWaitTimeSeconds) waitTimeSeconds = ClampInt(waitTimeSeconds, 0, MaxReceiveMessageWaitTimeSeconds)
// Long polling: ждём появления сообщения до waitTimeSeconds*10 итераций по 100ms // Long polling: ждём появления сообщения до waitTimeSeconds*10 итераций по 100ms
loops := waitTimeSeconds * 10 loops := waitTimeSeconds * 10
for loops > 0 { for loops > 0 {
models.SyncQueues.RLock() models.SyncQueues.RLock()
_, queueFound := models.SyncQueues.Queues[key] _, queueFound := models.SyncQueues.Queues[key]
if !queueFound { if !queueFound {
models.SyncQueues.RUnlock() models.SyncQueues.RUnlock()
return utils.CreateErrorResponseV1("QueueNotFound", true) return utils.CreateErrorResponseV1("QueueNotFound", true)
} }
messageFound := len(models.SyncQueues.Queues[key].Messages)-numberOfHiddenMessagesInQueue(*models.SyncQueues.Queues[key]) != 0 messageFound := len(models.SyncQueues.Queues[key].Messages)-numberOfHiddenMessagesInQueue(*models.SyncQueues.Queues[key]) != 0
models.SyncQueues.RUnlock() models.SyncQueues.RUnlock()
if !messageFound { if !messageFound {
continueTimer := time.NewTimer(100 * time.Millisecond) continueTimer := time.NewTimer(100 * time.Millisecond)
select { select {
case <-req.Context().Done(): case <-req.Context().Done():
continueTimer.Stop() continueTimer.Stop()
return http.StatusOK, models.ReceiveMessageResponse{ return http.StatusOK, models.ReceiveMessageResponse{
Xmlns: models.BaseXmlns, Xmlns: models.BaseXmlns,
Result: models.ReceiveMessageResult{}, Result: models.ReceiveMessageResult{},
Metadata: models.BaseResponseMetadata, Metadata: models.BaseResponseMetadata,
} }
case <-continueTimer.C: case <-continueTimer.C:
continueTimer.Stop() continueTimer.Stop()
} }
loops-- loops--
} else { } else {
break break
} }
} }
log.Debugf("Getting Message from Queue:%s (tenant: %s)", queueName, t.ID) log.Debugf("Getting Message from Queue:%s (tenant: %s)", queueName, t.ID)
models.SyncQueues.Lock() models.SyncQueues.Lock()
defer models.SyncQueues.Unlock() defer models.SyncQueues.Unlock()
if len(models.SyncQueues.Queues[key].Messages) > 0 { if len(models.SyncQueues.Queues[key].Messages) > 0 {
numMsg := 0 numMsg := 0
messages = make([]*models.ResultMessage, 0) messages = make([]*models.ResultMessage, 0)
for i := range models.SyncQueues.Queues[key].Messages { for i := range models.SyncQueues.Queues[key].Messages {
if numMsg >= maxNumberOfMessages { if numMsg >= maxNumberOfMessages {
break break
} }
if models.SyncQueues.Queues[key].Messages[i].ReceiptHandle != "" { if models.SyncQueues.Queues[key].Messages[i].ReceiptHandle != "" {
continue continue
} }
msg := &models.SyncQueues.Queues[key].Messages[i] msg := &models.SyncQueues.Queues[key].Messages[i]
if !msg.IsReadyForReceipt() { if !msg.IsReadyForReceipt() {
continue continue
} }
if models.SyncQueues.Queues[key].IsFIFO { if models.SyncQueues.Queues[key].IsFIFO {
if models.SyncQueues.Queues[key].IsLocked(msg.GroupID) { if models.SyncQueues.Queues[key].IsLocked(msg.GroupID) {
continue continue
} }
models.SyncQueues.Queues[key].LockGroup(msg.GroupID) models.SyncQueues.Queues[key].LockGroup(msg.GroupID)
} }
randomId := uuid.NewString() randomId := uuid.NewString()
msg.ReceiptHandle = msg.Uuid + "#" + randomId msg.ReceiptHandle = msg.Uuid + "#" + randomId
msg.ReceiptTime = time.Now().UTC() msg.ReceiptTime = time.Now().UTC()
if requestBody.VisibilityTimeout != 0 { if requestBody.VisibilityTimeout != 0 {
msg.VisibilityTimeout = time.Now().Add(time.Duration(requestBody.VisibilityTimeout) * time.Second) msg.VisibilityTimeout = time.Now().Add(time.Duration(requestBody.VisibilityTimeout) * time.Second)
} else { } else {
msg.VisibilityTimeout = time.Now().Add(time.Duration(models.SyncQueues.Queues[key].VisibilityTimeout) * time.Second) msg.VisibilityTimeout = time.Now().Add(time.Duration(models.SyncQueues.Queues[key].VisibilityTimeout) * time.Second)
} }
messages = append(messages, buildResultMessage(msg)) messages = append(messages, buildResultMessage(msg))
numMsg++ numMsg++
} }
respStruct = models.ReceiveMessageResponse{ respStruct = models.ReceiveMessageResponse{
"http://queue.amazonaws.com/doc/2012-11-05/", "http://queue.amazonaws.com/doc/2012-11-05/",
models.ReceiveMessageResult{Messages: messages}, models.ReceiveMessageResult{Messages: messages},
models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"}, models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"},
} }
} else { } else {
log.Warning("No messages in Queue:", queueName) log.Warning("No messages in Queue:", queueName)
respStruct = models.ReceiveMessageResponse{ respStruct = models.ReceiveMessageResponse{
Xmlns: "http://queue.amazonaws.com/doc/2012-11-05/", Xmlns: "http://queue.amazonaws.com/doc/2012-11-05/",
Result: models.ReceiveMessageResult{}, Result: models.ReceiveMessageResult{},
Metadata: models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"}, Metadata: models.ResponseMetadata{RequestId: "00000000-0000-0000-0000-000000000000"},
} }
} }
return http.StatusOK, respStruct return http.StatusOK, respStruct
} }
func buildResultMessage(m *models.SqsMessage) *models.ResultMessage { func buildResultMessage(m *models.SqsMessage) *models.ResultMessage {
return &models.ResultMessage{ return &models.ResultMessage{
MessageId: m.Uuid, MessageId: m.Uuid,
Body: m.MessageBody, Body: m.MessageBody,
ReceiptHandle: m.ReceiptHandle, ReceiptHandle: m.ReceiptHandle,
MD5OfBody: utils.GetMD5Hash(m.MessageBody), MD5OfBody: utils.GetMD5Hash(m.MessageBody),
MD5OfMessageAttributes: m.MD5OfMessageAttributes, MD5OfMessageAttributes: m.MD5OfMessageAttributes,
MessageAttributes: m.MessageAttributes, MessageAttributes: m.MessageAttributes,
Attributes: map[string]string{ Attributes: map[string]string{
"ApproximateFirstReceiveTimestamp": fmt.Sprintf("%d", m.ReceiptTime.UnixNano()/int64(time.Millisecond)), "ApproximateFirstReceiveTimestamp": fmt.Sprintf("%d", m.ReceiptTime.UnixNano()/int64(time.Millisecond)),
"SenderId": models.CurrentEnvironment.AccountID, "SenderId": models.CurrentEnvironment.AccountID,
"ApproximateReceiveCount": fmt.Sprintf("%d", m.NumberOfReceives+1), "ApproximateReceiveCount": fmt.Sprintf("%d", m.NumberOfReceives+1),
"SentTimestamp": fmt.Sprintf("%d", time.Now().UTC().UnixNano()/int64(time.Millisecond)), "SentTimestamp": fmt.Sprintf("%d", time.Now().UTC().UnixNano()/int64(time.Millisecond)),
}, },
} }
} }
+10 -8
View File
@@ -63,25 +63,27 @@ func SendMessageV1(req *http.Request) (int, interfaces.AbstractResponseBody) {
key := tenantQueueKey(t.AccessKey, queueName) 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) 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 && if maxMessageSize > 0 && len(messageBody) > maxMessageSize {
len(messageBody) > models.SyncQueues.Queues[key].MaximumMessageSize {
return utils.CreateErrorResponseV1("MessageTooBig", true) return utils.CreateErrorResponseV1("MessageTooBig", true)
} }
// Fix #11: лимит сообщений в очереди — защита от OOM // 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) { if currentMsgCount >= MaxMessagesForQueue(queueIsFIFO) {
return utils.CreateErrorResponseV1("OverLimit", true) return utils.CreateErrorResponseV1("OverLimit", true)
} }
delaySecs := models.SyncQueues.Queues[key].DelaySeconds
if requestBody.DelaySeconds != 0 { if requestBody.DelaySeconds != 0 {
delaySecs = requestBody.DelaySeconds delaySecs = requestBody.DelaySeconds
} }
+126 -119
View File
@@ -3,143 +3,150 @@
package gosqs package gosqs
import ( import (
"net/http" "net/http"
"strings" "strings"
"time" "time"
"github.com/google/uuid" "github.com/google/uuid"
"shared-sqs/app/interfaces" "shared-sqs/app/interfaces"
"shared-sqs/app/models" "shared-sqs/app/models"
"shared-sqs/app/utils" "shared-sqs/app/persistence"
"github.com/gorilla/mux" "shared-sqs/app/utils"
log "github.com/sirupsen/logrus"
"github.com/gorilla/mux"
log "github.com/sirupsen/logrus"
) )
func SendMessageBatchV1(req *http.Request) (int, interfaces.AbstractResponseBody) { func SendMessageBatchV1(req *http.Request) (int, interfaces.AbstractResponseBody) {
requestBody := models.NewSendMessageBatchRequest() requestBody := models.NewSendMessageBatchRequest()
ok := utils.REQUEST_TRANSFORMER(requestBody, req, false) ok := utils.REQUEST_TRANSFORMER(requestBody, req, false)
if !ok { if !ok {
log.Error("Invalid Request - SendMessageBatchV1") log.Error("Invalid Request - SendMessageBatchV1")
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
t := getTenantFromContext(req) t := getTenantFromContext(req)
if t == nil { if t == nil {
return utils.CreateErrorResponseV1("InvalidClientTokenId", true) return utils.CreateErrorResponseV1("InvalidClientTokenId", true)
} }
queueUrl := requestBody.QueueUrl queueUrl := requestBody.QueueUrl
queueName := "" queueName := ""
if queueUrl == "" { if queueUrl == "" {
vars := mux.Vars(req) vars := mux.Vars(req)
queueName = vars["queueName"] queueName = vars["queueName"]
} else { } else {
uriSegments := strings.Split(queueUrl, "/") uriSegments := strings.Split(queueUrl, "/")
queueName = uriSegments[len(uriSegments)-1] queueName = uriSegments[len(uriSegments)-1]
} }
key := tenantQueueKey(t.AccessKey, queueName) key := tenantQueueKey(t.AccessKey, queueName)
if _, ok := models.SyncQueues.Queues[key]; !ok { models.SyncQueues.RLock()
return utils.CreateErrorResponseV1("QueueNotFound", true) 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 { if len(sendEntries) == 0 {
return utils.CreateErrorResponseV1("EmptyBatchRequest", true) return utils.CreateErrorResponseV1("EmptyBatchRequest", true)
} }
if len(sendEntries) > 10 { if len(sendEntries) > 10 {
return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true) return utils.CreateErrorResponseV1("TooManyEntriesInBatchRequest", true)
} }
ids := map[string]struct{}{} ids := map[string]struct{}{}
for _, v := range sendEntries { for _, v := range sendEntries {
if _, ok := ids[v.Id]; ok { if _, ok := ids[v.Id]; ok {
return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true) return utils.CreateErrorResponseV1("BatchEntryIdsNotDistinct", true)
} }
// Валидация длины BatchEntryId (макс 80 chars) // Валидация длины BatchEntryId (макс 80 chars)
if err := ValidateBatchEntryID(v.Id); err != nil { if err := ValidateBatchEntryID(v.Id); err != nil {
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
// Валидация DeduplicationID и GroupID (макс 128 chars) // Валидация DeduplicationID и GroupID (макс 128 chars)
if err := ValidateDeduplicationID(v.MessageDeduplicationId); err != nil { if err := ValidateDeduplicationID(v.MessageDeduplicationId); err != nil {
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
if err := ValidateGroupID(v.MessageGroupId); err != nil { if err := ValidateGroupID(v.MessageGroupId); err != nil {
return utils.CreateErrorResponseV1("InvalidParameterValue", true) return utils.CreateErrorResponseV1("InvalidParameterValue", true)
} }
ids[v.Id] = struct{}{} ids[v.Id] = struct{}{}
} }
sentEntries := make([]models.SendMessageBatchResultEntry, 0) sentEntries := make([]models.SendMessageBatchResultEntry, 0)
log.Debugf("Batch sending to Queue: %s (tenant: %s)", queueName, t.ID) log.Debugf("Batch sending to Queue: %s (tenant: %s)", queueName, t.ID)
// Fix #2: проверяем размер каждого сообщения в batch (Critical — batch size bypass) // Fix #2: проверяем размер каждого сообщения в batch (Critical — batch size bypass)
maxMsgSize := models.SyncQueues.Queues[key].MaximumMessageSize if maxMsgSize <= 0 {
if maxMsgSize <= 0 { maxMsgSize = MaxMessageSizeDefault
maxMsgSize = MaxMessageSizeDefault }
} for _, entry := range sendEntries {
for _, entry := range sendEntries { if len(entry.MessageBody) > maxMsgSize {
if len(entry.MessageBody) > maxMsgSize { return utils.CreateErrorResponseV1("MessageTooBig", true)
return utils.CreateErrorResponseV1("MessageTooBig", true) }
} // Валидация количества message attributes (макс 10 по AWS)
// Валидация количества message attributes (макс 10 по AWS) if len(entry.MessageAttributes) > MaxMessageAttributes {
if len(entry.MessageAttributes) > MaxMessageAttributes { return utils.CreateErrorResponseV1("InvalidParameterValue", true)
return utils.CreateErrorResponseV1("InvalidParameterValue", true) }
} }
}
// Fix #11: лимит сообщений в очереди — защита от OOM // Fix #11: лимит сообщений в очереди — защита от OOM
models.SyncQueues.RLock() if currentMsgCount+len(sendEntries) > MaxMessagesForQueue(queueIsFIFO) {
currentMsgCount := len(models.SyncQueues.Queues[key].Messages) return utils.CreateErrorResponseV1("OverLimit", true)
queueIsFIFO := models.SyncQueues.Queues[key].IsFIFO }
models.SyncQueues.RUnlock()
if currentMsgCount+len(sendEntries) > MaxMessagesForQueue(queueIsFIFO) {
return utils.CreateErrorResponseV1("OverLimit", true)
}
for _, sendEntry := range sendEntries { models.SyncQueues.Lock()
msg := models.SqsMessage{MessageBody: sendEntry.MessageBody} queue = models.SyncQueues.Queues[key]
if len(sendEntry.MessageAttributes) > 0 { for _, sendEntry := range sendEntries {
msg.MessageAttributes = sendEntry.MessageAttributes msg := models.SqsMessage{MessageBody: sendEntry.MessageBody}
msg.MD5OfMessageAttributes = utils.HashAttributes(sendEntry.MessageAttributes) if len(sendEntry.MessageAttributes) > 0 {
} msg.MessageAttributes = sendEntry.MessageAttributes
msg.MD5OfMessageBody = utils.GetMD5Hash(sendEntry.MessageBody) msg.MD5OfMessageAttributes = utils.HashAttributes(sendEntry.MessageAttributes)
msg.GroupID = sendEntry.MessageGroupId }
msg.DeduplicationID = sendEntry.MessageDeduplicationId msg.MD5OfMessageBody = utils.GetMD5Hash(sendEntry.MessageBody)
msg.Uuid = uuid.NewString() msg.GroupID = sendEntry.MessageGroupId
msg.SentTime = time.Now() msg.DeduplicationID = sendEntry.MessageDeduplicationId
msg.Uuid = uuid.NewString()
msg.SentTime = time.Now()
models.SyncQueues.Lock() fifoSeqNumber := ""
fifoSeqNumber := "" if queue.IsFIFO {
if models.SyncQueues.Queues[key].IsFIFO { fifoSeqNumber = queue.NextSequenceNumber(sendEntry.MessageGroupId)
fifoSeqNumber = models.SyncQueues.Queues[key].NextSequenceNumber(sendEntry.MessageGroupId) }
} if !queue.IsDuplicate(sendEntry.MessageDeduplicationId) {
if !models.SyncQueues.Queues[key].IsDuplicate(sendEntry.MessageDeduplicationId) { queue.Messages = append(queue.Messages, msg)
models.SyncQueues.Queues[key].Messages = append(models.SyncQueues.Queues[key].Messages, msg) } else {
} else { log.Debugf("Duplicate deduplicationId [%s] in queue [%s]", sendEntry.MessageDeduplicationId, queueName)
log.Debugf("Duplicate deduplicationId [%s] in queue [%s]", sendEntry.MessageDeduplicationId, queueName) }
} queue.InitDuplicatation(sendEntry.MessageDeduplicationId)
models.SyncQueues.Queues[key].InitDuplicatation(sendEntry.MessageDeduplicationId)
models.SyncQueues.Unlock()
sentEntries = append(sentEntries, models.SendMessageBatchResultEntry{ sentEntries = append(sentEntries, models.SendMessageBatchResultEntry{
Id: sendEntry.Id, Id: sendEntry.Id,
MessageId: msg.Uuid, MessageId: msg.Uuid,
MD5OfMessageBody: msg.MD5OfMessageBody, MD5OfMessageBody: msg.MD5OfMessageBody,
MD5OfMessageAttributes: msg.MD5OfMessageAttributes, MD5OfMessageAttributes: msg.MD5OfMessageAttributes,
SequenceNumber: fifoSeqNumber, SequenceNumber: fifoSeqNumber,
}) })
log.Infof("%s: Queue: %s, Message: %s", time.Now().Format("2006-01-02 15:04:05"), queueName, msg.MessageBody) 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{ respStruct := models.SendMessageBatchResponse{
Xmlns: models.BaseXmlns, Xmlns: models.BaseXmlns,
Result: models.SendMessageBatchResult{Entry: sentEntries}, Result: models.SendMessageBatchResult{Entry: sentEntries},
Metadata: models.BaseResponseMetadata, Metadata: models.BaseResponseMetadata,
} }
return http.StatusOK, respStruct return http.StatusOK, respStruct
} }
+7 -7
View File
@@ -12,16 +12,16 @@ import (
// ── AWS SQS лимиты ──────────────────────────────────────────────────────────── // ── AWS SQS лимиты ────────────────────────────────────────────────────────────
const ( const (
// Очередь // Очередь
MaxQueueNameLength = 80 MaxQueueNameLength = 80
MaxMessageSizeDefault = 262144 // 256 KB MaxMessageSizeDefault = 262144 // 256 KB
MaxMessageSizeLimit = 262144 MaxMessageSizeLimit = 262144
MinMessageSizeLimit = 1024 // 1 KB MinMessageSizeLimit = 1024 // 1 KB
// Атрибуты очереди // Атрибуты очереди
MaxDelaySeconds = 900 // 15 min MaxDelaySeconds = 900 // 15 min
MaxVisibilityTimeout = 43200 // 12 hours MaxVisibilityTimeout = 43200 // 12 hours
MaxReceiveMessageWaitTimeSeconds = 20 // long polling cap MaxReceiveMessageWaitTimeSeconds = 20 // long polling cap
MinMessageRetentionPeriod = 60 // 1 min MinMessageRetentionPeriod = 60 // 1 min
MaxMessageRetentionPeriod = 1209600 // 14 days MaxMessageRetentionPeriod = 1209600 // 14 days
// Receive // Receive
@@ -29,7 +29,7 @@ const (
MinNumberOfMessagesLimit = 1 MinNumberOfMessagesLimit = 1
// Message attributes // Message attributes
MaxMessageAttributes = 10 MaxMessageAttributes = 10
MaxMessageAttributeSize = 262144 // 256 KB суммарно (тело + атрибуты) MaxMessageAttributeSize = 262144 // 256 KB суммарно (тело + атрибуты)
// Deduplication / GroupID // Deduplication / GroupID
+5 -1
View File
@@ -147,7 +147,11 @@ func extractAction(req *http.Request) string {
switch protocol { switch protocol {
case AwsJsonProtocol: case AwsJsonProtocol:
action := req.Header.Get("X-Amz-Target") 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: case AwsQueryProtocol:
return req.FormValue("Action") return req.FormValue("Action")
} }
+4
View File
@@ -202,6 +202,10 @@ func (s *TenantStore) CreateFromJWT(tenantID, sub, email string, maxQueues int)
s.mu.Unlock() s.mu.Unlock()
return existing, nil return existing, nil
} }
if len(s.byID) >= MaxTenantsGlobal {
s.mu.Unlock()
return nil, fmt.Errorf("global tenant limit reached (%d)", MaxTenantsGlobal)
}
accessKey, err := generateAccessKey() accessKey, err := generateAccessKey()
if err != nil { if err != nil {