package service import ( "context" "encoding/base64" "net/http" "net/http/httptest" "net/url" "strings" "testing" "time" "github.com/QuantumNous/new-api/common" ) func TestRefreshCodexOAuthTokenReportsHTTPStatusForNonJSONError(t *testing.T) { t.Parallel() server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusUnauthorized) _, _ = w.Write([]byte("upstream unavailable")) })) defer server.Close() _, err := refreshCodexOAuthToken(context.Background(), server.Client(), server.URL, "client-id", "refresh-token") if err == nil { t.Fatal("expected refresh failure") } if !strings.Contains(err.Error(), "status=401") { t.Fatalf("expected error to contain upstream status, got %q", err) } } func TestRefreshCodexOAuthTokenSendsRefreshGrantAndParsesResponse(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost || r.Header.Get("Content-Type") != "application/x-www-form-urlencoded" { t.Fatalf("unexpected request: method=%s content-type=%q", r.Method, r.Header.Get("Content-Type")) } if err := r.ParseForm(); err != nil { t.Fatalf("parse form: %v", err) } if r.Form.Get("grant_type") != "refresh_token" || r.Form.Get("client_id") != "client-id" || r.Form.Get("refresh_token") != "refresh-token" { t.Fatalf("unexpected refresh form: %#v", r.Form) } _, _ = w.Write([]byte(`{"access_token":"access","refresh_token":"next-refresh","expires_in":60}`)) })) defer server.Close() before := time.Now() result, err := refreshCodexOAuthToken(context.Background(), server.Client(), server.URL, "client-id", " refresh-token ") if err != nil { t.Fatalf("refresh token: %v", err) } if result.AccessToken != "access" || result.RefreshToken != "next-refresh" || !result.ExpiresAt.After(before.Add(59*time.Second)) { t.Fatalf("unexpected refresh result: %#v", result) } } func TestExchangeCodexAuthorizationCodeSendsPKCEForm(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if err := r.ParseForm(); err != nil { t.Fatalf("parse form: %v", err) } if r.Form.Get("grant_type") != "authorization_code" || r.Form.Get("code") != "code" || r.Form.Get("code_verifier") != "verifier" || r.Form.Get("redirect_uri") != "http://localhost/callback" { t.Fatalf("unexpected authorization-code form: %#v", r.Form) } _, _ = w.Write([]byte(`{"access_token":"access","refresh_token":"refresh","expires_in":60}`)) })) defer server.Close() result, err := exchangeCodexAuthorizationCode(context.Background(), server.Client(), server.URL, "client-id", " code ", " verifier ", "http://localhost/callback") if err != nil { t.Fatalf("exchange authorization code: %v", err) } if result.AccessToken != "access" || result.RefreshToken != "refresh" { t.Fatalf("unexpected exchange result: %#v", result) } } func TestExtractCodexClaimsFromJWT(t *testing.T) { claims := map[string]any{ "email": " user@example.com ", codexJWTClaimPath: map[string]any{ "chatgpt_account_id": " account-123 ", }, } token := newTestJWT(t, claims) accountID, ok := ExtractCodexAccountIDFromJWT(token) if !ok || accountID != "account-123" { t.Fatalf("expected trimmed account ID, got %q, %t", accountID, ok) } email, ok := ExtractEmailFromJWT(token) if !ok || email != "user@example.com" { t.Fatalf("expected trimmed email, got %q, %t", email, ok) } } func TestExtractCodexClaimsRejectMalformedOrEmptyValues(t *testing.T) { if _, ok := ExtractCodexAccountIDFromJWT("not-a-jwt"); ok { t.Fatal("malformed token must not yield account ID") } if _, ok := ExtractEmailFromJWT(newTestJWT(t, map[string]any{"email": " "})); ok { t.Fatal("empty email must not be accepted") } if _, ok := ExtractCodexAccountIDFromJWT(newTestJWT(t, map[string]any{ codexJWTClaimPath: map[string]any{"chatgpt_account_id": ""}, })); ok { t.Fatal("empty account ID must not be accepted") } } func TestCreateCodexOAuthAuthorizationFlowUsesPKCEAndState(t *testing.T) { flow, err := CreateCodexOAuthAuthorizationFlow() if err != nil { t.Fatalf("create authorization flow: %v", err) } if len(flow.State) != 32 { t.Fatalf("expected 32-character state, got %q", flow.State) } if flow.Verifier == "" || flow.Challenge == "" { t.Fatal("expected PKCE verifier and challenge") } u, err := url.Parse(flow.AuthorizeURL) if err != nil { t.Fatalf("parse authorize URL: %v", err) } q := u.Query() if q.Get("state") != flow.State || q.Get("code_challenge") != flow.Challenge { t.Fatalf("authorize URL does not include generated state and challenge: %s", flow.AuthorizeURL) } if q.Get("code_challenge_method") != "S256" || q.Get("redirect_uri") != codexOAuthRedirectURI { t.Fatalf("unexpected PKCE or redirect parameters: %s", flow.AuthorizeURL) } } func TestFetchCodexWhamUsageSendsRequiredHeaders(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/backend-api/wham/usage" { t.Fatalf("unexpected path: %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer access-token" { t.Fatalf("unexpected authorization: %q", r.Header.Get("Authorization")) } if r.Header.Get("chatgpt-account-id") != "account-123" || r.Header.Get("originator") != "codex_cli_rs" { t.Fatalf("missing Codex headers: %#v", r.Header) } w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"limit": 1}`)) })) defer server.Close() status, body, err := FetchCodexWhamUsage(context.Background(), server.Client(), server.URL+"/", " access-token ", " account-123 ") if err != nil { t.Fatalf("fetch usage: %v", err) } if status != http.StatusOK || string(body) != `{"limit": 1}` { t.Fatalf("unexpected usage response: status=%d body=%q", status, body) } } func TestFetchCodexWhamUsageRejectsMissingInputs(t *testing.T) { tests := []struct { name string client *http.Client baseURL string accessKey string accountID string }{ {name: "nil client", baseURL: "https://example.com", accessKey: "token", accountID: "account"}, {name: "empty base URL", client: http.DefaultClient, accessKey: "token", accountID: "account"}, {name: "empty access token", client: http.DefaultClient, baseURL: "https://example.com", accountID: "account"}, {name: "empty account ID", client: http.DefaultClient, baseURL: "https://example.com", accessKey: "token"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if _, _, err := FetchCodexWhamUsage(context.Background(), tt.client, tt.baseURL, tt.accessKey, tt.accountID); err == nil { t.Fatal("expected input validation error") } }) } } func newTestJWT(t *testing.T, claims map[string]any) string { t.Helper() payload, err := common.Marshal(claims) if err != nil { t.Fatalf("marshal claims: %v", err) } return "header." + base64.RawURLEncoding.EncodeToString(payload) + ".signature" }