Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions api/v1/admin/marketplace/marketplace.go
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,11 @@ func updateMarketplace(w http.ResponseWriter, r *http.Request) {
return
}

if slug == "csfloat" && opts.IsActive != nil && *opts.IsActive == false {
render.Errorf(w, r, errors.Forbidden, "cannot deactivate csfloat marketplace")
return
}

var updatedMarketplace *models.Marketplace
err := factory.RunInTransactionPrivate(func(tx repository.PrivateTransaction) error {
if err := tx.Marketplace().Update(slug, &opts); err != nil {
Expand Down
58 changes: 58 additions & 0 deletions api/v1/admin/marketplace/marketplace_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -646,6 +646,64 @@ func TestUpdateMarketplace(t *testing.T) {
}
},
},
{
name: "cannotDeactivateCSFloat",
setup: func(t *testing.T, db *gorm.DB, f repository.Factory, keygen isecret.KeyGenerator) (*http.Request, string, string, error) {
_, authKey, formattedKey := testutil.SetupMarketplaceWithKey(t, db, "csfloat", keygen, models.PermissionAdmin)

isActive := false
reqBody := dto.MarketplaceUpdates{
IsActive: &isActive,
}

payload, err := json.Marshal(reqBody)
if err != nil {
return nil, "", "", err
}

r := httptest.NewRequest(http.MethodPatch, "/csfloat", bytes.NewBuffer(payload))
r.Header.Set("Content-Type", "application/json")
r.Header.Set("Authorization", "Bearer "+formattedKey)

return r, "csfloat", authKey.ID, nil
},
validateFunc: func(t *testing.T, db *gorm.DB, authKeyID string, resp *http.Response) {
if resp.StatusCode != http.StatusForbidden {
t.Errorf("wanted status code %d, got %d", http.StatusForbidden, resp.StatusCode)
}

defer resp.Body.Close()
var respData errors.Error
if err := json.NewDecoder(resp.Body).Decode(&respData); err != nil {
t.Fatalf("failed to decode response body: %v", err)
}

if diff := cmp.Diff(errors.Forbidden, respData, cmpopts.IgnoreFields(errors.Error{}, "status", "wrapped", "Details")); diff != "" {
t.Error(diff)
}

if respData.Details != "cannot deactivate csfloat marketplace" {
t.Errorf("wanted details %q, got %q", "cannot deactivate csfloat marketplace", respData.Details)
}

var storedMarketplace models.Marketplace
if err := db.Where("slug = ?", "csfloat").First(&storedMarketplace).Error; err != nil {
t.Fatalf("First(): %v", err)
}
if !storedMarketplace.IsActive {
t.Error("expected csfloat marketplace to remain active")
}

var auditCount int64
if err := db.Model(&models.AdminAudit{}).Where("target_action = ? AND target_resource_type = ? AND target_resource = ?",
models.TargetActionUpdateMarketplace, models.TargetResourceTypeMarketplace, "csfloat").Count(&auditCount).Error; err != nil {
t.Fatalf("Count(): %v", err)
}
if auditCount != 0 {
t.Errorf("wanted 0 update audit records, got %d", auditCount)
}
},
},
{
name: "emptySlug",
setup: func(t *testing.T, db *gorm.DB, f repository.Factory, keygen isecret.KeyGenerator) (*http.Request, string, string, error) {
Expand Down
2 changes: 1 addition & 1 deletion api/v1/marketplace/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ func Router() chi.Router {
r.Use(rwmiddleware.RequirePermissions(models.PermissionManage))

r.Route("/keys", func(r chi.Router) {
r.Use(ratelimit.ThrottleByAPIKey(time.Minute, 100))
r.Use(ratelimit.ThrottleByMarketplace(time.Minute, 100))

r.Get("/", listKeys)
r.Post("/", createKey)
Expand Down
12 changes: 11 additions & 1 deletion api/v1/reversals/reversals.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@ import (
"github.com/go-chi/chi/v5"
)

const maxBatchSize = 1_000

func createReversals(w http.ResponseWriter, r *http.Request) {
factory, ok := r.Context().Value(middleware.FactoryContextKey).(repository.Factory)
if !ok {
Expand Down Expand Up @@ -49,8 +51,16 @@ func createReversals(w http.ResponseWriter, r *http.Request) {
render.Error(w, r, &errors.JSONDecode)
return
}
if len(req.Data) == 0 {
render.Errorf(w, r, errors.BadRequest, "data cannot be empty")
return
}
if len(req.Data) > maxBatchSize {
render.Errorf(w, r, errors.BadRequest, "data exceeds max length of %d", maxBatchSize)
return
}

reversals := make([]*models.Reversal, 0)
reversals := make([]*models.Reversal, 0, len(req.Data))
for _, reversal := range req.Data {
reversals = append(reversals, &models.Reversal{
SteamID: reversal.SteamID,
Expand Down
77 changes: 77 additions & 0 deletions api/v1/reversals/reversals_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,83 @@ func TestCreateReversals(t *testing.T) {
}
},
},
{
name: "emptyData",
setup: func(t *testing.T, db *gorm.DB, f repository.Factory, keygen isecret.KeyGenerator) (*http.Request, []*reversal, error) {
_, _, formattedKey := testutil.SetupMarketplaceWithKey(t, db, "test-marketplace", keygen, models.PermissionWrite)

body, err := json.Marshal(&req{Data: []*reversal{}})
if err != nil {
return nil, nil, err
}

r := httptest.NewRequest(http.MethodPost, "/", bytes.NewBuffer(body))
r.Header.Set("Content-Type", "application/json")
r.Header.Set("Authorization", "Bearer "+formattedKey)
return r, nil, nil
},
validateFunc: func(t *testing.T, db *gorm.DB, data []*reversal, resp *http.Response) {
if resp.StatusCode != http.StatusBadRequest {
t.Errorf("wanted status code %d, got %d", http.StatusBadRequest, resp.StatusCode)
}

defer resp.Body.Close()
var respData errors.Error
if err := json.NewDecoder(resp.Body).Decode(&respData); err != nil {
t.Fatalf("failed to decode response body: %v", err)
}

if respData.Details != "data cannot be empty" {
t.Errorf("wanted details %q, got %q", "data cannot be empty", respData.Details)
}
},
},
{
name: "dataExceedsMaxLength",
setup: func(t *testing.T, db *gorm.DB, f repository.Factory, keygen isecret.KeyGenerator) (*http.Request, []*reversal, error) {
_, _, formattedKey := testutil.SetupMarketplaceWithKey(t, db, "test-marketplace", keygen, models.PermissionWrite)

data := make([]*reversal, maxBatchSize+1)
for i := range data {
data[i] = &reversal{
SteamID: models.SteamID(76561197960287930 + i),
}
}
body, err := json.Marshal(&req{Data: data})
if err != nil {
return nil, nil, err
}

r := httptest.NewRequest(http.MethodPost, "/", bytes.NewBuffer(body))
r.Header.Set("Content-Type", "application/json")
r.Header.Set("Authorization", "Bearer "+formattedKey)
return r, data, nil
},
validateFunc: func(t *testing.T, db *gorm.DB, data []*reversal, resp *http.Response) {
if resp.StatusCode != http.StatusBadRequest {
t.Errorf("wanted status code %d, got %d", http.StatusBadRequest, resp.StatusCode)
}

defer resp.Body.Close()
var respData errors.Error
if err := json.NewDecoder(resp.Body).Decode(&respData); err != nil {
t.Fatalf("failed to decode response body: %v", err)
}

wantDetails := fmt.Sprintf("data exceeds max length of %d", maxBatchSize)
if respData.Details != wantDetails {
t.Errorf("wanted details %q, got %q", wantDetails, respData.Details)
}

var count int64
if err := db.Model(&models.Reversal{}).Count(&count).Error; err != nil {
t.Fatalf("failed to count reversals: %v", err)
}
if count != 0 {
t.Errorf("wanted 0 reversals, got %d", count)
}
},
},
{
name: "createFailedInvalidReversal",
setup: func(t *testing.T, db *gorm.DB, f repository.Factory, keygen isecret.KeyGenerator) (*http.Request, []*reversal, error) {
Expand Down
8 changes: 4 additions & 4 deletions api/v1/reversals/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,19 +16,19 @@ func Router() chi.Router {

r.With(
middleware.RequirePermissions(models.PermissionWrite),
ratelimit.ThrottleByAPIKey(time.Hour, 2_000),
ratelimit.ThrottleByMarketplace(time.Hour, 2_000),
).Post("/", createReversals)

r.With(
middleware.RequirePermissions(models.PermissionDelete),
ratelimit.ThrottleByAPIKey(time.Hour, 2_000),
ratelimit.ThrottleByMarketplace(time.Hour, 2_000),
).Delete("/{id}", expungeReversal)

r.Route("/", func(r chi.Router) {
r.Use(middleware.RequirePermissions(models.PermissionExport))

r.With(ratelimit.ThrottleByAPIKey(time.Minute, 300)).Get("/", listReversalsHandler)
r.With(ratelimit.ThrottleByAPIKey(time.Minute, 60)).Get("/export", exportReversals)
r.With(ratelimit.ThrottleByMarketplace(time.Minute, 300)).Get("/", listReversalsHandler)
r.With(ratelimit.ThrottleByMarketplace(time.Minute, 60)).Get("/export", exportReversals)
})
return r
}
15 changes: 15 additions & 0 deletions ratelimit/ratelimit.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,21 @@ func ThrottleByAPIKey(dur time.Duration, limit uint64) func(http.Handler) http.H
return newLimiter(dur, limit, byAPIKey)
}

func byMarketplace(r *http.Request) (string, error) {
key, ok := r.Context().Value(middleware.KeyContextKey).(*models.Key)
if !ok {
return "", fmt.Errorf("key not found in context")
}
if key.MarketplaceSlug == "" {
return "", fmt.Errorf("marketplace slug not found in key")
}
return key.MarketplaceSlug, nil
}

func ThrottleByMarketplace(dur time.Duration, limit uint64) func(http.Handler) http.Handler {
return newLimiter(dur, limit, byMarketplace)
}

func newThrottlerWithLimiter(keyFunc httplimit.KeyFunc, store limiter.Store) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
fn := func(w http.ResponseWriter, r *http.Request) {
Expand Down
75 changes: 75 additions & 0 deletions ratelimit/ratelimit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -388,3 +388,78 @@ func TestThrottleByAPIKey(t *testing.T) {
})
}
}

func TestThrottleByMarketplace(t *testing.T) {
t.Parallel()

db := testutil.NewTestDB(t)
keygen := secret.NewKeyGenerator(constants.EnvironmentDevelopment)
f, err := factory.NewFactoryWithConfig(&factory.Config{
PrivateDB: db,
PublicDB: db,
KeyGen: keygen,
})
if err != nil {
t.Fatalf("NewFactoryWithConfig(): %v", err)
}

testMarketplace, _, firstFormattedKey := testutil.SetupMarketplaceWithKey(t, db, "test-marketplace", keygen, models.PermissionWrite)
secondSecretKey, err := keygen.GenerateSecretKey()
if err != nil {
t.Fatalf("GenerateSecretKey(): %v", err)
}
secondKeyID, err := secondSecretKey.ID()
if err != nil {
t.Fatalf("ID(): %v", err)
}
secondKey := &models.Key{
ID: secondKeyID,
Environment: keygen.Environment(),
MarketplaceSlug: testMarketplace.Slug,
Permissions: models.PermissionWrite,
}
testutil.Insert(t, db, secondKey)
secondFormattedKey, err := secondSecretKey.Format()
if err != nil {
t.Fatalf("Format(): %v", err)
}

_, _, otherMarketplaceFormattedKey := testutil.SetupMarketplaceWithKey(t, db, "other-marketplace", keygen, models.PermissionWrite)

fn := func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}
next := http.HandlerFunc(fn)

factoryMiddleware := middleware.FactoryMiddleware(f)
throttleMiddleware := ThrottleByMarketplace(time.Minute, 1)
finalHandler := factoryMiddleware(
middleware.AuthMiddleware(
throttleMiddleware(next),
),
)

requestWithKey := func(formattedKey string) *http.Request {
r := httptest.NewRequest(http.MethodGet, "http://testing", nil)
r.Header.Set("Authorization", "Bearer "+formattedKey)
return r
}

w := httptest.NewRecorder()
finalHandler.ServeHTTP(w, requestWithKey(firstFormattedKey))
if w.Code != http.StatusOK {
t.Fatalf("first marketplace request got status code %d, wanted %d", w.Code, http.StatusOK)
}

w = httptest.NewRecorder()
finalHandler.ServeHTTP(w, requestWithKey(secondFormattedKey))
if w.Code != http.StatusTooManyRequests {
t.Fatalf("second same-marketplace key got status code %d, wanted %d", w.Code, http.StatusTooManyRequests)
}

w = httptest.NewRecorder()
finalHandler.ServeHTTP(w, requestWithKey(otherMarketplaceFormattedKey))
if w.Code != http.StatusOK {
t.Fatalf("different marketplace got status code %d, wanted %d", w.Code, http.StatusOK)
}
}
9 changes: 8 additions & 1 deletion repository/private/key.go
Original file line number Diff line number Diff line change
Expand Up @@ -103,5 +103,12 @@ func (k *keyRepository) List(opts *dto.KeyListOptions) ([]*models.Key, error) {

func (k *keyRepository) ValidateKey(secretKey string) (*models.Key, error) {
hashedKey := secret.Sha256Hash(secretKey)
return k.Read(hashedKey)
key, err := k.Read(hashedKey)
if err != nil {
return nil, err
}
if key.Marketplace == nil || !key.Marketplace.IsActive {
return nil, gorm.ErrRecordNotFound
}
Comment thread
cursor[bot] marked this conversation as resolved.
return key, nil
}
51 changes: 49 additions & 2 deletions repository/private/key_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -461,8 +461,9 @@ func TestKeyRepository_ValidateKey(t *testing.T) {
keyRepo := NewKeyRepository(db, keygen)

testMarketplace := &models.Marketplace{
Slug: "test-marketplace",
Name: "Test Marketplace",
Slug: "test-marketplace",
Name: "Test Marketplace",
IsActive: true,
}
testutil.Insert(t, db, testMarketplace)

Expand Down Expand Up @@ -521,3 +522,49 @@ func TestKeyRepository_ValidateKey_Errors(t *testing.T) {
t.Fatalf("ValidateKey(): got error %v, wanted %v", err, gorm.ErrRecordNotFound)
}
}

func TestKeyRepository_ValidateKey_InactiveMarketplace(t *testing.T) {
t.Parallel()

db := testutil.NewTestDB(t)
keygen := secret.NewKeyGenerator(constants.EnvironmentDevelopment)
keyRepo := NewKeyRepository(db, keygen)

testMarketplace := &models.Marketplace{
Slug: "test-marketplace",
Name: "Test Marketplace",
IsActive: false,
}
testutil.Insert(t, db, testMarketplace)

secretKey, err := keygen.GenerateSecretKey()
if err != nil {
t.Fatalf("GenerateSecretKey(): %v", err)
}

id, err := secretKey.ID()
if err != nil {
t.Fatalf("ID(): %v", err)
}

testKey := &models.Key{
ID: id,
Environment: keygen.Environment(),
MarketplaceSlug: testMarketplace.Slug,
Permissions: models.PermissionRead,
}
testutil.Insert(t, db, testKey)

formattedKey, err := secretKey.Format()
if err != nil {
t.Fatalf("Format(): %v", err)
}

_, err = keyRepo.ValidateKey(formattedKey)
if err == nil {
t.Fatal("ValidateKey(): got nil error, wanted error")
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("ValidateKey(): got error %v, wanted %v", err, gorm.ErrRecordNotFound)
}
}
Loading