This commit is contained in:
@@ -0,0 +1,319 @@
|
||||
package cliproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func authResponse(items string) string {
|
||||
return `{"auth_files":` + items + `}`
|
||||
}
|
||||
|
||||
func validAuth(index string) string {
|
||||
return `{"auth_index":"` + index + `","provider":"codex","type":"codex","label":"Main","email":"user@example.com","status":"ready","id_token":{"chatgpt_account_id":"acct-123","plan_type":"new_unknown_plan"}}`
|
||||
}
|
||||
|
||||
func TestAuthFiltersProviderStatusAndRequiresUniqueMatch(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
want string
|
||||
}{
|
||||
{name: "valid", body: authResponse(`[` + validAuth("auth/one") + `]`), want: ""},
|
||||
{name: "wrong provider", body: authResponse(`[{"auth_index":"auth/one","provider":"gemini","type":"gemini","id_token":{"chatgpt_account_id":"acct"}}]`), want: "未找到"},
|
||||
{name: "disabled", body: authResponse(`[{"auth_index":"auth/one","provider":"codex","disabled":true,"id_token":{"chatgpt_account_id":"acct"}}]`), want: "未找到"},
|
||||
{name: "quota unavailable remains readable", body: authResponse(`[{"auth_index":"auth/one","type":"codex","status":"error","unavailable":true,"id_token":{"chatgpt_account_id":"acct-123","plan_type":"new_unknown_plan"}}]`), want: ""},
|
||||
{name: "duplicate", body: authResponse(`[` + validAuth("auth/one") + `,` + validAuth("auth/one") + `]`), want: "不唯一"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if got := r.URL.Query().Get("auth_index"); got != "auth/one" {
|
||||
t.Fatalf("auth_index = %q", got)
|
||||
}
|
||||
_, _ = w.Write([]byte(tt.body))
|
||||
}))
|
||||
defer server.Close()
|
||||
client := New(server.URL, "management-secret")
|
||||
auth, err := client.Auth(context.Background(), "auth/one")
|
||||
if tt.want == "" {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if auth.AuthIndex != "auth/one" || auth.AccountID != "acct-123" || auth.PlanType == nil || *auth.PlanType != "new_unknown_plan" {
|
||||
t.Fatalf("auth = %#v", auth)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), tt.want) {
|
||||
t.Fatalf("error = %v; want %q", err, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSnapshotUsesManagementAuthorizationAndAPICallRequestStructure(t *testing.T) {
|
||||
var mu sync.Mutex
|
||||
var upstreamURLs []string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer management-secret" {
|
||||
t.Fatalf("Authorization = %q", got)
|
||||
}
|
||||
switch r.URL.Path {
|
||||
case "/v0/management/auth-files":
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth 1") + `]`)))
|
||||
case "/v0/management/api-call":
|
||||
var request struct {
|
||||
AuthIndex string `json:"auth_index"`
|
||||
Method string `json:"method"`
|
||||
URL string `json:"url"`
|
||||
Header map[string]string `json:"header"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if request.AuthIndex != "auth 1" || request.Method != http.MethodGet || request.Header["Authorization"] != "Bearer $TOKEN$" || request.Header["Chatgpt-Account-Id"] != "acct-123" || request.Header["Accept"] != "application/json" || request.Header["Content-Type"] != "application/json" || request.Header["User-Agent"] != codexUserAgent {
|
||||
t.Fatalf("api-call request = %#v", request)
|
||||
}
|
||||
mu.Lock()
|
||||
upstreamURLs = append(upstreamURLs, request.URL)
|
||||
mu.Unlock()
|
||||
var body string
|
||||
switch request.URL {
|
||||
case usageURL:
|
||||
body = `{"rate_limit":{"primary_window":{"used_percent":12.75,"limit_window_seconds":301,"reset_after_seconds":60}}}`
|
||||
case profileURL:
|
||||
body = `{"lifetime_tokens":1000,"peak_daily_tokens":250,"longest_running_turn_sec":90,"current_streak_days":3,"longest_streak_days":8,"daily_usage_buckets":[{"start_date":"2026-08-14","tokens":77}]}`
|
||||
case resetCreditsURL:
|
||||
body = `{"available_count":1,"credits":[{"expires_at":"2026-08-20T00:00:00Z"}]}`
|
||||
default:
|
||||
t.Fatalf("unexpected upstream URL %q", request.URL)
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": body})
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := New(server.URL, "management-secret")
|
||||
fixedNow := time.Date(2026, 8, 14, 12, 0, 0, 0, time.UTC)
|
||||
client.now = func() time.Time { return fixedNow }
|
||||
snapshot, err := client.Snapshot(context.Background(), "auth 1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(snapshot.Limits) != 1 || snapshot.Limits[0].UsedPercent != 12.75 || snapshot.Limits[0].WindowDurationMinutes != 6 || snapshot.Limits[0].ResetsAt != fixedNow.Unix()+60 {
|
||||
t.Fatalf("limits = %#v", snapshot.Limits)
|
||||
}
|
||||
if snapshot.Limits[0].PlanType == nil || *snapshot.Limits[0].PlanType != "new_unknown_plan" {
|
||||
t.Fatalf("plan type was not preserved: %#v", snapshot.Limits[0].PlanType)
|
||||
}
|
||||
if snapshot.Summary.LifetimeTokens == nil || *snapshot.Summary.LifetimeTokens != 1000 || snapshot.Summary.PeakDailyTokens == nil || *snapshot.Summary.PeakDailyTokens != 250 || len(snapshot.Usage) != 1 || snapshot.Usage[0].TotalTokens != 77 {
|
||||
t.Fatalf("profile = %#v usage = %#v", snapshot.Summary, snapshot.Usage)
|
||||
}
|
||||
if snapshot.ResetCredits == nil || snapshot.ResetCredits.AvailableCount != 1 || len(snapshot.ResetCredits.ExpiresAt) != 1 || snapshot.ResetCredits.ExpiresAt[0] != time.Date(2026, 8, 20, 0, 0, 0, 0, time.UTC).Unix() {
|
||||
t.Fatalf("reset credits = %#v", snapshot.ResetCredits)
|
||||
}
|
||||
if len(upstreamURLs) != 3 {
|
||||
t.Fatalf("upstream URLs = %#v", upstreamURLs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageParsesAllLimitKindsAndClampsPercentages(t *testing.T) {
|
||||
plan := "enterprise-new"
|
||||
body := []byte(`{
|
||||
"rate_limit":{"primary_window":{"used_percent":-1.25,"limit_window_seconds":300,"reset_at":1800000000}},
|
||||
"code_review_rate_limit":{"secondary_window":{"used_percent":101.5,"limit_window_seconds":604801,"reset_at":1800000100}},
|
||||
"additional_rate_limits":[{"metered_feature":"spark","limit_name":"Spark usage","primary_window":{"used_percent":45.5,"limit_window_seconds":61,"reset_at":1800000200}}],
|
||||
"rate_limit_reset_credits":{"available_count":2}
|
||||
}`)
|
||||
limits, fallback, effectivePlan, err := parseUsage(body, &plan, time.Now())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(limits) != 3 || limits[0].LimitID != "codex" || limits[0].UsedPercent != 0 || limits[1].LimitID != "code_review" || limits[1].UsedPercent != 100 || limits[1].WindowDurationMinutes != 10081 || limits[2].LimitID != "spark" || limits[2].LimitName == nil || *limits[2].LimitName != "Spark usage" || limits[2].WindowDurationMinutes != 2 || fallback != 2 || effectivePlan == nil || *effectivePlan != plan {
|
||||
t.Fatalf("limits = %#v fallback = %d", limits, fallback)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageSkipsWindowsWithoutUsedPercentage(t *testing.T) {
|
||||
limits, _, _, err := parseUsage([]byte(`{"rate_limit":{"primary_window":{"limit_window_seconds":18000,"reset_at":1800000000}}}`), nil, time.Now())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(limits) != 0 {
|
||||
t.Fatalf("limits = %#v; missing used_percent must not become zero usage", limits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageReadsNestedAdditionalLimitsAndPlanFromPayload(t *testing.T) {
|
||||
body := []byte(`{"data":{"plan_type":"prolite","additional_rate_limits":[{"metered_feature":"codex_bengalfox","limit_name":"GPT-5.3-Codex-Spark","rate_limit":{"primary_window":{"used_percent":9.5,"limit_window_seconds":18000,"reset_at":1800000000}}}]}}`)
|
||||
limits, _, plan, err := parseUsage(body, nil, time.Now())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(limits) != 1 || limits[0].LimitID != "codex_bengalfox" || limits[0].UsedPercent != 9.5 || plan == nil || *plan != "prolite" || limits[0].PlanType == nil || *limits[0].PlanType != "prolite" {
|
||||
t.Fatalf("limits = %#v plan = %v", limits, plan)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileMissingOrNullBucketsAreNotMarkedAvailable(t *testing.T) {
|
||||
for _, body := range []string{`{}`, `{"stats":{}}`, `{"stats":{"daily_usage_buckets":null}}`} {
|
||||
summary, usage, usageAvailable, err := parseProfile([]byte(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if usageAvailable || usage == nil || len(usage) != 0 || usageSummaryAvailable(summary) {
|
||||
t.Fatalf("body=%s summary=%#v usage=%#v available=%v", body, summary, usage, usageAvailable)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileInvalidOptionalMetricsRemainUnavailable(t *testing.T) {
|
||||
summary, usage, usageAvailable, err := parseProfile([]byte(`{"stats":{"lifetime_tokens":"unknown","daily_usage_buckets":[{"start_date":"2026-08-14","tokens":"bad"}]}}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if summary.LifetimeTokens != nil || !usageAvailable || len(usage) != 0 {
|
||||
t.Fatalf("summary = %#v usage = %#v", summary, usage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetCreditsIgnoreNonAvailableDetails(t *testing.T) {
|
||||
credits := parseResetCredits([]byte(`{"available_count":2,"credits":[{"status":"redeemed","expires_at":"2026-08-19T00:00:00Z"},{"status":"available","expires_at":"2026-08-20T00:00:00Z"}]}`), 0)
|
||||
if credits == nil || credits.AvailableCount != 2 || len(credits.ExpiresAt) != 1 || credits.ExpiresAt[0] != time.Date(2026, 8, 20, 0, 0, 0, 0, time.UTC).Unix() {
|
||||
t.Fatalf("credits = %#v", credits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileFailureDoesNotHideRateLimits(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v0/management/auth-files" {
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth-1") + `]`)))
|
||||
return
|
||||
}
|
||||
var request struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
_ = json.NewDecoder(r.Body).Decode(&request)
|
||||
switch request.URL {
|
||||
case usageURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": `{"rate_limit":{"primary_window":{"used_percent":50,"limit_window_seconds":18000,"reset_at":1800000000}}}`})
|
||||
case profileURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 503, "body": `{"error":"profile unavailable"}`})
|
||||
case resetCreditsURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 404, "body": `{}`})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
snapshot, err := New(server.URL, "management-secret").Snapshot(context.Background(), "auth-1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(snapshot.Limits) != 1 || snapshot.ProfileAvailable || snapshot.Usage == nil || len(snapshot.Usage) != 0 || snapshot.Summary.LifetimeTokens != nil {
|
||||
t.Fatalf("snapshot = %#v", snapshot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetCreditsSuccessfulZeroOverridesUsageFallback(t *testing.T) {
|
||||
if credits := parseResetCredits([]byte(`{"available_count":0,"credits":[]}`), 3); credits != nil {
|
||||
t.Fatalf("credits = %#v; successful detail response is authoritative", credits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOptionalResetFailureFallsBackToUsageCount(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v0/management/auth-files" {
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth-1") + `]`)))
|
||||
return
|
||||
}
|
||||
var request struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
_ = json.NewDecoder(r.Body).Decode(&request)
|
||||
switch request.URL {
|
||||
case usageURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": `{"rate_limit_reset_credits":{"available_count":3}}`})
|
||||
case profileURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": `{}`})
|
||||
case resetCreditsURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 503, "body": `{"secret":"upstream detail"}`})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
client := New(server.URL, "management-secret")
|
||||
snapshot, err := client.Snapshot(context.Background(), "auth-1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if snapshot.ResetCredits == nil || snapshot.ResetCredits.AvailableCount != 3 || snapshot.ResetCredits.ExpiresAt == nil {
|
||||
t.Fatalf("reset credits = %#v", snapshot.ResetCredits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagementRequestDoesNotFollowRedirectsWithSecret(t *testing.T) {
|
||||
reachedRedirect := false
|
||||
target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
reachedRedirect = true
|
||||
if r.Header.Get("Authorization") != "" {
|
||||
t.Fatal("management Authorization header reached redirect target")
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer target.Close()
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, target.URL, http.StatusFound)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := New(server.URL, "management-secret")
|
||||
_, err := client.Auth(context.Background(), "auth-1")
|
||||
if err == nil || !strings.Contains(err.Error(), "302") {
|
||||
t.Fatalf("error = %v", err)
|
||||
}
|
||||
if reachedRedirect {
|
||||
t.Fatal("redirect was unexpectedly followed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeUpstreamErrorTypeOnlyReturnsBoundedIdentifiers(t *testing.T) {
|
||||
if got := safeUpstreamErrorType(json.RawMessage(`"{\"error\":{\"type\":\"token_expired\"}}"`)); got != "token_expired" {
|
||||
t.Fatalf("type = %q", got)
|
||||
}
|
||||
if got := safeUpstreamErrorType(json.RawMessage(`{"error":{"type":"secret bearer value"}}`)); got != "" {
|
||||
t.Fatalf("unsafe type = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpstreamErrorDoesNotLeakBodyOrManagementKey(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v0/management/auth-files" {
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth-1") + `]`)))
|
||||
return
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 401, "body": `{"access_token":"raw-token","detail":"private detail"}`})
|
||||
}))
|
||||
defer server.Close()
|
||||
client := New(server.URL, "management-secret")
|
||||
_, err := client.Snapshot(context.Background(), "auth-1")
|
||||
if err == nil {
|
||||
t.Fatal("snapshot unexpectedly succeeded")
|
||||
}
|
||||
message := err.Error()
|
||||
for _, secret := range []string{"raw-token", "private detail", "management-secret"} {
|
||||
if strings.Contains(message, secret) {
|
||||
t.Fatalf("error leaked %q: %s", secret, message)
|
||||
}
|
||||
}
|
||||
if !strings.Contains(message, "401") {
|
||||
t.Fatalf("error = %q", message)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user