|
- 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"
- }
|