diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/limits/limits.go | 116 | ||||
| -rw-r--r-- | internal/limits/limits_test.go | 38 |
2 files changed, 114 insertions, 40 deletions
diff --git a/internal/limits/limits.go b/internal/limits/limits.go index 6e0bf74..255e5e0 100644 --- a/internal/limits/limits.go +++ b/internal/limits/limits.go @@ -64,20 +64,37 @@ const ( ) const acquireScript = ` -local req = tonumber(ARGV[1]) -local tok = tonumber(ARGV[2]) -local conc = tonumber(ARGV[3]) -local estimate = tonumber(ARGV[4]) -local ttl = tonumber(ARGV[5]) -local r = 0 -local t = 0 +local project_req = tonumber(ARGV[1]) +local project_tok = tonumber(ARGV[2]) +local project_conc = tonumber(ARGV[3]) +local key_req = tonumber(ARGV[4]) +local key_tok = tonumber(ARGV[5]) +local estimate = tonumber(ARGV[6]) +local ttl = tonumber(ARGV[7]) +local pr = 0 +local pt = 0 local c = 0 -if req > 0 then r = redis.call('INCR', KEYS[1]); if r == 1 then redis.call('PEXPIRE', KEYS[1], ttl) end end -if tok > 0 then t = redis.call('INCRBY', KEYS[2], estimate); if t == estimate then redis.call('PEXPIRE', KEYS[2], ttl) end end -if conc > 0 then c = redis.call('INCR', KEYS[3]); redis.call('PEXPIRE', KEYS[3], 3600000) end -if (req > 0 and r > req) then if req > 0 then redis.call('DECR', KEYS[1]) end; if tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end; if conc > 0 then redis.call('DECR', KEYS[3]) end; return {0,1} end -if (tok > 0 and t > tok) then if req > 0 then redis.call('DECR', KEYS[1]) end; if tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end; if conc > 0 then redis.call('DECR', KEYS[3]) end; return {0,2} end -if (conc > 0 and c > conc) then if req > 0 then redis.call('DECR', KEYS[1]) end; if tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end; if conc > 0 then redis.call('DECR', KEYS[3]) end; return {0,3} end +local kr = 0 +local kt = 0 +if project_req > 0 then pr = redis.call('INCR', KEYS[1]); if pr == 1 then redis.call('PEXPIRE', KEYS[1], ttl) end end +if project_tok > 0 then pt = redis.call('INCRBY', KEYS[2], estimate); if pt == estimate then redis.call('PEXPIRE', KEYS[2], ttl) end end +if project_conc > 0 then c = redis.call('INCR', KEYS[3]); redis.call('PEXPIRE', KEYS[3], 3600000) end +if key_req > 0 then kr = redis.call('INCR', KEYS[4]); if kr == 1 then redis.call('PEXPIRE', KEYS[4], ttl) end end +if key_tok > 0 then kt = redis.call('INCRBY', KEYS[5], estimate); if kt == estimate then redis.call('PEXPIRE', KEYS[5], ttl) end end +local reason = 0 +if project_req > 0 and pr > project_req then reason = 1 +elseif project_tok > 0 and pt > project_tok then reason = 2 +elseif project_conc > 0 and c > project_conc then reason = 3 +elseif key_req > 0 and kr > key_req then reason = 4 +elseif key_tok > 0 and kt > key_tok then reason = 5 end +if reason > 0 then + if project_req > 0 then redis.call('DECR', KEYS[1]) end + if project_tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end + if project_conc > 0 then redis.call('DECR', KEYS[3]) end + if key_req > 0 then redis.call('DECR', KEYS[4]) end + if key_tok > 0 then redis.call('DECRBY', KEYS[5], estimate) end + return {0,reason} +end return {1,0} ` @@ -139,15 +156,15 @@ func (l *Limiter) Policy(projectID string) (domain.LimitPolicy, bool) { } func (l *Limiter) Acquire(ctx context.Context, principal domain.Principal, body []byte) (Lease, error) { - policy, ok := l.Policy(principal.ProjectID) - if !ok || (policy.RequestsPerMinute == 0 && policy.TokensPerMinute == 0 && policy.Concurrent == 0) { + policy, _ := l.Policy(principal.ProjectID) + if policy.RequestsPerMinute == 0 && policy.TokensPerMinute == 0 && policy.Concurrent == 0 && principal.RequestsPerMinute == 0 && principal.TokensPerMinute == 0 { return noopLease{}, nil } estimate := EstimateTokens(body, l.defaultMaxOutput) minute := time.Now().Unix() / 60 if l.redis != nil && time.Now().UnixNano() >= l.redisRetryAt.Load() { redisContext, cancel := context.WithTimeout(ctx, redisCommandTimeout) - lease, err := l.acquireRedis(redisContext, principal.ProjectID, minute, policy, estimate) + lease, err := l.acquireRedis(redisContext, principal, minute, policy, estimate) cancel() if err == nil { return lease, nil @@ -158,7 +175,7 @@ func (l *Limiter) Acquire(ctx context.Context, principal domain.Principal, body } l.markRedis(false, err) } - return l.acquireLocal(principal.ProjectID, minute, policy, estimate) + return l.acquireLocal(principal, minute, policy, estimate) } func EstimateTokens(body []byte, defaultMaxOutput int64) int64 { @@ -192,10 +209,13 @@ func EstimateTokens(body []byte, defaultMaxOutput int64) int64 { return input + maxOutput } -func (l *Limiter) acquireRedis(ctx context.Context, project string, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { - base := l.prefix + ":" + project + ":" + strconv.FormatInt(minute, 10) - keys := []string{base + ":requests", base + ":tokens", l.prefix + ":" + project + ":concurrent"} - values, err := l.redis.Eval(ctx, acquireScript, keys, policy.RequestsPerMinute, policy.TokensPerMinute, policy.Concurrent, estimate, 125000).Result() +func (l *Limiter) acquireRedis(ctx context.Context, principal domain.Principal, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { + projectBase := l.prefix + ":project:" + principal.ProjectID + ":" + strconv.FormatInt(minute, 10) + keyBase := l.prefix + ":key:" + principal.KeyID + ":" + strconv.FormatInt(minute, 10) + concurrencyKey := l.prefix + ":project:" + principal.ProjectID + ":concurrent" + keys := []string{projectBase + ":requests", projectBase + ":tokens", concurrencyKey, keyBase + ":requests", keyBase + ":tokens"} + values, err := l.redis.Eval(ctx, acquireScript, keys, policy.RequestsPerMinute, policy.TokensPerMinute, policy.Concurrent, + principal.RequestsPerMinute, principal.TokensPerMinute, estimate, 125000).Result() if err != nil { return nil, err } @@ -207,19 +227,19 @@ func (l *Limiter) acquireRedis(ctx context.Context, project string, minute int64 reason, _ := toInt64(items[1]) if allowed == 0 { switch reason { - case 1: + case 1, 4: return nil, ErrRequestsExceeded - case 2: + case 2, 5: return nil, ErrTokensExceeded default: return nil, ErrConcurrencyLimit } } l.markRedis(true, nil) - return &redisLease{limiter: l, key: keys[2]}, nil + return &redisLease{limiter: l, key: concurrencyKey}, nil } -func (l *Limiter) acquireLocal(project string, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { +func (l *Limiter) acquireLocal(principal domain.Principal, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { l.localMu.Lock() defer l.localMu.Unlock() for key, value := range l.local { @@ -227,28 +247,44 @@ func (l *Limiter) acquireLocal(project string, minute int64, policy domain.Limit delete(l.local, key) } } - window := l.local[project] + projectKey := "project:" + principal.ProjectID + keyKey := "key:" + principal.KeyID + projectWindow := l.localWindow(projectKey, minute) + keyWindow := l.localWindow(keyKey, minute) + if policy.RequestsPerMinute > 0 && projectWindow.requests >= policy.RequestsPerMinute { + return nil, ErrRequestsExceeded + } + if policy.TokensPerMinute > 0 && (estimate > policy.TokensPerMinute || projectWindow.tokens > policy.TokensPerMinute-estimate) { + return nil, ErrTokensExceeded + } + if policy.Concurrent > 0 && projectWindow.concurrent >= policy.Concurrent { + return nil, ErrConcurrencyLimit + } + if principal.RequestsPerMinute > 0 && keyWindow.requests >= principal.RequestsPerMinute { + return nil, ErrRequestsExceeded + } + if principal.TokensPerMinute > 0 && (estimate > principal.TokensPerMinute || keyWindow.tokens > principal.TokensPerMinute-estimate) { + return nil, ErrTokensExceeded + } + projectWindow.requests++ + projectWindow.tokens += estimate + projectWindow.concurrent++ + keyWindow.requests++ + keyWindow.tokens += estimate + return &localLease{limiter: l, project: projectKey}, nil +} + +func (l *Limiter) localWindow(key string, minute int64) *localWindow { + window := l.local[key] if window == nil { window = &localWindow{minute: minute} - l.local[project] = window + l.local[key] = window } else if window.minute != minute { window.minute = minute window.requests = 0 window.tokens = 0 } - if policy.RequestsPerMinute > 0 && window.requests >= policy.RequestsPerMinute { - return nil, ErrRequestsExceeded - } - if policy.TokensPerMinute > 0 && (estimate > policy.TokensPerMinute || window.tokens > policy.TokensPerMinute-estimate) { - return nil, ErrTokensExceeded - } - if policy.Concurrent > 0 && window.concurrent >= policy.Concurrent { - return nil, ErrConcurrencyLimit - } - window.requests++ - window.tokens += estimate - window.concurrent++ - return &localLease{limiter: l, project: project}, nil + return window } func (l *Limiter) releaseLocal(project string) { diff --git a/internal/limits/limits_test.go b/internal/limits/limits_test.go index 78b346e..8e48d42 100644 --- a/internal/limits/limits_test.go +++ b/internal/limits/limits_test.go @@ -56,6 +56,44 @@ func TestLocalTokenAndConcurrencyLimits(t *testing.T) { retry.Release() } +func TestLocalKeyLimitDoesNotConsumeProjectQuotaWhenRejected(t *testing.T) { + limiter := New("", "test", 0, nil) + limiter.ReplacePolicies([]domain.LimitPolicy{{ProjectID: "project-1", RequestsPerMinute: 2}}) + firstKey := domain.Principal{ProjectID: "project-1", KeyID: "key-1", RequestsPerMinute: 1} + lease, err := limiter.Acquire(context.Background(), firstKey, []byte(`{}`)) + if err != nil { + t.Fatal(err) + } + lease.Release() + if _, err := limiter.Acquire(context.Background(), firstKey, []byte(`{}`)); !errors.Is(err, ErrRequestsExceeded) { + t.Fatalf("second key request error = %v, want request limit", err) + } + secondKey := domain.Principal{ProjectID: "project-1", KeyID: "key-2", RequestsPerMinute: 1} + lease, err = limiter.Acquire(context.Background(), secondKey, []byte(`{}`)) + if err != nil { + t.Fatalf("key rejection consumed project quota: %v", err) + } + lease.Release() + if _, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "project-1", KeyID: "key-3"}, []byte(`{}`)); !errors.Is(err, ErrRequestsExceeded) { + t.Fatalf("project request limit error = %v", err) + } +} + +func TestLocalKeyTokenLimit(t *testing.T) { + limiter := New("", "test", 0, nil) + body := []byte(`{"max_tokens":4}`) + estimate := EstimateTokens(body, 0) + principal := domain.Principal{ProjectID: "project-1", KeyID: "key-1", TokensPerMinute: estimate} + lease, err := limiter.Acquire(context.Background(), principal, body) + if err != nil { + t.Fatal(err) + } + lease.Release() + if _, err := limiter.Acquire(context.Background(), principal, body); !errors.Is(err, ErrTokensExceeded) { + t.Fatalf("second key token request error = %v, want token limit", err) + } +} + func TestEstimateTokensUsesLargestExplicitOutputLimit(t *testing.T) { body := []byte(`{"max_tokens":10,"max_completion_tokens":25}`) want := int64((len(body)+3)/4 + 25) |
