summaryrefslogtreecommitdiff
path: root/internal/limits/limits_test.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 00:26:25 +1200
committerChia <Chia@93.nz>2026-08-05 00:33:31 +1200
commit1a3d7f9a8a181df48f0e911cbe17a3fad3ab9ac9 (patch)
tree8c92e1e7326fc67ed077a0a878697f1be14b43da /internal/limits/limits_test.go
parent5b651488b081b65fda8a323f228e139adb79a35d (diff)
add some scriptsmain
Diffstat (limited to 'internal/limits/limits_test.go')
-rw-r--r--internal/limits/limits_test.go103
1 files changed, 103 insertions, 0 deletions
diff --git a/internal/limits/limits_test.go b/internal/limits/limits_test.go
new file mode 100644
index 0000000..78b346e
--- /dev/null
+++ b/internal/limits/limits_test.go
@@ -0,0 +1,103 @@
+package limits
+
+import (
+ "context"
+ "errors"
+ "testing"
+ "time"
+
+ "aigw/internal/domain"
+)
+
+func TestLocalRequestLimit(t *testing.T) {
+ limiter := New("", "test", 0, nil)
+ limiter.ReplacePolicies([]domain.LimitPolicy{{ProjectID: "project-1", RequestsPerMinute: 2}})
+ principal := domain.Principal{ProjectID: "project-1"}
+ for i := 0; i < 2; i++ {
+ lease, err := limiter.Acquire(context.Background(), principal, []byte(`{}`))
+ if err != nil {
+ t.Fatal(err)
+ }
+ lease.Release()
+ }
+ if _, err := limiter.Acquire(context.Background(), principal, []byte(`{}`)); !errors.Is(err, ErrRequestsExceeded) {
+ t.Fatalf("error = %v, want request limit", err)
+ }
+}
+
+func TestLocalTokenAndConcurrencyLimits(t *testing.T) {
+ limiter := New("", "test", 0, nil)
+ body := []byte(`{"max_tokens":4}`)
+ estimate := EstimateTokens(body, 0)
+ limiter.ReplacePolicies([]domain.LimitPolicy{
+ {ProjectID: "tokens", TokensPerMinute: estimate},
+ {ProjectID: "concurrency", Concurrent: 1},
+ })
+ lease, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "tokens"}, body)
+ if err != nil {
+ t.Fatal(err)
+ }
+ lease.Release()
+ if _, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "tokens"}, body); !errors.Is(err, ErrTokensExceeded) {
+ t.Fatalf("error = %v, want token limit", err)
+ }
+ held, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "concurrency"}, body)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "concurrency"}, body); !errors.Is(err, ErrConcurrencyLimit) {
+ t.Fatalf("error = %v, want concurrency limit", err)
+ }
+ held.Release()
+ retry, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "concurrency"}, body)
+ if err != nil {
+ t.Fatalf("acquire after release: %v", err)
+ }
+ retry.Release()
+}
+
+func TestEstimateTokensUsesLargestExplicitOutputLimit(t *testing.T) {
+ body := []byte(`{"max_tokens":10,"max_completion_tokens":25}`)
+ want := int64((len(body)+3)/4 + 25)
+ if got := EstimateTokens(body, 4); got != want {
+ t.Fatalf("EstimateTokens = %d, want %d", got, want)
+ }
+}
+
+func TestEstimateTokensHonorsExplicitLimitBelowDefault(t *testing.T) {
+ body := []byte(`{"max_tokens":10}`)
+ want := int64((len(body)+3)/4 + 10)
+ if got := EstimateTokens(body, 4096); got != want {
+ t.Fatalf("EstimateTokens = %d, want %d", got, want)
+ }
+}
+
+func TestRedisFailureFallsBackQuicklyAndOpensCircuit(t *testing.T) {
+ limiter := New("redis://127.0.0.1:1/0", "test", 0, nil)
+ defer limiter.Close()
+ limiter.ReplacePolicies([]domain.LimitPolicy{{ProjectID: "project-1", RequestsPerMinute: 2}})
+ principal := domain.Principal{ProjectID: "project-1"}
+
+ started := time.Now()
+ lease, err := limiter.Acquire(context.Background(), principal, []byte(`{}`))
+ if err != nil {
+ t.Fatal(err)
+ }
+ lease.Release()
+ if elapsed := time.Since(started); elapsed > time.Second {
+ t.Fatalf("Redis fallback took %v, want under 1s", elapsed)
+ }
+ if limiter.redisRetryAt.Load() <= time.Now().UnixNano() {
+ t.Fatal("Redis failure did not open the retry circuit")
+ }
+
+ started = time.Now()
+ lease, err = limiter.Acquire(context.Background(), principal, []byte(`{}`))
+ if err != nil {
+ t.Fatal(err)
+ }
+ lease.Release()
+ if elapsed := time.Since(started); elapsed > 100*time.Millisecond {
+ t.Fatalf("open-circuit local fallback took %v, want under 100ms", elapsed)
+ }
+}