diff --git a/README.md b/README.md index 26a35c0..15627e0 100644 --- a/README.md +++ b/README.md @@ -44,6 +44,7 @@ Config file: `~/.clawlet/config.json` clawlet currently supports these LLM providers: - **OpenAI** (`openai/`, API key: `env.OPENAI_API_KEY`) +- **OpenAI Codex (OAuth)** (`openai-codex/`, no API key; login: `clawlet provider login openai-codex`) - **OpenRouter** (`openrouter//`, API key: `env.OPENROUTER_API_KEY`) - **Anthropic** (`anthropic/`, API key: `env.ANTHROPIC_API_KEY`) - **Gemini** (`gemini/`, API key: `env.GEMINI_API_KEY` or `env.GOOGLE_API_KEY`) @@ -89,7 +90,21 @@ Minimal config (Local via vLLM using the same `ollama/` route): } ``` -clawlet will fill in sensible defaults for missing sections (tools, gateway, cron, heartbeat, channels). +OpenAI Codex (OAuth): + +```bash +# one-time login +clawlet provider login openai-codex + +# headless environment (SSH / container) +clawlet provider login openai-codex --device-code +``` + +```json +{ + "agents": { "defaults": { "model": "openai-codex/gpt-5.1-codex" } } +} +``` ### Option: Memory search setup diff --git a/cmd/clawlet/cmd_provider.go b/cmd/clawlet/cmd_provider.go new file mode 100644 index 0000000..ac090eb --- /dev/null +++ b/cmd/clawlet/cmd_provider.go @@ -0,0 +1,65 @@ +package main + +import ( + "context" + "fmt" + + "github.com/mosaxiv/clawlet/llm" + "github.com/urfave/cli/v3" +) + +const oauthProviderOpenAICodex = "openai-codex" + +func cmdProvider() *cli.Command { + return &cli.Command{ + Name: "provider", + Usage: "provider authentication utilities", + Commands: []*cli.Command{ + { + Name: "login", + Usage: "authenticate an OAuth provider", + ArgsUsage: "", + Flags: []cli.Flag{ + &cli.BoolFlag{ + Name: "device-code", + Usage: "use OAuth device code flow (for headless environments)", + }, + }, + Action: func(ctx context.Context, cmd *cli.Command) error { + if cmd.Args().Len() < 1 { + return cli.Exit("usage: clawlet provider login ", 2) + } + switch cmd.Args().Get(0) { + case oauthProviderOpenAICodex: + return loginOpenAICodex(ctx, cmd.Bool("device-code")) + default: + return cli.Exit(fmt.Sprintf("unsupported oauth provider: %s (supported: %s)", cmd.Args().Get(0), oauthProviderOpenAICodex), 1) + } + }, + }, + }, + } +} + +func loginOpenAICodex(ctx context.Context, useDeviceCode bool) error { + if tok, err := llm.LoadCodexOAuthToken(); err == nil && tok.Valid() { + fmt.Printf("already authenticated with OpenAI Codex (%s)\n", tok.AccountID) + return nil + } + fmt.Println("starting OpenAI Codex OAuth login...") + var err error + if useDeviceCode { + err = llm.LoginCodexOAuthDeviceCode(ctx) + } else { + err = llm.LoginCodexOAuthInteractive(ctx) + } + if err != nil { + return err + } + tok, err := llm.LoadCodexOAuthToken() + if err != nil { + return err + } + fmt.Printf("authenticated with OpenAI Codex (%s)\n", tok.AccountID) + return nil +} diff --git a/cmd/clawlet/config.go b/cmd/clawlet/config.go index d1eb4b0..7b78715 100644 --- a/cmd/clawlet/config.go +++ b/cmd/clawlet/config.go @@ -120,7 +120,7 @@ func resolveWorkspace(flagValue string) (string, error) { func providerNeedsAPIKey(provider string) bool { switch strings.ToLower(strings.TrimSpace(provider)) { - case "ollama": + case "ollama", "openai-codex": return false default: return true diff --git a/cmd/clawlet/main.go b/cmd/clawlet/main.go index 42d8e7a..e590495 100644 --- a/cmd/clawlet/main.go +++ b/cmd/clawlet/main.go @@ -19,6 +19,7 @@ func main() { cmdStatus(), cmdAgent(), cmdGateway(), + cmdProvider(), cmdChannels(), cmdCron(), }, diff --git a/config/config.go b/config/config.go index 7e6ec63..cc37504 100644 --- a/config/config.go +++ b/config/config.go @@ -259,6 +259,7 @@ const ( DefaultMemorySearchHybridTextWeight = 0.3 DefaultMemorySearchCandidateMultiplier = 4 DefaultOpenAIBaseURL = "https://api.openai.com/v1" + DefaultOpenAICodexBaseURL = "https://chatgpt.com/backend-api" DefaultOpenRouterBaseURL = "https://openrouter.ai/api/v1" DefaultAnthropicBaseURL = "https://api.anthropic.com" DefaultGeminiBaseURL = "https://generativelanguage.googleapis.com/v1beta" @@ -541,6 +542,8 @@ func (cfg *Config) ApplyLLMRouting() (provider string, configuredModel string) { cfg.LLM.BaseURL = DefaultGeminiBaseURL case "ollama": cfg.LLM.BaseURL = DefaultOllamaBaseURL + case "openai-codex": + cfg.LLM.BaseURL = DefaultOpenAICodexBaseURL default: cfg.LLM.BaseURL = DefaultOpenAIBaseURL } @@ -570,6 +573,8 @@ func (cfg *Config) ApplyLLMRouting() (provider string, configuredModel string) { switch provider { case "openai": cfg.LLM.BaseURL = DefaultOpenAIBaseURL + case "openai-codex": + cfg.LLM.BaseURL = DefaultOpenAICodexBaseURL case "openrouter": cfg.LLM.BaseURL = DefaultOpenRouterBaseURL case "anthropic": @@ -602,6 +607,9 @@ func (cfg *Config) ApplyLLMRouting() (provider string, configuredModel string) { func parseRoutedModel(s string) (provider string, model string) { s = strings.TrimSpace(s) + if after, ok := strings.CutPrefix(s, "openai-codex/"); ok { + return "openai-codex", after + } if after, ok := strings.CutPrefix(s, "openai/"); ok { return "openai", after } diff --git a/config/config_test.go b/config/config_test.go index d8e2e60..40721f5 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -186,6 +186,27 @@ func TestApplyLLMRouting_OllamaLocal(t *testing.T) { } } +func TestApplyLLMRouting_OpenAICodex(t *testing.T) { + cfg := Default() + cfg.Agents.Defaults.Model = "openai-codex/gpt-5.1-codex" + cfg.LLM.BaseURL = "" + cfg.LLM.APIKey = "" + + provider, _ := cfg.ApplyLLMRouting() + if provider != "openai-codex" { + t.Fatalf("provider=%q", provider) + } + if cfg.LLM.BaseURL != DefaultOpenAICodexBaseURL { + t.Fatalf("baseURL=%q", cfg.LLM.BaseURL) + } + if cfg.LLM.APIKey != "" { + t.Fatalf("apiKey=%q", cfg.LLM.APIKey) + } + if cfg.LLM.Model != "gpt-5.1-codex" { + t.Fatalf("model=%q", cfg.LLM.Model) + } +} + func TestApplyLLMRouting_LocalAlias(t *testing.T) { cfg := Default() cfg.Agents.Defaults.Model = "local/qwen2.5:14b" diff --git a/llm/client.go b/llm/client.go index 1d1a48a..ccc1400 100644 --- a/llm/client.go +++ b/llm/client.go @@ -48,6 +48,8 @@ func (c *Client) Chat(ctx context.Context, messages []Message, tools []ToolDefin return c.chatAnthropic(ctx, messages, tools) case "gemini": return c.chatGemini(ctx, messages, tools) + case "openai-codex": + return c.chatOpenAICodex(ctx, messages, tools) default: return nil, fmt.Errorf("unsupported llm provider: %s", strings.TrimSpace(c.Provider)) } diff --git a/llm/openai_codex.go b/llm/openai_codex.go new file mode 100644 index 0000000..dfeb332 --- /dev/null +++ b/llm/openai_codex.go @@ -0,0 +1,464 @@ +package llm + +import ( + "bufio" + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" +) + +const ( + defaultCodexBaseURL = "https://chatgpt.com/backend-api" + defaultCodexModel = "gpt-5.2" + defaultCodexInstructions = "You are Codex, a coding assistant." +) + +type codexRequest struct { + Model string `json:"model"` + Store bool `json:"store"` + Stream bool `json:"stream"` + Instructions string `json:"instructions"` + Input []codexInputItem `json:"input"` + Text codexTextConfig `json:"text"` + Include []string `json:"include,omitempty"` + PromptCacheKey string `json:"prompt_cache_key,omitempty"` + ToolChoice string `json:"tool_choice,omitempty"` + ParallelToolCalls bool `json:"parallel_tool_calls,omitempty"` + Tools []codexTool `json:"tools,omitempty"` +} + +type codexTextConfig struct { + Verbosity string `json:"verbosity,omitempty"` +} + +type codexTool struct { + Type string `json:"type"` + Name string `json:"name"` + Description string `json:"description,omitempty"` + Parameters json.RawMessage `json:"parameters"` +} + +type codexInputItem struct { + Type string `json:"type,omitempty"` + Role string `json:"role,omitempty"` + Status string `json:"status,omitempty"` + ID string `json:"id,omitempty"` + Content []codexInputContent `json:"content,omitempty"` + CallID string `json:"call_id,omitempty"` + Name string `json:"name,omitempty"` + Arguments string `json:"arguments,omitempty"` + Output string `json:"output,omitempty"` +} + +type codexInputContent struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` +} + +func (c *Client) chatOpenAICodex(ctx context.Context, messages []Message, tools []ToolDefinition) (*ChatResult, error) { + tok, err := LoadCodexOAuthToken() + if err != nil { + return nil, err + } + + systemPrompt, inputItems := toCodexInput(messages) + if strings.TrimSpace(systemPrompt) == "" { + systemPrompt = defaultCodexInstructions + } + endpoint := codexResponsesEndpoint(c.BaseURL) + model := resolveCodexModel(c.Model) + reqBody := codexRequest{ + Model: model, + Store: false, + Stream: true, + Instructions: systemPrompt, + Input: inputItems, + Text: codexTextConfig{ + Verbosity: "medium", + }, + Include: []string{"reasoning.encrypted_content"}, + PromptCacheKey: codexPromptCacheKey(messages), + ToolChoice: "auto", + ParallelToolCalls: true, + } + + if len(tools) > 0 { + convertedTools, err := toCodexTools(tools) + if err != nil { + return nil, err + } + reqBody.Tools = convertedTools + } + + b, err := json.Marshal(reqBody) + if err != nil { + return nil, err + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(b)) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+tok.AccessToken) + req.Header.Set("chatgpt-account-id", tok.AccountID) + req.Header.Set("OpenAI-Beta", "responses=experimental") + req.Header.Set("originator", codexOAuthOriginator) + req.Header.Set("User-Agent", "clawlet (go)") + req.Header.Set("Accept", "text/event-stream") + req.Header.Set("Content-Type", "application/json") + for k, v := range c.Headers { + if strings.TrimSpace(k) == "" { + continue + } + req.Header.Set(k, v) + } + + resp, err := c.HTTP.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + raw, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<20)) + return nil, fmt.Errorf("codex http %d: %s", resp.StatusCode, codexFriendlyError(resp.StatusCode, strings.TrimSpace(string(raw)))) + } + + return consumeCodexSSE(resp.Body) +} + +type codexSSEEvent struct { + Type string `json:"type"` + Delta string `json:"delta"` + CallID string `json:"call_id"` + Arguments json.RawMessage `json:"arguments"` + Item struct { + Type string `json:"type"` + ID string `json:"id"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } `json:"item"` +} + +type codexToolCallBuffer struct { + ItemID string + Name string + Arguments string +} + +func consumeCodexSSE(r io.Reader) (*ChatResult, error) { + out := &ChatResult{} + buffers := map[string]*codexToolCallBuffer{} + + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 0, 64*1024), 2<<20) + dataLines := make([]string, 0, 2) + + flush := func() error { + if len(dataLines) == 0 { + return nil + } + data := strings.TrimSpace(strings.Join(dataLines, "\n")) + dataLines = dataLines[:0] + if data == "" || data == "[DONE]" { + return nil + } + return handleCodexSSEData(data, out, buffers) + } + + for scanner.Scan() { + line := scanner.Text() + if line == "" { + if err := flush(); err != nil { + return nil, err + } + continue + } + if after, ok := strings.CutPrefix(line, "data:"); ok { + dataLines = append(dataLines, strings.TrimSpace(after)) + } + } + if err := scanner.Err(); err != nil { + return nil, err + } + if err := flush(); err != nil { + return nil, err + } + + return out, nil +} + +func handleCodexSSEData(data string, out *ChatResult, buffers map[string]*codexToolCallBuffer) error { + var evt codexSSEEvent + if err := json.Unmarshal([]byte(data), &evt); err != nil { + // Ignore non-JSON chunks. + return nil + } + switch evt.Type { + case "response.output_text.delta": + out.Content += evt.Delta + case "response.output_item.added": + if evt.Item.Type != "function_call" { + return nil + } + callID := strings.TrimSpace(evt.Item.CallID) + if callID == "" { + return nil + } + buffers[callID] = &codexToolCallBuffer{ + ItemID: strings.TrimSpace(evt.Item.ID), + Name: strings.TrimSpace(evt.Item.Name), + Arguments: rawToCodexArgString(evt.Item.Arguments), + } + case "response.function_call_arguments.delta": + callID := strings.TrimSpace(evt.CallID) + if callID == "" { + return nil + } + buf := buffers[callID] + if buf == nil { + buf = &codexToolCallBuffer{} + buffers[callID] = buf + } + buf.Arguments += evt.Delta + case "response.function_call_arguments.done": + callID := strings.TrimSpace(evt.CallID) + if callID == "" { + return nil + } + buf := buffers[callID] + if buf == nil { + buf = &codexToolCallBuffer{} + buffers[callID] = buf + } + if args := rawToCodexArgString(evt.Arguments); args != "" { + buf.Arguments = args + } + case "response.output_item.done": + if evt.Item.Type != "function_call" { + return nil + } + callID := strings.TrimSpace(evt.Item.CallID) + if callID == "" { + return nil + } + buf := buffers[callID] + if buf == nil { + buf = &codexToolCallBuffer{} + } + if name := strings.TrimSpace(evt.Item.Name); name != "" { + buf.Name = name + } + if itemID := strings.TrimSpace(evt.Item.ID); itemID != "" { + buf.ItemID = itemID + } + if args := rawToCodexArgString(evt.Item.Arguments); args != "" { + buf.Arguments = args + } + + itemID := strings.TrimSpace(buf.ItemID) + if itemID == "" { + itemID = "fc_0" + } + out.ToolCalls = append(out.ToolCalls, ToolCall{ + ID: callID + "|" + itemID, + Name: strings.TrimSpace(buf.Name), + Arguments: codexArgumentsToJSON(buf.Arguments), + }) + delete(buffers, callID) + case "error", "response.failed": + return fmt.Errorf("codex response failed") + } + return nil +} + +func toCodexTools(tools []ToolDefinition) ([]codexTool, error) { + out := make([]codexTool, 0, len(tools)) + for _, t := range tools { + name := strings.TrimSpace(t.Function.Name) + if name == "" { + continue + } + params, err := schemaToRawJSON(t.Function.Parameters) + if err != nil { + return nil, fmt.Errorf("codex tool schema %s: %w", name, err) + } + out = append(out, codexTool{ + Type: "function", + Name: name, + Description: t.Function.Description, + Parameters: params, + }) + } + return out, nil +} + +func toCodexInput(messages []Message) (string, []codexInputItem) { + systemPrompt := "" + input := make([]codexInputItem, 0, len(messages)) + + for i, m := range messages { + role := strings.ToLower(strings.TrimSpace(m.Role)) + switch role { + case "system": + systemPrompt = m.Content + case "user": + input = append(input, codexInputItem{ + Role: "user", + Content: []codexInputContent{ + {Type: "input_text", Text: m.Content}, + }, + }) + case "assistant": + if strings.TrimSpace(m.Content) != "" { + input = append(input, codexInputItem{ + Type: "message", + Role: "assistant", + Status: "completed", + ID: fmt.Sprintf("msg_%d", i), + Content: []codexInputContent{ + {Type: "output_text", Text: m.Content}, + }, + }) + } + for _, tc := range m.ToolCalls { + callID, itemID := splitCodexToolCallID(tc.ID) + if strings.TrimSpace(callID) == "" { + callID = fmt.Sprintf("call_%d", i) + } + if strings.TrimSpace(itemID) == "" { + itemID = fmt.Sprintf("fc_%d", i) + } + args := strings.TrimSpace(tc.Function.Arguments) + if args == "" { + args = "{}" + } + input = append(input, codexInputItem{ + Type: "function_call", + ID: itemID, + CallID: callID, + Name: tc.Function.Name, + Arguments: args, + }) + } + case "tool": + callID, _ := splitCodexToolCallID(m.ToolCallID) + input = append(input, codexInputItem{ + Type: "function_call_output", + CallID: callID, + Output: m.Content, + }) + } + } + + return systemPrompt, input +} + +func splitCodexToolCallID(id string) (callID string, itemID string) { + v := strings.TrimSpace(id) + if v == "" { + return "call_0", "" + } + if strings.Contains(v, "|") { + parts := strings.SplitN(v, "|", 2) + return parts[0], parts[1] + } + return v, "" +} + +func stripCodexModelPrefix(model string) string { + m := strings.TrimSpace(model) + if after, ok := strings.CutPrefix(m, "openai-codex/"); ok { + return after + } + return m +} + +func resolveCodexModel(model string) string { + m := strings.ToLower(stripCodexModelPrefix(model)) + if m == "" { + return defaultCodexModel + } + if strings.Contains(m, "/") { + return defaultCodexModel + } + if strings.HasPrefix(m, "gpt-") || strings.HasPrefix(m, "o3") || strings.HasPrefix(m, "o4") { + return m + } + return defaultCodexModel +} + +func codexResponsesEndpoint(baseURL string) string { + base := strings.TrimRight(strings.TrimSpace(baseURL), "/") + if base == "" { + base = defaultCodexBaseURL + } + if strings.HasSuffix(base, "/codex") { + return base + "/responses" + } + if strings.HasSuffix(base, "/codex/responses") { + return base + } + return base + "/codex/responses" +} + +func codexPromptCacheKey(messages []Message) string { + b, err := json.Marshal(messages) + if err != nil { + return "" + } + sum := sha256.Sum256(b) + return hex.EncodeToString(sum[:]) +} + +func rawToCodexArgString(v json.RawMessage) string { + if len(v) == 0 { + return "" + } + var s string + if err := json.Unmarshal(v, &s); err == nil { + return s + } + return string(v) +} + +func codexArgumentsToJSON(raw string) json.RawMessage { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return json.RawMessage(`{}`) + } + if json.Valid([]byte(trimmed)) { + return json.RawMessage(trimmed) + } + b, _ := json.Marshal(map[string]string{"raw": trimmed}) + return b +} + +func codexFriendlyError(statusCode int, raw string) string { + if statusCode == http.StatusTooManyRequests { + return "usage quota exceeded or rate limited; try again later" + } + if raw == "" { + return http.StatusText(statusCode) + } + return raw +} + +func schemaToRawJSON(s JSONSchema) (json.RawMessage, error) { + b, err := json.Marshal(s) + if err != nil { + return nil, err + } + trimmed := bytes.TrimSpace(b) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return json.RawMessage(`{"type":"object"}`), nil + } + return json.RawMessage(trimmed), nil +} diff --git a/llm/openai_codex_oauth.go b/llm/openai_codex_oauth.go new file mode 100644 index 0000000..d6794f8 --- /dev/null +++ b/llm/openai_codex_oauth.go @@ -0,0 +1,750 @@ +package llm + +import ( + "bufio" + "context" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "os" + "os/exec" + "path/filepath" + "runtime" + "strconv" + "strings" + "time" + + "github.com/mosaxiv/clawlet/paths" +) + +const ( + codexOAuthClientID = "app_EMoamEEZ73f0CkXaXp7hrann" + codexOAuthIssuer = "https://auth.openai.com" + codexOAuthAuthorize = "https://auth.openai.com/oauth/authorize" + codexOAuthTokenURL = "https://auth.openai.com/oauth/token" + codexOAuthRedirectURI = "http://localhost:1455/auth/callback" + codexOAuthScope = "openid profile email offline_access" + codexOAuthOriginator = "codex_cli_rs" + codexJWTClaimPath = "https://api.openai.com/auth" + codexTokenFileName = "codex.json" + codexMinTTLSeconds = int64(60) +) + +const codexOAuthSuccessHTML = "Authentication successful

Authentication successful. Return to your terminal to continue.

" + +type CodexOAuthToken struct { + AccessToken string + AccountID string +} + +func (t CodexOAuthToken) Valid() bool { + return strings.TrimSpace(t.AccessToken) != "" && strings.TrimSpace(t.AccountID) != "" +} + +type codexStoredToken struct { + Access string `json:"access"` + Refresh string `json:"refresh"` + Expires int64 `json:"expires"` + AccountID string `json:"account_id,omitempty"` +} + +type codexDeviceCodeResponse struct { + DeviceAuthID string + UserCode string + IntervalSec int + ExpiresInSec int +} + +var errCodexDeviceAuthPending = errors.New("device authorization pending") + +func LoadCodexOAuthToken() (CodexOAuthToken, error) { + tok, err := getCodexToken(codexMinTTLSeconds) + if err != nil { + return CodexOAuthToken{}, err + } + out := CodexOAuthToken{AccessToken: tok.Access, AccountID: tok.AccountID} + if !out.Valid() { + return CodexOAuthToken{}, fmt.Errorf("codex oauth token is invalid; run `clawlet provider login openai-codex`") + } + return out, nil +} + +func LoginCodexOAuthInteractive(ctx context.Context) error { + verifier, challenge, err := generatePKCE() + if err != nil { + return err + } + state, err := createState() + if err != nil { + return err + } + + authURL := buildCodexAuthorizeURL(state, challenge) + fmt.Println("Open the following URL in your browser if it does not open automatically:") + fmt.Println(authURL) + _ = openBrowser(authURL) + + codeCh := make(chan string, 1) + server, serverErr := startCodexLocalServer(state, codeCh) + if serverErr != nil { + fmt.Printf("warning: local callback server could not start (%v)\n", serverErr) + } + + if server != nil { + defer server.Close() + fmt.Println("Waiting for browser callback...") + } + + code := "" + waitCtx, cancel := context.WithTimeout(ctx, 120*time.Second) + defer cancel() + if server != nil { + select { + case code = <-codeCh: + case <-waitCtx.Done(): + } + } + + if strings.TrimSpace(code) == "" { + fmt.Print("Paste the callback URL or authorization code: ") + line, err := bufio.NewReader(os.Stdin).ReadString('\n') + if err != nil && !errors.Is(err, io.EOF) { + return fmt.Errorf("read authorization input: %w", err) + } + parsedCode, parsedState := parseAuthorizationInput(line) + if parsedState != "" && parsedState != state { + return fmt.Errorf("oauth state validation failed") + } + code = parsedCode + } + if strings.TrimSpace(code) == "" { + return fmt.Errorf("authorization code not found") + } + + fmt.Println("Exchanging authorization code for tokens...") + tok, err := exchangeAuthorizationCode(ctx, code, verifier, codexOAuthRedirectURI) + if err != nil { + return err + } + if err := saveStoredCodexToken(tok); err != nil { + return err + } + return nil +} + +func LoginCodexOAuthDeviceCode(ctx context.Context) error { + device, err := requestCodexDeviceCode(ctx) + if err != nil { + return err + } + + fmt.Printf("\nTo authenticate, open this URL in your browser:\n\n %s/codex/device\n\nThen enter this code: %s\n\nWaiting for authentication...\n", + codexOAuthIssuer, device.UserCode) + + tok, err := pollCodexDeviceCode(ctx, device) + if err != nil { + return err + } + if err := saveStoredCodexToken(tok); err != nil { + return err + } + return nil +} + +func getCodexToken(minTTLSeconds int64) (codexStoredToken, error) { + tok, err := loadStoredCodexToken() + if err != nil { + return codexStoredToken{}, err + } + nowMs := time.Now().UnixMilli() + if tok.Expires-nowMs > minTTLSeconds*1000 { + return tok, nil + } + + refreshed, err := refreshCodexToken(tok.Refresh) + if err != nil { + latest, loadErr := loadStoredCodexToken() + if loadErr == nil && latest.Expires-time.Now().UnixMilli() > 0 { + return latest, nil + } + return codexStoredToken{}, err + } + if strings.TrimSpace(refreshed.AccountID) == "" { + refreshed.AccountID = tok.AccountID + } + if err := saveStoredCodexToken(refreshed); err != nil { + return codexStoredToken{}, err + } + return refreshed, nil +} + +func exchangeAuthorizationCode(ctx context.Context, code, verifier, redirectURI string) (codexStoredToken, error) { + form := url.Values{} + form.Set("grant_type", "authorization_code") + form.Set("client_id", codexOAuthClientID) + form.Set("code", strings.TrimSpace(code)) + form.Set("code_verifier", verifier) + form.Set("redirect_uri", redirectURI) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, codexOAuthTokenURL, strings.NewReader(form.Encode())) + if err != nil { + return codexStoredToken{}, err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) + if err != nil { + return codexStoredToken{}, err + } + defer resp.Body.Close() + + body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + if resp.StatusCode != http.StatusOK { + return codexStoredToken{}, fmt.Errorf("token exchange failed: %d %s", resp.StatusCode, strings.TrimSpace(string(body))) + } + return parseTokenPayload(body, "token exchange response missing fields", true) +} + +func requestCodexDeviceCode(ctx context.Context) (codexDeviceCodeResponse, error) { + reqBody, err := json.Marshal(map[string]string{ + "client_id": codexOAuthClientID, + }) + if err != nil { + return codexDeviceCodeResponse{}, err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, codexOAuthIssuer+"/api/accounts/deviceauth/usercode", strings.NewReader(string(reqBody))) + if err != nil { + return codexDeviceCodeResponse{}, err + } + req.Header.Set("Content-Type", "application/json") + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) + if err != nil { + return codexDeviceCodeResponse{}, err + } + defer resp.Body.Close() + body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + if resp.StatusCode != http.StatusOK { + return codexDeviceCodeResponse{}, fmt.Errorf("device code request failed: %s", strings.TrimSpace(string(body))) + } + return parseDeviceCodeResponse(body) +} + +func parseDeviceCodeResponse(body []byte) (codexDeviceCodeResponse, error) { + var raw struct { + DeviceAuthID string `json:"device_auth_id"` + UserCode string `json:"user_code"` + Interval json.RawMessage `json:"interval"` + ExpiresIn json.RawMessage `json:"expires_in"` + } + if err := json.Unmarshal(body, &raw); err != nil { + return codexDeviceCodeResponse{}, err + } + intervalSec, err := parseFlexibleInt(raw.Interval) + if err != nil { + return codexDeviceCodeResponse{}, err + } + if intervalSec < 1 { + intervalSec = 5 + } + expiresInSec, err := parseFlexibleInt(raw.ExpiresIn) + if err != nil { + return codexDeviceCodeResponse{}, err + } + // Fallback to a practical timeout when server doesn't return expires_in. + if expiresInSec < 60 { + expiresInSec = 30 * 60 + } + if strings.TrimSpace(raw.DeviceAuthID) == "" || strings.TrimSpace(raw.UserCode) == "" { + return codexDeviceCodeResponse{}, fmt.Errorf("device code response missing fields") + } + return codexDeviceCodeResponse{ + DeviceAuthID: raw.DeviceAuthID, + UserCode: raw.UserCode, + IntervalSec: intervalSec, + ExpiresInSec: expiresInSec, + }, nil +} + +func parseFlexibleInt(raw json.RawMessage) (int, error) { + if len(raw) == 0 || string(raw) == "null" { + return 0, nil + } + var v int + if err := json.Unmarshal(raw, &v); err == nil { + return v, nil + } + var s string + if err := json.Unmarshal(raw, &s); err == nil { + s = strings.TrimSpace(s) + if s == "" { + return 0, nil + } + return strconv.Atoi(s) + } + return 0, fmt.Errorf("invalid integer value: %s", string(raw)) +} + +func pollCodexDeviceCode(ctx context.Context, device codexDeviceCodeResponse) (codexStoredToken, error) { + deadline := time.NewTimer(time.Duration(device.ExpiresInSec) * time.Second) + defer deadline.Stop() + ticker := time.NewTicker(time.Duration(device.IntervalSec) * time.Second) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return codexStoredToken{}, ctx.Err() + case <-deadline.C: + return codexStoredToken{}, fmt.Errorf("device code authentication timed out") + case <-ticker.C: + tok, done, err := tryPollCodexDeviceCode(ctx, device.DeviceAuthID, device.UserCode) + if err != nil { + if errors.Is(err, errCodexDeviceAuthPending) { + continue + } + return codexStoredToken{}, err + } + if done { + return tok, nil + } + } + } +} + +func tryPollCodexDeviceCode(ctx context.Context, deviceAuthID, userCode string) (codexStoredToken, bool, error) { + reqBody, err := json.Marshal(map[string]string{ + "device_auth_id": deviceAuthID, + "user_code": userCode, + }) + if err != nil { + return codexStoredToken{}, false, err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, codexOAuthIssuer+"/api/accounts/deviceauth/token", strings.NewReader(string(reqBody))) + if err != nil { + return codexStoredToken{}, false, err + } + req.Header.Set("Content-Type", "application/json") + + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) + if err != nil { + return codexStoredToken{}, false, err + } + defer resp.Body.Close() + + body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + if resp.StatusCode != http.StatusOK { + if codexDeviceAuthIsPending(body) { + return codexStoredToken{}, false, errCodexDeviceAuthPending + } + return codexStoredToken{}, false, fmt.Errorf("device auth token request failed: %d %s", resp.StatusCode, strings.TrimSpace(string(body))) + } + + var tokenResp struct { + AuthorizationCode string `json:"authorization_code"` + CodeVerifier string `json:"code_verifier"` + } + if err := json.Unmarshal(body, &tokenResp); err != nil { + return codexStoredToken{}, false, err + } + if strings.TrimSpace(tokenResp.AuthorizationCode) == "" || strings.TrimSpace(tokenResp.CodeVerifier) == "" { + return codexStoredToken{}, false, fmt.Errorf("device auth token response missing fields") + } + + tok, err := exchangeAuthorizationCode(ctx, tokenResp.AuthorizationCode, tokenResp.CodeVerifier, codexOAuthIssuer+"/deviceauth/callback") + if err != nil { + return codexStoredToken{}, false, err + } + return tok, true, nil +} + +func codexDeviceAuthIsPending(body []byte) bool { + raw := strings.ToLower(strings.TrimSpace(string(body))) + if raw == "" { + return true + } + if strings.Contains(raw, "pending") || + strings.Contains(raw, "authorization_pending") || + strings.Contains(raw, "slow_down") || + strings.Contains(raw, "deviceauth_authorization_unknown") || + strings.Contains(raw, "device authorization is unknown") { + return true + } + var payload struct { + ErrorCode string `json:"error_code"` + ErrorDescription string `json:"error_description"` + Message string `json:"message"` + ErrorRaw json.RawMessage `json:"error"` + } + if err := json.Unmarshal(body, &payload); err == nil { + var errorText string + var errorObj struct { + Message string `json:"message"` + Type string `json:"type"` + Code string `json:"code"` + } + _ = json.Unmarshal(payload.ErrorRaw, &errorText) + _ = json.Unmarshal(payload.ErrorRaw, &errorObj) + for _, v := range []string{ + errorText, + payload.ErrorCode, + payload.ErrorDescription, + payload.Message, + errorObj.Message, + errorObj.Type, + errorObj.Code, + } { + l := strings.ToLower(strings.TrimSpace(v)) + if strings.Contains(l, "pending") || + strings.Contains(l, "authorization_pending") || + strings.Contains(l, "slow_down") || + strings.Contains(l, "deviceauth_authorization_unknown") || + strings.Contains(l, "device authorization is unknown") { + return true + } + } + } + return false +} + +func refreshCodexToken(refreshToken string) (codexStoredToken, error) { + form := url.Values{} + form.Set("grant_type", "refresh_token") + form.Set("refresh_token", strings.TrimSpace(refreshToken)) + form.Set("client_id", codexOAuthClientID) + + req, err := http.NewRequest(http.MethodPost, codexOAuthTokenURL, strings.NewReader(form.Encode())) + if err != nil { + return codexStoredToken{}, err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) + if err != nil { + return codexStoredToken{}, err + } + defer resp.Body.Close() + + body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + if resp.StatusCode != http.StatusOK { + return codexStoredToken{}, fmt.Errorf("token refresh failed: %d %s", resp.StatusCode, strings.TrimSpace(string(body))) + } + tok, err := parseTokenPayload(body, "token refresh response missing fields", false) + if err != nil { + return codexStoredToken{}, err + } + if strings.TrimSpace(tok.Refresh) == "" { + tok.Refresh = strings.TrimSpace(refreshToken) + } + return tok, nil +} + +func parseTokenPayload(body []byte, missingErr string, requireRefreshToken bool) (codexStoredToken, error) { + var payload struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresIn int64 `json:"expires_in"` + IDToken string `json:"id_token"` + } + if err := json.Unmarshal(body, &payload); err != nil { + return codexStoredToken{}, err + } + if strings.TrimSpace(payload.AccessToken) == "" || payload.ExpiresIn <= 0 { + return codexStoredToken{}, errors.New(missingErr) + } + if requireRefreshToken && strings.TrimSpace(payload.RefreshToken) == "" { + return codexStoredToken{}, errors.New(missingErr) + } + accountID := decodeCodexAccountID(payload.IDToken) + if strings.TrimSpace(accountID) == "" { + accountID = decodeCodexAccountID(payload.AccessToken) + } + return codexStoredToken{ + Access: payload.AccessToken, + Refresh: payload.RefreshToken, + Expires: time.Now().UnixMilli() + payload.ExpiresIn*1000, + AccountID: accountID, + }, nil +} + +func decodeCodexAccountID(token string) string { + parts := strings.Split(token, ".") + if len(parts) != 3 { + return "" + } + payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return "" + } + var payload map[string]json.RawMessage + if err := json.Unmarshal(payloadBytes, &payload); err != nil { + return "" + } + + if accountID := rawJSONFieldString(payload["chatgpt_account_id"]); strings.TrimSpace(accountID) != "" { + return accountID + } + if accountID := rawJSONFieldString(payload["https://api.openai.com/auth.chatgpt_account_id"]); strings.TrimSpace(accountID) != "" { + return accountID + } + if authRaw, ok := payload[codexJWTClaimPath]; ok { + var auth struct { + ChatGPTAccountID string `json:"chatgpt_account_id"` + } + if err := json.Unmarshal(authRaw, &auth); err == nil && strings.TrimSpace(auth.ChatGPTAccountID) != "" { + return auth.ChatGPTAccountID + } + } + if orgsRaw, ok := payload["organizations"]; ok { + var orgs []struct { + ID string `json:"id"` + } + if err := json.Unmarshal(orgsRaw, &orgs); err == nil { + for _, org := range orgs { + if strings.TrimSpace(org.ID) != "" { + return org.ID + } + } + } + } + return "" +} + +func rawJSONFieldString(raw json.RawMessage) string { + if len(raw) == 0 { + return "" + } + var out string + if err := json.Unmarshal(raw, &out); err != nil { + return "" + } + return out +} + +func buildCodexAuthorizeURL(state, challenge string) string { + q := url.Values{} + q.Set("response_type", "code") + q.Set("client_id", codexOAuthClientID) + q.Set("redirect_uri", codexOAuthRedirectURI) + q.Set("scope", codexOAuthScope) + q.Set("code_challenge", challenge) + q.Set("code_challenge_method", "S256") + q.Set("state", state) + q.Set("id_token_add_organizations", "true") + q.Set("codex_cli_simplified_flow", "true") + q.Set("originator", codexOAuthOriginator) + return codexOAuthAuthorize + "?" + q.Encode() +} + +func generatePKCE() (verifier string, challenge string, err error) { + rnd := make([]byte, 32) + if _, err := rand.Read(rnd); err != nil { + return "", "", err + } + verifier = base64.RawURLEncoding.EncodeToString(rnd) + sum := sha256.Sum256([]byte(verifier)) + challenge = base64.RawURLEncoding.EncodeToString(sum[:]) + return verifier, challenge, nil +} + +func createState() (string, error) { + rnd := make([]byte, 16) + if _, err := rand.Read(rnd); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(rnd), nil +} + +func parseAuthorizationInput(raw string) (code string, state string) { + v := strings.TrimSpace(raw) + if v == "" { + return "", "" + } + if u, err := url.Parse(v); err == nil && u.RawQuery != "" { + q := u.Query() + if q.Get("code") != "" { + return q.Get("code"), q.Get("state") + } + } + if strings.Contains(v, "#") { + parts := strings.SplitN(v, "#", 2) + return strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]) + } + if strings.Contains(v, "code=") { + if q, err := url.ParseQuery(v); err == nil { + return q.Get("code"), q.Get("state") + } + } + return v, "" +} + +func openBrowser(u string) error { + var cmd *exec.Cmd + switch runtime.GOOS { + case "darwin": + cmd = exec.Command("open", u) + case "windows": + cmd = exec.Command("rundll32", "url.dll,FileProtocolHandler", u) + default: + cmd = exec.Command("xdg-open", u) + } + return cmd.Start() +} + +func startCodexLocalServer(expectedState string, codeCh chan<- string) (io.Closer, error) { + mux := http.NewServeMux() + mux.HandleFunc("/auth/callback", func(w http.ResponseWriter, r *http.Request) { + state := r.URL.Query().Get("state") + if state != expectedState { + http.Error(w, "State mismatch", http.StatusBadRequest) + return + } + code := r.URL.Query().Get("code") + if strings.TrimSpace(code) == "" { + http.Error(w, "Missing code", http.StatusBadRequest) + return + } + + select { + case codeCh <- code: + default: + } + + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.Header().Set("Connection", "close") + _, _ = w.Write([]byte(codexOAuthSuccessHTML)) + }) + + ln, err := net.Listen("tcp", "localhost:1455") + if err != nil { + return nil, err + } + srv := &http.Server{Handler: mux} + go func() { _ = srv.Serve(ln) }() + return closerFunc(func() error { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + return srv.Shutdown(ctx) + }), nil +} + +type closerFunc func() error + +func (f closerFunc) Close() error { return f() } + +func loadStoredCodexToken() (codexStoredToken, error) { + path, err := codexTokenPath() + if err != nil { + return codexStoredToken{}, err + } + tok, err := readStoredCodexToken(path) + if err == nil { + return tok, nil + } + + imported, importErr := importFromCodexCLI(path) + if importErr == nil { + return imported, nil + } + return codexStoredToken{}, fmt.Errorf("oauth credentials not found; run `clawlet provider login openai-codex`") +} + +func readStoredCodexToken(path string) (codexStoredToken, error) { + b, err := os.ReadFile(path) + if err != nil { + return codexStoredToken{}, err + } + var tok codexStoredToken + if err := json.Unmarshal(b, &tok); err != nil { + return codexStoredToken{}, err + } + if strings.TrimSpace(tok.Access) == "" || strings.TrimSpace(tok.Refresh) == "" || tok.Expires <= 0 { + return codexStoredToken{}, fmt.Errorf("invalid token file") + } + return tok, nil +} + +func importFromCodexCLI(destPath string) (codexStoredToken, error) { + codexHome := strings.TrimSpace(os.Getenv("CODEX_HOME")) + if codexHome == "" { + codexHome = filepath.Join(userHomeDir(), ".codex") + } + codexPath := filepath.Join(codexHome, "auth.json") + b, err := os.ReadFile(codexPath) + if err != nil { + return codexStoredToken{}, err + } + var parsed struct { + Tokens struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + AccountID string `json:"account_id"` + } `json:"tokens"` + } + if err := json.Unmarshal(b, &parsed); err != nil { + return codexStoredToken{}, err + } + if parsed.Tokens.AccessToken == "" || parsed.Tokens.RefreshToken == "" || parsed.Tokens.AccountID == "" { + return codexStoredToken{}, fmt.Errorf("invalid codex auth format") + } + expires := time.Now().UnixMilli() + int64(time.Hour/time.Millisecond) + if st, err := os.Stat(codexPath); err == nil { + expires = st.ModTime().UnixMilli() + int64(time.Hour/time.Millisecond) + } + tok := codexStoredToken{ + Access: parsed.Tokens.AccessToken, + Refresh: parsed.Tokens.RefreshToken, + Expires: expires, + AccountID: parsed.Tokens.AccountID, + } + if err := writeStoredCodexToken(destPath, tok); err != nil { + return codexStoredToken{}, err + } + return tok, nil +} + +func saveStoredCodexToken(tok codexStoredToken) error { + path, err := codexTokenPath() + if err != nil { + return err + } + return writeStoredCodexToken(path, tok) +} + +func writeStoredCodexToken(path string, tok codexStoredToken) error { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return err + } + b, err := json.MarshalIndent(tok, "", " ") + if err != nil { + return err + } + b = append(b, '\n') + if err := os.WriteFile(path, b, 0o600); err != nil { + return err + } + _ = os.Chmod(path, 0o600) + return nil +} + +func codexTokenPath() (string, error) { + cfgDir, err := paths.ConfigDir() + if err != nil { + return "", err + } + return filepath.Join(cfgDir, "auth", codexTokenFileName), nil +} + +func userHomeDir() string { + home, err := os.UserHomeDir() + if err != nil { + return "" + } + return home +} diff --git a/llm/openai_codex_test.go b/llm/openai_codex_test.go new file mode 100644 index 0000000..5ecd9cb --- /dev/null +++ b/llm/openai_codex_test.go @@ -0,0 +1,235 @@ +package llm + +import ( + "encoding/base64" + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestCodexResponsesEndpoint(t *testing.T) { + if got := codexResponsesEndpoint(""); got != "https://chatgpt.com/backend-api/codex/responses" { + t.Fatalf("endpoint=%q", got) + } + if got := codexResponsesEndpoint("https://chatgpt.com/backend-api"); got != "https://chatgpt.com/backend-api/codex/responses" { + t.Fatalf("endpoint=%q", got) + } + if got := codexResponsesEndpoint("https://chatgpt.com/backend-api/codex/responses"); got != "https://chatgpt.com/backend-api/codex/responses" { + t.Fatalf("endpoint=%q", got) + } + if got := codexResponsesEndpoint("https://chatgpt.com/backend-api/codex"); got != "https://chatgpt.com/backend-api/codex/responses" { + t.Fatalf("endpoint=%q", got) + } +} + +func TestResolveCodexModel(t *testing.T) { + if got := resolveCodexModel(""); got != defaultCodexModel { + t.Fatalf("model=%q", got) + } + if got := resolveCodexModel("openai-codex/gpt-5.1-codex"); got != "gpt-5.1-codex" { + t.Fatalf("model=%q", got) + } + if got := resolveCodexModel("anthropic/claude-sonnet"); got != defaultCodexModel { + t.Fatalf("model=%q", got) + } +} + +func TestToCodexInput_ToolMapping(t *testing.T) { + msgs := []Message{ + {Role: "system", Content: "sys"}, + {Role: "user", Content: "hello"}, + { + Role: "assistant", + Content: "calling tool", + ToolCalls: []ToolCallPayload{ + { + ID: "call_1|fc_1", + Type: "function", + Function: ToolCallPayloadFunc{ + Name: "read_file", + Arguments: `{"path":"README.md"}`, + }, + }, + }, + }, + {Role: "tool", ToolCallID: "call_1|fc_1", Name: "read_file", Content: `{"ok":true}`}, + } + + system, input := toCodexInput(msgs) + if system != "sys" { + t.Fatalf("system=%q", system) + } + if len(input) != 4 { + t.Fatalf("input=%d", len(input)) + } + if input[2].Type != "function_call" { + t.Fatalf("type=%q", input[2].Type) + } + if input[2].CallID != "call_1" { + t.Fatalf("call_id=%q", input[2].CallID) + } +} + +func TestConsumeCodexSSE_ToolCall(t *testing.T) { + stream := strings.Join([]string{ + `data: {"type":"response.output_item.added","item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"read_file","arguments":""}}`, + "", + `data: {"type":"response.output_text.delta","delta":"Hello"}`, + "", + `data: {"type":"response.function_call_arguments.delta","call_id":"call_1","delta":"{\"path\":\"README.md\"}"}`, + "", + `data: {"type":"response.output_item.done","item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"read_file","arguments":"{\"path\":\"README.md\"}"}}`, + "", + `data: {"type":"response.completed","response":{"status":"completed"}}`, + "", + }, "\n") + + out, err := consumeCodexSSE(strings.NewReader(stream)) + if err != nil { + t.Fatalf("consume: %v", err) + } + if out.Content != "Hello" { + t.Fatalf("content=%q", out.Content) + } + if len(out.ToolCalls) != 1 { + t.Fatalf("tool_calls=%d", len(out.ToolCalls)) + } + if out.ToolCalls[0].ID != "call_1|fc_1" { + t.Fatalf("tool_id=%q", out.ToolCalls[0].ID) + } + var args map[string]string + if err := json.Unmarshal(out.ToolCalls[0].Arguments, &args); err != nil { + t.Fatalf("args json: %v", err) + } + if args["path"] != "README.md" { + t.Fatalf("path=%q", args["path"]) + } +} + +func TestParseAuthorizationInput(t *testing.T) { + code, state := parseAuthorizationInput("http://localhost:1455/auth/callback?code=abc&state=xyz") + if code != "abc" || state != "xyz" { + t.Fatalf("code=%q state=%q", code, state) + } + code, state = parseAuthorizationInput("abc#xyz") + if code != "abc" || state != "xyz" { + t.Fatalf("code=%q state=%q", code, state) + } + code, state = parseAuthorizationInput("just-code") + if code != "just-code" || state != "" { + t.Fatalf("code=%q state=%q", code, state) + } +} + +func TestLoadCodexOAuthToken_FromStoredToken(t *testing.T) { + dir := t.TempDir() + t.Setenv("HOME", dir) + path := filepath.Join(dir, ".clawlet", "auth", "codex.json") + + stored := codexStoredToken{ + Access: "access-token", + Refresh: "refresh-token", + Expires: time.Now().Add(10 * time.Minute).UnixMilli(), + AccountID: "acct_123", + } + b, err := json.Marshal(stored) + if err != nil { + t.Fatal(err) + } + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, b, 0o600); err != nil { + t.Fatal(err) + } + + tok, err := LoadCodexOAuthToken() + if err != nil { + t.Fatalf("load: %v", err) + } + if tok.AccessToken != "access-token" { + t.Fatalf("access=%q", tok.AccessToken) + } + if tok.AccountID != "acct_123" { + t.Fatalf("account=%q", tok.AccountID) + } +} + +func TestDecodeCodexAccountID_FromNestedClaim(t *testing.T) { + payload := map[string]interface{}{ + "https://api.openai.com/auth": map[string]interface{}{ + "chatgpt_account_id": "acct_nested", + }, + } + token := "x." + base64.RawURLEncoding.EncodeToString(mustJSON(payload)) + ".y" + if got := decodeCodexAccountID(token); got != "acct_nested" { + t.Fatalf("account_id=%q", got) + } +} + +func TestParseDeviceCodeResponse(t *testing.T) { + body := []byte(`{"device_auth_id":"dev-1","user_code":"ABC-DEF","interval":"7","expires_in":"1800"}`) + parsed, err := parseDeviceCodeResponse(body) + if err != nil { + t.Fatalf("parse: %v", err) + } + if parsed.DeviceAuthID != "dev-1" { + t.Fatalf("device_auth_id=%q", parsed.DeviceAuthID) + } + if parsed.UserCode != "ABC-DEF" { + t.Fatalf("user_code=%q", parsed.UserCode) + } + if parsed.IntervalSec != 7 { + t.Fatalf("interval=%d", parsed.IntervalSec) + } + if parsed.ExpiresInSec != 1800 { + t.Fatalf("expires_in=%d", parsed.ExpiresInSec) + } +} + +func TestParseTokenPayload_RefreshTokenOptionalOnRefreshFlow(t *testing.T) { + body := []byte(`{"access_token":"acc","expires_in":3600}`) + tok, err := parseTokenPayload(body, "missing", false) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tok.Access != "acc" { + t.Fatalf("access=%q", tok.Access) + } + if tok.Refresh != "" { + t.Fatalf("refresh=%q", tok.Refresh) + } +} + +func TestParseTokenPayload_RequiresRefreshTokenOnAuthCodeFlow(t *testing.T) { + body := []byte(`{"access_token":"acc","expires_in":3600}`) + if _, err := parseTokenPayload(body, "missing", true); err == nil { + t.Fatal("expected error, got nil") + } +} + +func TestCodexDeviceAuthIsPending(t *testing.T) { + if !codexDeviceAuthIsPending([]byte(`{"error":"authorization_pending"}`)) { + t.Fatal("expected pending=true") + } + if !codexDeviceAuthIsPending([]byte(`{"error":{"message":"Device authorization is unknown. Please try again.","type":"invalid_request_error","code":"deviceauth_authorization_unknown"}}`)) { + t.Fatal("expected pending=true for nested deviceauth_authorization_unknown") + } + if !codexDeviceAuthIsPending([]byte(`{"error":{"message":"Device authorization is unknown. Please try again."}}`)) { + t.Fatal("expected pending=true for unknown message only") + } + if codexDeviceAuthIsPending([]byte(`{"error":"access_denied"}`)) { + t.Fatal("expected pending=false") + } +} + +func mustJSON(v interface{}) []byte { + b, err := json.Marshal(v) + if err != nil { + panic(err) + } + return b +}