diff --git a/agent/loop.go b/agent/loop.go index 0a64480..20d56b7 100644 --- a/agent/loop.go +++ b/agent/loop.go @@ -78,6 +78,7 @@ func Run(ctx context.Context, cfg Config, state State) error { var content strings.Builder var toolCalls []backend.ToolCall var usage backend.Usage + var stopReason string var done bool for ev := range ch { @@ -88,6 +89,7 @@ func Run(ctx context.Context, cfg Config, state State) error { toolCalls = append(toolCalls, ev.ToolCalls...) if ev.Done { usage = ev.Usage + stopReason = ev.StopReason done = true break } @@ -96,6 +98,12 @@ func Run(ctx context.Context, cfg Config, state State) error { if !done { return fmt.Errorf("step %d: stream ended without done event", step) } + switch stopReason { + case "stop", "tool_calls", "length", "": + // normal + default: + return fmt.Errorf("step %d: %s", step, stopReason) + } totalToolCalls += len(toolCalls) // Announce and execute tool calls. diff --git a/backend/anthropic.go b/backend/anthropic.go new file mode 100644 index 0000000..e4ac77e --- /dev/null +++ b/backend/anthropic.go @@ -0,0 +1,337 @@ +package backend + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" +) + +const anthropicDefaultMaxTokens = 8192 + +// AnthropicBackend speaks the Anthropic Messages API. +type AnthropicBackend struct { + apiKey string + client *http.Client +} + +func NewAnthropic(apiKey string) *AnthropicBackend { + return &AnthropicBackend{apiKey: apiKey, client: &http.Client{}} +} + +// -- wire types -- + +type anthropicRequest struct { + Model string `json:"model"` + MaxTokens int `json:"max_tokens"` + System string `json:"system,omitempty"` + Messages []anthropicMessage `json:"messages"` + Tools []anthropicTool `json:"tools,omitempty"` + Stream bool `json:"stream"` + Temperature *float64 `json:"temperature,omitempty"` +} + +type anthropicMessage struct { + Role string `json:"role"` + Content []anthropicContentBlock `json:"content"` +} + +type anthropicContentBlock struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Input json.RawMessage `json:"input,omitempty"` + ToolUseID string `json:"tool_use_id,omitempty"` + Content string `json:"content,omitempty"` // tool_result text +} + +type anthropicTool struct { + Name string `json:"name"` + Description string `json:"description,omitempty"` + InputSchema json.RawMessage `json:"input_schema"` +} + +// -- implementation -- + +func (b *AnthropicBackend) ChatStream(ctx context.Context, model string, messages []Message, tools []Tool, params GenerationParams) (<-chan StreamEvent, error) { + system, wireMessages := buildAnthropicMessages(messages) + + maxTokens := params.MaxTokens + if maxTokens == 0 { + maxTokens = anthropicDefaultMaxTokens + } + + areq := anthropicRequest{ + Model: model, + MaxTokens: maxTokens, + System: system, + Messages: wireMessages, + Stream: true, + Temperature: params.Temperature, + } + for _, t := range tools { + schema := t.Parameters + if schema == nil { + schema = json.RawMessage(`{"type":"object","properties":{}}`) + } + areq.Tools = append(areq.Tools, anthropicTool{ + Name: t.Name, + Description: t.Description, + InputSchema: schema, + }) + } + + data, err := json.Marshal(areq) + if err != nil { + return nil, err + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, "https://api.anthropic.com/v1/messages", bytes.NewReader(data)) + if err != nil { + return nil, err + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("X-Api-Key", b.apiKey) + httpReq.Header.Set("Anthropic-Version", "2023-06-01") + + resp, err := b.client.Do(httpReq) + if err != nil { + return nil, err + } + if resp.StatusCode == http.StatusTooManyRequests { + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return nil, &RateLimitError{RetryAfter: parseRetryAfter(resp.Header.Get("Retry-After")), Message: string(body)} + } + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return nil, fmt.Errorf("anthropic HTTP %d: %s", resp.StatusCode, body) + } + + ch := make(chan StreamEvent, 8) + go func() { + defer close(ch) + defer resp.Body.Close() + streamAnthropicSSE(ctx, resp.Body, ch) + }() + return ch, nil +} + +// buildAnthropicMessages converts ollie messages to Anthropic wire format. +// System messages are extracted into the top-level system field. +// Consecutive tool messages are batched into a single user message with +// multiple tool_result blocks (Anthropic requires strictly alternating roles). +func buildAnthropicMessages(messages []Message) (system string, out []anthropicMessage) { + for i := 0; i < len(messages); { + m := messages[i] + switch m.Role { + case "system": + if system != "" { + system += "\n\n" + } + system += m.Content + i++ + case "user": + out = append(out, anthropicMessage{ + Role: "user", + Content: []anthropicContentBlock{{Type: "text", Text: m.Content}}, + }) + i++ + case "assistant": + msg := anthropicMessage{Role: "assistant"} + if m.Content != "" { + msg.Content = append(msg.Content, anthropicContentBlock{Type: "text", Text: m.Content}) + } + for _, tc := range m.ToolCalls { + input := tc.Arguments + if input == nil { + input = json.RawMessage("{}") + } + msg.Content = append(msg.Content, anthropicContentBlock{ + Type: "tool_use", + ID: tc.ID, + Name: tc.Name, + Input: input, + }) + } + out = append(out, msg) + i++ + case "tool": + // Batch all consecutive tool messages into one user message. + var blocks []anthropicContentBlock + for i < len(messages) && messages[i].Role == "tool" { + blocks = append(blocks, anthropicContentBlock{ + Type: "tool_result", + ToolUseID: messages[i].ToolCallID, + Content: messages[i].Content, + }) + i++ + } + out = append(out, anthropicMessage{Role: "user", Content: blocks}) + default: + i++ + } + } + return +} + +// streamAnthropicSSE reads the Anthropic SSE stream and sends StreamEvents to ch. +func streamAnthropicSSE(ctx context.Context, body io.Reader, ch chan<- StreamEvent) { + scanner := bufio.NewScanner(body) + scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) + + type toolAccum struct { + id, name string + args strings.Builder + } + tools := map[int]*toolAccum{} + + var inputTokens, outputTokens int + var stopReason string + var curEvent, curData string + done := false + + process := func(typ, data string) { + if done { + return + } + raw := []byte(data) + switch typ { + case "message_start": + var v struct { + Message struct { + Usage struct { + InputTokens int `json:"input_tokens"` + } `json:"usage"` + } `json:"message"` + } + json.Unmarshal(raw, &v) //nolint:errcheck + inputTokens = v.Message.Usage.InputTokens + + case "content_block_start": + var v struct { + Index int `json:"index"` + ContentBlock struct { + Type string `json:"type"` + ID string `json:"id"` + Name string `json:"name"` + } `json:"content_block"` + } + json.Unmarshal(raw, &v) //nolint:errcheck + if v.ContentBlock.Type == "tool_use" { + tools[v.Index] = &toolAccum{id: v.ContentBlock.ID, name: v.ContentBlock.Name} + } + + case "content_block_delta": + var v struct { + Index int `json:"index"` + Delta struct { + Type string `json:"type"` + Text string `json:"text"` + PartialJSON string `json:"partial_json"` + } `json:"delta"` + } + json.Unmarshal(raw, &v) //nolint:errcheck + switch v.Delta.Type { + case "text_delta": + if v.Delta.Text != "" { + select { + case ch <- StreamEvent{Content: v.Delta.Text}: + case <-ctx.Done(): + done = true + } + } + case "input_json_delta": + if t := tools[v.Index]; t != nil { + t.args.WriteString(v.Delta.PartialJSON) + } + } + + case "message_delta": + var v struct { + Delta struct { + StopReason string `json:"stop_reason"` + } `json:"delta"` + Usage struct { + OutputTokens int `json:"output_tokens"` + } `json:"usage"` + } + json.Unmarshal(raw, &v) //nolint:errcheck + outputTokens = v.Usage.OutputTokens + stopReason = v.Delta.StopReason + + case "message_stop": + ev := StreamEvent{ + Done: true, + StopReason: mapAnthropicStopReason(stopReason), + Usage: Usage{InputTokens: inputTokens, OutputTokens: outputTokens}, + } + for _, t := range tools { + ev.ToolCalls = append(ev.ToolCalls, ToolCall{ + ID: t.id, + Name: t.name, + Arguments: json.RawMessage(t.args.String()), + }) + } + ch <- ev + done = true + + case "error": + var v struct { + Error struct { + Message string `json:"message"` + } `json:"error"` + } + json.Unmarshal(raw, &v) //nolint:errcheck + ch <- StreamEvent{Done: true, StopReason: "error: " + v.Error.Message} + done = true + } + } + + for scanner.Scan() { + if done { + break + } + line := scanner.Text() + if line == "" { + if curEvent != "" && curData != "" { + process(curEvent, curData) + } + curEvent = "" + curData = "" + continue + } + if strings.HasPrefix(line, "event: ") { + curEvent = strings.TrimPrefix(line, "event: ") + } else if strings.HasPrefix(line, "data: ") { + curData = strings.TrimPrefix(line, "data: ") + } + } + // Handle any trailing event not terminated by a blank line. + if !done && curEvent != "" && curData != "" { + process(curEvent, curData) + } + + if err := scanner.Err(); err != nil && !done { + ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)} + } +} + +func mapAnthropicStopReason(reason string) string { + switch reason { + case "end_turn", "stop_sequence": + return "stop" + case "max_tokens": + return "length" + case "tool_use": + return "tool_calls" + default: + return reason + } +} diff --git a/backend/codewhisperer.go b/backend/codewhisperer.go new file mode 100644 index 0000000..a3a724e --- /dev/null +++ b/backend/codewhisperer.go @@ -0,0 +1,807 @@ +package backend + +// CodeWhispererBackend implements Backend for Amazon CodeWhisperer / Kiro. +// +// Auth is configured via the apiKey parameter to NewCodeWhisperer: +// - Empty string → read from Kiro CLI SQLite database at the default path +// ($XDG_DATA_HOME/kiro-cli/data.sqlite3 on Linux). +// - "sqlite:///path/to/data.sqlite3" → read from the specified SQLite file. +// - Any other string → used directly as a bearer token (static auth). +// +// SQLite auth requires the sqlite3 CLI to be present in PATH. + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/url" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "sync" + "time" +) + +// ── Public constructor ──────────────────────────────────────────────────────── + +type CodeWhispererBackend struct { + endpoint string // overrides auth-derived endpoint if non-empty + extraHeaders map[string]string + authSource kiroAuthSource + authInitErr error + httpClient *http.Client +} + +// NewCodeWhisperer returns a CodeWhisperer backend. See package doc for apiKey semantics. +func NewCodeWhisperer(apiKey string) (*CodeWhispererBackend, error) { + if strings.TrimSpace(apiKey) == "" { + apiKey = kiroDefaultSQLiteKey() + } + authSource, err := newKiroAuthSource(apiKey) + return &CodeWhispererBackend{ + authSource: authSource, + authInitErr: err, + httpClient: &http.Client{}, + }, nil +} + +// ── Backend interface ───────────────────────────────────────────────────────── + +func (b *CodeWhispererBackend) ChatStream( + ctx context.Context, + model string, + messages []Message, + tools []Tool, + params GenerationParams, +) (<-chan StreamEvent, error) { + if b.authInitErr != nil { + return nil, b.authInitErr + } + + profileARN, err := b.authSource.ProfileARN(ctx) + if err != nil { + return nil, fmt.Errorf("kiro profile ARN: %w", err) + } + + req, err := buildKiroRequest(model, messages, tools, profileARN) + if err != nil { + return nil, err + } + + ch := make(chan StreamEvent, 8) + go func() { + defer close(ch) + b.runStream(ctx, req, ch) + }() + return ch, nil +} + +func (b *CodeWhispererBackend) runStream(ctx context.Context, req *kiroGenerateRequest, ch chan<- StreamEvent) { + for attempt := range 2 { + token, err := b.authSource.AccessToken(ctx) + if err != nil { + ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("auth: %v", err)} + return + } + endpoint, err := b.resolveEndpoint(ctx) + if err != nil { + ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("endpoint: %v", err)} + return + } + + client := newKiroAPIClient(endpoint, token, b.extraHeaders, b.httpClient) + + stopped := false + result, streamErr := client.Stream(ctx, *req, kiroStreamCallbacks{ + OnAssistantDelta: func(event kiroAssistantEvent) { + if stopped || event.Content == "" { + return + } + select { + case ch <- StreamEvent{Content: event.Content}: + case <-ctx.Done(): + stopped = true + } + }, + }) + + if streamErr != nil { + if attempt == 0 && b.shouldRefresh(streamErr) { + if refreshErr := b.authSource.Refresh(ctx); refreshErr != nil { + ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("token refresh: %v", refreshErr)} + return + } + continue + } + ch <- StreamEvent{Done: true, StopReason: streamErr.Error()} + return + } + + final := StreamEvent{Done: true, StopReason: "stop"} + if len(result.ToolUses) > 0 { + final.StopReason = "tool_calls" + for _, tu := range result.ToolUses { + final.ToolCalls = append(final.ToolCalls, ToolCall{ + ID: tu.ToolUseID, + Name: tu.Name, + Arguments: tu.Input, + }) + } + } + ch <- final + return + } +} + +func (b *CodeWhispererBackend) resolveEndpoint(ctx context.Context) (string, error) { + if ep := strings.TrimSpace(b.endpoint); ep != "" { + return ep, nil + } + return b.authSource.DefaultEndpoint(ctx) +} + +func (b *CodeWhispererBackend) shouldRefresh(err error) bool { + if !b.authSource.CanRefresh() { + return false + } + var apiErr *kiroAPIError + return errors.As(err, &apiErr) && apiErr.isUnauthorized() +} + +// ── Request building ────────────────────────────────────────────────────────── + +// kiroSanitizeModel returns "auto" for any model string that isn't a native +// Kiro/CodeWhisperer model (e.g. Ollama "name:tag" or OpenRouter "org/name"). +func kiroSanitizeModel(model string) string { + if model == "" || strings.ContainsAny(model, "/:") { + return "auto" + } + return model +} + +func buildKiroRequest(model string, messages []Message, tools []Tool, profileARN string) (*kiroGenerateRequest, error) { + model = kiroSanitizeModel(model) + if len(messages) == 0 { + return nil, fmt.Errorf("at least one message is required") + } + + encodedTools := encodeKiroToolDefinitions(tools) + + lastRole := strings.ToLower(strings.TrimSpace(messages[len(messages)-1].Role)) + + var history []kiroChatMessage + var current kiroChatMessage + var err error + + switch lastRole { + case "user", "system": + history, err = encodeKiroHistory(messages[:len(messages)-1]) + if err != nil { + return nil, err + } + current = kiroChatMessage{UserInputMessage: encodeKiroUserMessage(messages[len(messages)-1], model, encodedTools, nil)} + + case "tool": + // Find the last non-tool message. + cutoff := len(messages) - 1 + for cutoff >= 0 && strings.EqualFold(messages[cutoff].Role, "tool") { + cutoff-- + } + history, err = encodeKiroHistory(messages[:cutoff+1]) + if err != nil { + return nil, err + } + toolResults := encodeKiroToolResults(messages[cutoff+1:]) + current = kiroChatMessage{UserInputMessage: emptyKiroUserMessage(model, encodedTools, toolResults)} + + default: + return nil, fmt.Errorf("last message role %q is not supported (expected user, system, or tool)", lastRole) + } + + convID, err := kiroNewUUID() + if err != nil { + return nil, err + } + contID, err := kiroNewUUID() + if err != nil { + return nil, err + } + + return &kiroGenerateRequest{ + ProfileARN: profileARN, + ConversationState: kiroConversationState{ + ConversationID: convID, + History: history, + CurrentMessage: current, + ChatTriggerType: kiroTriggerType, + AgentContinuationID: contID, + AgentTaskType: kiroAgentTaskType, + }, + }, nil +} + +func encodeKiroHistory(messages []Message) ([]kiroChatMessage, error) { + out := make([]kiroChatMessage, 0, len(messages)) + for i := 0; i < len(messages); i++ { + m := messages[i] + switch strings.ToLower(strings.TrimSpace(m.Role)) { + case "user", "system": + out = append(out, kiroChatMessage{UserInputMessage: encodeKiroUserMessage(m, "", nil, nil)}) + case "assistant": + out = append(out, kiroChatMessage{AssistantResponseMessage: encodeKiroAssistantMessage(m)}) + case "tool": + // Batch consecutive tool messages together. + start := i + for i+1 < len(messages) && strings.EqualFold(messages[i+1].Role, "tool") { + i++ + } + results := encodeKiroToolResults(messages[start : i+1]) + out = append(out, kiroChatMessage{UserInputMessage: emptyKiroUserMessage("", nil, results)}) + } + } + return out, nil +} + +func encodeKiroUserMessage(m Message, modelID string, tools []kiroTool, toolResults []kiroToolResult) *kiroUserInputMessage { + msg := &kiroUserInputMessage{ + Content: m.Content, + Origin: kiroDefaultOrigin, + ModelID: modelID, + } + ctx := &kiroUserInputContext{} + if len(tools) > 0 { + ctx.Tools = tools + } + if len(toolResults) > 0 { + ctx.ToolResults = toolResults + } + if len(ctx.Tools) > 0 || len(ctx.ToolResults) > 0 { + msg.UserInputMessageContext = ctx + } + return msg +} + +func emptyKiroUserMessage(modelID string, tools []kiroTool, toolResults []kiroToolResult) *kiroUserInputMessage { + msg := &kiroUserInputMessage{ + Content: "", + Origin: kiroDefaultOrigin, + ModelID: modelID, + } + ctx := &kiroUserInputContext{} + if len(tools) > 0 { + ctx.Tools = tools + } + if len(toolResults) > 0 { + ctx.ToolResults = toolResults + } + if len(ctx.Tools) > 0 || len(ctx.ToolResults) > 0 { + msg.UserInputMessageContext = ctx + } + return msg +} + +func encodeKiroAssistantMessage(m Message) *kiroAssistantResponseMessage { + msg := &kiroAssistantResponseMessage{Content: m.Content} + for _, tc := range m.ToolCalls { + input := tc.Arguments + if len(input) == 0 || !json.Valid(input) { + input = json.RawMessage("{}") + } + toolUseID := strings.TrimSpace(tc.ID) + if toolUseID == "" { + toolUseID, _ = kiroNewUUID() + } + msg.ToolUses = append(msg.ToolUses, kiroToolUse{ + ToolUseID: toolUseID, + Name: tc.Name, + Input: input, + }) + } + return msg +} + +func encodeKiroToolDefinitions(tools []Tool) []kiroTool { + if len(tools) == 0 { + return nil + } + out := make([]kiroTool, 0, len(tools)) + for _, t := range tools { + schema := t.Parameters + if schema == nil { + schema = json.RawMessage("{}") + } + out = append(out, kiroTool{ + ToolSpecification: &kiroToolSpecification{ + Name: t.Name, + Description: t.Description, + InputSchema: kiroInputSchema{JSON: schema}, + }, + }) + } + return out +} + +func encodeKiroToolResults(messages []Message) []kiroToolResult { + results := make([]kiroToolResult, 0, len(messages)) + for _, m := range messages { + if !strings.EqualFold(m.Role, "tool") { + continue + } + toolUseID := strings.TrimSpace(m.ToolCallID) + results = append(results, kiroToolResult{ + ToolUseID: toolUseID, + Content: []kiroToolResultContent{encodeKiroToolResultContent(m.Content)}, + }) + } + return results +} + +func encodeKiroToolResultContent(content string) kiroToolResultContent { + content = strings.TrimSpace(content) + if content != "" && json.Valid([]byte(content)) { + return kiroToolResultContent{JSON: json.RawMessage(content)} + } + return kiroToolResultContent{Text: content} +} + +// ── Auth sources ────────────────────────────────────────────────────────────── + +type kiroAuthSource interface { + DefaultEndpoint(ctx context.Context) (string, error) + ProfileARN(ctx context.Context) (string, error) + AccessToken(ctx context.Context) (string, error) + Refresh(ctx context.Context) error + CanRefresh() bool +} + +// staticKiroAuth holds a pre-issued bearer token. +type staticKiroAuth struct { + token string +} + +func (s *staticKiroAuth) DefaultEndpoint(_ context.Context) (string, error) { + return "", fmt.Errorf("base URL not configured and no profile ARN available with static auth") +} + +func (s *staticKiroAuth) ProfileARN(_ context.Context) (string, error) { return "", nil } + +func (s *staticKiroAuth) AccessToken(_ context.Context) (string, error) { + t := strings.TrimSpace(s.token) + if t == "" { + return "", fmt.Errorf("token is empty") + } + return t, nil +} + +func (s *staticKiroAuth) Refresh(_ context.Context) error { + return fmt.Errorf("static tokens cannot be refreshed") +} + +func (s *staticKiroAuth) CanRefresh() bool { return false } + +// sqliteKiroAuth reads tokens from the Kiro CLI SQLite database and refreshes via OIDC. +type sqliteKiroAuth struct { + path string + state kiroSQLiteState + loaded bool + initMu sync.Mutex + mu sync.Mutex +} + +type kiroSQLiteState struct { + ProfileARN string // from OIDC profile state; empty for social/personal accounts + Region string + AccessToken string + RefreshToken string + ExpiresAt string + ClientID string + ClientSecret string + IsSocial bool // true for personal (GitHub/Google) accounts +} + +func (s *sqliteKiroAuth) DefaultEndpoint(ctx context.Context) (string, error) { + state, err := s.load(ctx) + if err != nil { + return "", err + } + if arn := strings.TrimSpace(state.ProfileARN); arn != "" { + return kiroDefaultEndpoint(arn) + } + if region := strings.TrimSpace(state.Region); region != "" { + return kiroEndpointForRegion(region), nil + } + return "", fmt.Errorf("cannot derive endpoint: no profile ARN or region in Kiro auth state") +} + +func (s *sqliteKiroAuth) ProfileARN(ctx context.Context) (string, error) { + state, err := s.load(ctx) + if err != nil { + return "", err + } + // Personal/social accounts don't send profileArn in requests. + if state.IsSocial { + return "", nil + } + return strings.TrimSpace(state.ProfileARN), nil +} + +func (s *sqliteKiroAuth) AccessToken(ctx context.Context) (string, error) { + state, err := s.load(ctx) + if err != nil { + return "", err + } + if state.AccessToken != "" && !kiroTokenExpiringSoon(state.ExpiresAt) { + return state.AccessToken, nil + } + if !s.canRefreshState(state) { + if state.AccessToken != "" { + return state.AccessToken, nil // expired but can't refresh — use it anyway + } + return "", fmt.Errorf("sqlite auth state %q: no access token and cannot refresh", s.path) + } + if err := s.Refresh(ctx); err != nil { + if state.AccessToken != "" { + return state.AccessToken, nil // refresh failed — use stale token + } + return "", err + } + state, err = s.load(ctx) + if err != nil { + return "", err + } + return state.AccessToken, nil +} + +func (s *sqliteKiroAuth) Refresh(ctx context.Context) error { + s.mu.Lock() + defer s.mu.Unlock() + + state, err := s.load(ctx) + if err != nil { + return err + } + if !s.canRefreshState(state) { + return fmt.Errorf("sqlite auth state %q: insufficient data for token refresh", s.path) + } + + // Re-read current token — another process may have already refreshed it. + currentAccess, currentRefresh, currentExpiry, err := kiroReadSQLiteTokens(ctx, s.path) + if err == nil && currentAccess != "" && currentAccess != state.AccessToken && !kiroTokenExpiringSoon(currentExpiry) { + return nil // already refreshed by another process + } + + oidcEndpoint := fmt.Sprintf("https://oidc.%s.amazonaws.com/token", strings.TrimSpace(state.Region)) + client := newKiroOIDCClient(oidcEndpoint, nil) + + refreshToken := currentRefresh + if refreshToken == "" { + refreshToken = state.RefreshToken + } + resp, err := client.RefreshToken(ctx, kiroRefreshTokenRequest{ + ClientID: state.ClientID, + ClientSecret: state.ClientSecret, + GrantType: "refresh_token", + RefreshToken: refreshToken, + }) + if err != nil { + return err + } + + newAccess := strings.TrimSpace(resp.AccessToken) + newRefresh := strings.TrimSpace(resp.RefreshToken) + var expiresAt string + if resp.ExpiresIn != nil && *resp.ExpiresIn > 0 { + expiresAt = time.Now().Add(time.Duration(*resp.ExpiresIn) * time.Second).UTC().Format(time.RFC3339) + } + return kiroUpdateSQLiteTokens(ctx, s.path, newAccess, newRefresh, expiresAt) +} + +func (s *sqliteKiroAuth) CanRefresh() bool { + state, err := s.load(context.Background()) + return err == nil && s.canRefreshState(state) +} + +func (s *sqliteKiroAuth) canRefreshState(state kiroSQLiteState) bool { + return strings.TrimSpace(state.RefreshToken) != "" && + strings.TrimSpace(state.Region) != "" && + strings.TrimSpace(state.ClientID) != "" && + strings.TrimSpace(state.ClientSecret) != "" +} + +func (s *sqliteKiroAuth) load(ctx context.Context) (kiroSQLiteState, error) { + s.initMu.Lock() + if !s.loaded { + state, err := kiroLoadSQLiteState(ctx, s.path) + if err != nil { + s.initMu.Unlock() + return kiroSQLiteState{}, err + } + s.state = state + s.loaded = true + } + s.initMu.Unlock() + + // Always re-read tokens — they may have been refreshed by another process. + access, refresh, expiresAt, err := kiroReadSQLiteTokens(ctx, s.path) + if err != nil { + return s.state, nil // fall back to cached state + } + state := s.state + state.AccessToken = access + state.ExpiresAt = expiresAt + if refresh != "" { + state.RefreshToken = refresh + } + return state, nil +} + +const kiroTokenExpiryBuffer = 2 * time.Minute + +func kiroTokenExpiringSoon(expiresAt string) bool { + if expiresAt == "" { + return false + } + // Social auth uses nanosecond precision; OIDC auth uses second precision. + t, err := time.Parse(time.RFC3339Nano, expiresAt) + if err != nil { + t, err = time.Parse(time.RFC3339, expiresAt) + if err != nil { + return false + } + } + return time.Until(t) < kiroTokenExpiryBuffer +} + +// ── SQLite helpers ──────────────────────────────────────────────────────────── + +// Kiro supports two auth flows whose tokens live under different SQLite keys: +// - OIDC (org/enterprise accounts): kirocli:odic:token + kirocli:odic:device-registration +// - Social (personal accounts, e.g. GitHub login): kirocli:social:token +// +// We probe OIDC first; if that key is absent we fall back to social. + +const ( + kiroSQLiteProfileQuery = "SELECT value FROM state WHERE key = 'api.codewhisperer.profile' LIMIT 1;" + kiroSQLiteOIDCTokenQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:odic:token' LIMIT 1;" + kiroSQLiteDeviceRegQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:odic:device-registration' LIMIT 1;" + kiroSQLiteSocialTokenQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:social:token' LIMIT 1;" +) + +type kiroSQLiteProfileState struct { + ARN string `json:"arn"` +} + +// OIDC auth token (org/enterprise accounts) +type kiroSQLiteOIDCTokenState struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresAt string `json:"expires_at"` + Region string `json:"region"` +} + +type kiroSQLiteDeviceRegState struct { + ClientID string `json:"client_id"` + ClientSecret string `json:"client_secret"` + Region string `json:"region"` +} + +// Social auth token (personal accounts — GitHub, etc.) +type kiroSQLiteSocialTokenState struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresAt string `json:"expires_at"` + ProfileARN string `json:"profile_arn"` +} + +func kiroLoadSQLiteState(ctx context.Context, path string) (kiroSQLiteState, error) { + // Try OIDC token first. + if oidcVal, ok, err := kiroQuerySQLiteOptional(ctx, path, kiroSQLiteOIDCTokenQuery); err != nil { + return kiroSQLiteState{}, err + } else if ok { + return kiroLoadOIDCState(ctx, path, oidcVal) + } + + // Fall back to social token (personal accounts). + socialVal, err := kiroQuerySQLite(ctx, path, kiroSQLiteSocialTokenQuery) + if err != nil { + return kiroSQLiteState{}, fmt.Errorf("no OIDC or social token found in %q", path) + } + return kiroLoadSocialState(socialVal) +} + +func kiroLoadOIDCState(ctx context.Context, path, tokenVal string) (kiroSQLiteState, error) { + var token kiroSQLiteOIDCTokenState + if err := json.Unmarshal([]byte(tokenVal), &token); err != nil { + return kiroSQLiteState{}, fmt.Errorf("decode OIDC token from sqlite %q: %w", path, err) + } + + var profile kiroSQLiteProfileState + if v, ok, err := kiroQuerySQLiteOptional(ctx, path, kiroSQLiteProfileQuery); err != nil { + return kiroSQLiteState{}, err + } else if ok { + _ = json.Unmarshal([]byte(v), &profile) + } + + var dev kiroSQLiteDeviceRegState + if devVal, ok, err := kiroQuerySQLiteOptional(ctx, path, kiroSQLiteDeviceRegQuery); err != nil { + return kiroSQLiteState{}, err + } else if ok { + _ = json.Unmarshal([]byte(devVal), &dev) + } + + region := strings.TrimSpace(token.Region) + if region == "" { + region = strings.TrimSpace(dev.Region) + } + + return kiroSQLiteState{ + ProfileARN: strings.TrimSpace(profile.ARN), + Region: region, + AccessToken: strings.TrimSpace(token.AccessToken), + RefreshToken: strings.TrimSpace(token.RefreshToken), + ExpiresAt: strings.TrimSpace(token.ExpiresAt), + ClientID: strings.TrimSpace(dev.ClientID), + ClientSecret: strings.TrimSpace(dev.ClientSecret), + }, nil +} + +func kiroLoadSocialState(tokenVal string) (kiroSQLiteState, error) { + var token kiroSQLiteSocialTokenState + if err := json.Unmarshal([]byte(tokenVal), &token); err != nil { + return kiroSQLiteState{}, fmt.Errorf("decode social token: %w", err) + } + return kiroSQLiteState{ + ProfileARN: strings.TrimSpace(token.ProfileARN), // used only for endpoint derivation + AccessToken: strings.TrimSpace(token.AccessToken), + RefreshToken: strings.TrimSpace(token.RefreshToken), + ExpiresAt: strings.TrimSpace(token.ExpiresAt), + IsSocial: true, + // No Region, ClientID, ClientSecret — social auth can't use OIDC refresh. + }, nil +} + +// kiroReadSQLiteTokens re-reads the current tokens from SQLite (for cross-process +// refresh detection). Tries OIDC key first, then social. +func kiroReadSQLiteTokens(ctx context.Context, path string) (access, refresh, expiresAt string, err error) { + // Try OIDC token. + if out, err2 := kiroRunSQLite(ctx, "-readonly", "-batch", "-noheader", path, kiroSQLiteOIDCTokenQuery); err2 == nil { + if value := strings.TrimSpace(string(out)); value != "" { + var state kiroSQLiteOIDCTokenState + if err2 := json.Unmarshal([]byte(value), &state); err2 == nil { + return strings.TrimSpace(state.AccessToken), strings.TrimSpace(state.RefreshToken), strings.TrimSpace(state.ExpiresAt), nil + } + } + } + + // Fall back to social token. + out, err2 := kiroRunSQLite(ctx, "-readonly", "-batch", "-noheader", path, kiroSQLiteSocialTokenQuery) + if err2 != nil { + return "", "", "", fmt.Errorf("read token from sqlite %q: %w", path, err2) + } + value := strings.TrimSpace(string(out)) + if value == "" { + return "", "", "", fmt.Errorf("read token from sqlite %q: no rows", path) + } + var state kiroSQLiteSocialTokenState + if err2 := json.Unmarshal([]byte(value), &state); err2 != nil { + return "", "", "", fmt.Errorf("decode social token from sqlite %q: %w", path, err2) + } + return strings.TrimSpace(state.AccessToken), strings.TrimSpace(state.RefreshToken), strings.TrimSpace(state.ExpiresAt), nil +} + +func kiroUpdateSQLiteTokens(ctx context.Context, path, access, refresh, expiresAt string) error { + access = strings.TrimSpace(access) + if access == "" { + return fmt.Errorf("access token must not be empty") + } + var sb strings.Builder + sb.WriteString("BEGIN IMMEDIATE;\n") + sb.WriteString("UPDATE auth_kv SET value = json_set(value, '$.access_token', ") + sb.WriteString(kiroSQLiteQuote(access)) + if r := strings.TrimSpace(refresh); r != "" { + sb.WriteString(", '$.refresh_token', ") + sb.WriteString(kiroSQLiteQuote(r)) + } + if e := strings.TrimSpace(expiresAt); e != "" { + sb.WriteString(", '$.expires_at', ") + sb.WriteString(kiroSQLiteQuote(e)) + } + sb.WriteString(") WHERE key = 'kirocli:odic:token';\nCOMMIT;\n") + _, err := kiroRunSQLite(ctx, "-batch", path, sb.String()) + return err +} + +func kiroQuerySQLite(ctx context.Context, path, query string) (string, error) { + out, err := kiroRunSQLite(ctx, "-readonly", "-batch", "-noheader", path, query) + if err != nil { + return "", fmt.Errorf("sqlite query on %q: %w", path, err) + } + value := strings.TrimSpace(string(out)) + if value == "" { + return "", fmt.Errorf("sqlite query on %q returned no rows", path) + } + return value, nil +} + +func kiroQuerySQLiteOptional(ctx context.Context, path, query string) (string, bool, error) { + v, err := kiroQuerySQLite(ctx, path, query) + if err == nil { + return v, true, nil + } + if strings.Contains(err.Error(), "returned no rows") { + return "", false, nil + } + return "", false, err +} + +func kiroRunSQLite(ctx context.Context, args ...string) ([]byte, error) { + cmd := exec.CommandContext(ctx, "sqlite3", args...) + out, err := cmd.CombinedOutput() + if err != nil { + msg := strings.TrimSpace(string(out)) + if msg != "" { + return nil, fmt.Errorf("%s", msg) + } + return nil, err + } + return out, nil +} + +func kiroSQLiteQuote(s string) string { + return "'" + strings.ReplaceAll(s, "'", "''") + "'" +} + +// ── Auth source factory ─────────────────────────────────────────────────────── + +func newKiroAuthSource(apiKey string) (kiroAuthSource, error) { + apiKey = strings.TrimSpace(apiKey) + if !strings.HasPrefix(strings.ToLower(apiKey), "sqlite://") { + return &staticKiroAuth{token: apiKey}, nil + } + // Parse sqlite:// URI + parsed, err := url.Parse(apiKey) + if err != nil { + return nil, fmt.Errorf("parse sqlite URL %q: %w", apiKey, err) + } + path := parsed.Path + if parsed.Host != "" && !strings.HasPrefix(apiKey, "sqlite:///") { + path = "/" + parsed.Host + parsed.Path + } + path, err = url.PathUnescape(path) + if err != nil { + return nil, fmt.Errorf("decode sqlite path: %w", err) + } + if strings.TrimSpace(path) == "" { + return nil, fmt.Errorf("sqlite URL %q does not contain a database path", apiKey) + } + return &sqliteKiroAuth{path: path}, nil +} + +func kiroDefaultSQLiteKey() string { + return "sqlite://" + filepath.Join(kiroUserDataDir(), "kiro-cli", "data.sqlite3") +} + +func kiroUserDataDir() string { + switch runtime.GOOS { + case "windows": + if d := os.Getenv("LocalAppData"); d != "" { + return d + } + case "darwin": + if h, _ := os.UserHomeDir(); h != "" { + return filepath.Join(h, "Library", "Application Support") + } + default: + if d := os.Getenv("XDG_DATA_HOME"); d != "" { + return d + } + if h, _ := os.UserHomeDir(); h != "" { + return filepath.Join(h, ".local", "share") + } + } + return os.TempDir() +} diff --git a/backend/codewhisperer_internal.go b/backend/codewhisperer_internal.go new file mode 100644 index 0000000..af6b540 --- /dev/null +++ b/backend/codewhisperer_internal.go @@ -0,0 +1,610 @@ +package backend + +// Low-level plumbing for the Amazon CodeWhisperer / Kiro backend: +// - Wire types for the GenerateAssistantResponse API +// - Binary AWS event stream decoder +// - HTTP API client +// - SQLite-backed token storage helpers +// - OIDC token refresh client + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/base64" + "encoding/binary" + "encoding/json" + "fmt" + "hash/crc32" + "io" + "maps" + "net/http" + "strings" +) + +// ── Wire types ──────────────────────────────────────────────────────────────── + +const ( + kiroGenerateTarget = "AmazonCodeWhispererStreamingService.GenerateAssistantResponse" + kiroDefaultOrigin = "KIRO_CLI" + kiroTriggerType = "MANUAL" + kiroAgentTaskType = "vibe" + kiroAmzSDKRequest = "attempt=1; max=3" + + kiroUserAgent = "aws-sdk-rust/1.3.14 ua/2.1 api/codewhispererstreaming/0.1.14474 os/linux lang/rust/1.92.0 md/appVersion-1.27.2 app/AmazonQ-For-CLI" + kiroAmzUserAgent = "aws-sdk-rust/1.3.14 ua/2.1 api/codewhispererstreaming/0.1.14474 os/linux lang/rust/1.92.0 m/F app/AmazonQ-For-CLI" + + kiroOIDCUserAgent = "aws-sdk-rust/1.3.10 os/linux lang/rust/1.92.0" + kiroOIDCAmzUserAgent = "aws-sdk-rust/1.3.10 ua/2.1 api/ssooidc/1.92.0 os/linux lang/rust/1.92.0 m/E app/AmazonQ-For-CLI" +) + +type kiroRefreshTokenRequest struct { + ClientID string `json:"clientId"` + ClientSecret string `json:"clientSecret"` + GrantType string `json:"grantType"` + RefreshToken string `json:"refreshToken"` +} + +type kiroRefreshTokenResponse struct { + AccessToken string `json:"accessToken"` + RefreshToken string `json:"refreshToken,omitempty"` + ExpiresIn *int `json:"expiresIn,omitempty"` +} + +type kiroGenerateRequest struct { + ConversationState kiroConversationState `json:"conversationState"` + ProfileARN string `json:"profileArn,omitempty"` +} + +type kiroConversationState struct { + ConversationID string `json:"conversationId,omitempty"` + History []kiroChatMessage `json:"history,omitempty"` + CurrentMessage kiroChatMessage `json:"currentMessage"` + ChatTriggerType string `json:"chatTriggerType,omitempty"` + AgentContinuationID string `json:"agentContinuationId,omitempty"` + AgentTaskType string `json:"agentTaskType,omitempty"` +} + +type kiroChatMessage struct { + UserInputMessage *kiroUserInputMessage `json:"userInputMessage,omitempty"` + AssistantResponseMessage *kiroAssistantResponseMessage `json:"assistantResponseMessage,omitempty"` +} + +type kiroUserInputMessage struct { + Content string `json:"content"` + UserInputMessageContext *kiroUserInputContext `json:"userInputMessageContext,omitempty"` + Origin string `json:"origin,omitempty"` + ModelID string `json:"modelId,omitempty"` +} + +type kiroUserInputContext struct { + ToolResults []kiroToolResult `json:"toolResults,omitempty"` + Tools []kiroTool `json:"tools,omitempty"` +} + +type kiroAssistantResponseMessage struct { + MessageID string `json:"messageId,omitempty"` + Content string `json:"content"` + ToolUses []kiroToolUse `json:"toolUses,omitempty"` +} + +type kiroTool struct { + ToolSpecification *kiroToolSpecification `json:"toolSpecification,omitempty"` +} + +type kiroToolSpecification struct { + InputSchema kiroInputSchema `json:"inputSchema"` + Name string `json:"name"` + Description string `json:"description,omitempty"` +} + +type kiroInputSchema struct { + JSON json.RawMessage `json:"json"` +} + +type kiroToolUse struct { + ToolUseID string `json:"toolUseId"` + Name string `json:"name"` + Input json.RawMessage `json:"input"` +} + +type kiroToolResult struct { + ToolUseID string `json:"toolUseId"` + Content []kiroToolResultContent `json:"content"` + Status string `json:"status,omitempty"` +} + +type kiroToolResultContent struct { + Text string `json:"text,omitempty"` + JSON json.RawMessage `json:"json,omitempty"` +} + +// stream event types +type kiroMetadataEvent struct { + ConversationID string `json:"conversationId,omitempty"` +} + +type kiroAssistantEvent struct { + Content string `json:"content"` +} + +type kiroToolUseEvent struct { + ToolUseID string `json:"toolUseId"` + Name string `json:"name"` + Input string `json:"input,omitempty"` + Stop bool `json:"stop,omitempty"` +} + +type kiroStreamResult struct { + ToolUses []kiroToolUse +} + +// ── Binary AWS event stream decoder ────────────────────────────────────────── + +type kiroEventFrame struct { + Headers map[string]any + MessageType string + EventType string + Payload []byte +} + +func kiroReadEventFrame(r io.Reader) (*kiroEventFrame, error) { + prelude := make([]byte, 12) + if _, err := io.ReadFull(r, prelude); err != nil { + if err == io.EOF || err == io.ErrUnexpectedEOF { + return nil, io.EOF + } + return nil, fmt.Errorf("read eventstream prelude: %w", err) + } + + totalLength := binary.BigEndian.Uint32(prelude[0:4]) + headersLength := binary.BigEndian.Uint32(prelude[4:8]) + preludeCRC := binary.BigEndian.Uint32(prelude[8:12]) + + if totalLength < 16 { + return nil, fmt.Errorf("invalid eventstream frame length %d", totalLength) + } + if crc32.ChecksumIEEE(prelude[:8]) != preludeCRC { + return nil, fmt.Errorf("eventstream prelude crc mismatch") + } + + remaining := make([]byte, totalLength-12) + if _, err := io.ReadFull(r, remaining); err != nil { + return nil, fmt.Errorf("read eventstream remainder: %w", err) + } + + message := append(prelude, remaining...) + expectedCRC := binary.BigEndian.Uint32(message[len(message)-4:]) + if crc32.ChecksumIEEE(message[:len(message)-4]) != expectedCRC { + return nil, fmt.Errorf("eventstream message crc mismatch") + } + + headersStart := 12 + headersEnd := headersStart + int(headersLength) + if headersEnd > len(message)-4 { + return nil, fmt.Errorf("invalid eventstream headers length %d", headersLength) + } + + headers, err := kiroDecodeHeaders(message[headersStart:headersEnd]) + if err != nil { + return nil, err + } + + frame := &kiroEventFrame{ + Headers: headers, + Payload: append([]byte(nil), message[headersEnd:len(message)-4]...), + } + if v, ok := kiroHeaderString(headers, ":message-type"); ok { + frame.MessageType = v + } + if v, ok := kiroHeaderString(headers, ":event-type"); ok { + frame.EventType = v + } else if v, ok := kiroHeaderString(headers, ":exception-type"); ok { + frame.EventType = v + } + return frame, nil +} + +func kiroDecodeHeaders(data []byte) (map[string]any, error) { + headers := make(map[string]any) + for offset := 0; offset < len(data); { + nameLen := int(data[offset]) + offset++ + if offset+nameLen > len(data) { + return nil, fmt.Errorf("eventstream header name exceeds buffer") + } + name := string(data[offset : offset+nameLen]) + offset += nameLen + if offset >= len(data) { + return nil, fmt.Errorf("eventstream header missing type") + } + val, consumed, err := kiroDecodeHeaderValue(data[offset:]) + if err != nil { + return nil, err + } + offset += consumed + headers[name] = val + } + return headers, nil +} + +func kiroDecodeHeaderValue(data []byte) (any, int, error) { + if len(data) == 0 { + return nil, 0, fmt.Errorf("empty eventstream header value") + } + switch data[0] { + case 0: + return true, 1, nil + case 1: + return false, 1, nil + case 2: + if len(data) < 2 { + return nil, 0, fmt.Errorf("truncated int8 eventstream header") + } + return int8(data[1]), 2, nil + case 3: + if len(data) < 3 { + return nil, 0, fmt.Errorf("truncated int16 eventstream header") + } + return int16(binary.BigEndian.Uint16(data[1:3])), 3, nil + case 4: + if len(data) < 5 { + return nil, 0, fmt.Errorf("truncated int32 eventstream header") + } + return int32(binary.BigEndian.Uint32(data[1:5])), 5, nil + case 5, 8: + if len(data) < 9 { + return nil, 0, fmt.Errorf("truncated int64 eventstream header") + } + return int64(binary.BigEndian.Uint64(data[1:9])), 9, nil + case 6: + if len(data) < 3 { + return nil, 0, fmt.Errorf("truncated bytes eventstream header") + } + size := int(binary.BigEndian.Uint16(data[1:3])) + if len(data) < 3+size { + return nil, 0, fmt.Errorf("truncated bytes eventstream header payload") + } + return base64.StdEncoding.EncodeToString(data[3 : 3+size]), 3 + size, nil + case 7: + if len(data) < 3 { + return nil, 0, fmt.Errorf("truncated string eventstream header") + } + size := int(binary.BigEndian.Uint16(data[1:3])) + if len(data) < 3+size { + return nil, 0, fmt.Errorf("truncated string eventstream header payload") + } + return string(data[3 : 3+size]), 3 + size, nil + case 9: + if len(data) < 17 { + return nil, 0, fmt.Errorf("truncated uuid eventstream header") + } + return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", + data[1:5], data[5:7], data[7:9], data[9:11], data[11:17]), 17, nil + default: + return nil, 0, fmt.Errorf("unsupported eventstream header type %d", data[0]) + } +} + +func kiroHeaderString(headers map[string]any, name string) (string, bool) { + v, ok := headers[name] + if !ok { + return "", false + } + s, ok := v.(string) + return s, ok +} + +// ── API client ──────────────────────────────────────────────────────────────── + +type kiroAPIClient struct { + endpoint string + token string + httpClient *http.Client + extraHeaders map[string]string +} + +type kiroAPIError struct { + StatusCode int + Code string + Message string +} + +func (e *kiroAPIError) Error() string { + parts := make([]string, 0, 2) + if e.Code != "" { + parts = append(parts, e.Code) + } + if e.Message != "" { + parts = append(parts, e.Message) + } + if len(parts) == 0 { + return fmt.Sprintf("HTTP %d", e.StatusCode) + } + return strings.Join(parts, ": ") +} + +func (e *kiroAPIError) isUnauthorized() bool { + return e.StatusCode == http.StatusUnauthorized || e.StatusCode == http.StatusForbidden +} + +type kiroStreamCallbacks struct { + OnAssistantDelta func(kiroAssistantEvent) +} + +func newKiroAPIClient(endpoint, token string, extraHeaders map[string]string, httpClient *http.Client) *kiroAPIClient { + if httpClient == nil { + httpClient = &http.Client{} + } + return &kiroAPIClient{ + endpoint: strings.TrimRight(endpoint, "/") + "/", + token: token, + httpClient: httpClient, + extraHeaders: cloneKiroHeaders(extraHeaders), + } +} + +func (c *kiroAPIClient) Stream( + ctx context.Context, + req kiroGenerateRequest, + callbacks kiroStreamCallbacks, +) (*kiroStreamResult, error) { + body, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, c.endpoint, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("build request: %w", err) + } + if err := c.applyHeaders(httpReq); err != nil { + return nil, err + } + + resp, err := c.httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("send request: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, kiroDecodeHTTPError(resp) + } + + result := &kiroStreamResult{} + accumulators := make(map[string]*kiroToolUseAccumulator) + order := make([]string, 0) + + for { + frame, err := kiroReadEventFrame(resp.Body) + if err == io.EOF { + break + } + if err != nil { + return nil, err + } + + if frame.MessageType == "exception" { + return nil, kiroDecodeStreamException(frame) + } + if frame.MessageType != "" && frame.MessageType != "event" { + continue + } + + switch frame.EventType { + case "assistantResponseEvent": + var event kiroAssistantEvent + if err := json.Unmarshal(frame.Payload, &event); err != nil { + return nil, fmt.Errorf("decode assistant response event: %w", err) + } + if callbacks.OnAssistantDelta != nil { + callbacks.OnAssistantDelta(event) + } + case "toolUseEvent": + var event kiroToolUseEvent + if err := json.Unmarshal(frame.Payload, &event); err != nil { + return nil, fmt.Errorf("decode tool use event: %w", err) + } + if _, ok := accumulators[event.ToolUseID]; !ok { + accumulators[event.ToolUseID] = &kiroToolUseAccumulator{ + ToolUseID: event.ToolUseID, + Name: event.Name, + } + order = append(order, event.ToolUseID) + } + acc := accumulators[event.ToolUseID] + if event.Name != "" { + acc.Name = event.Name + } + acc.Input.WriteString(event.Input) + } + } + + toolUses := make([]kiroToolUse, 0, len(order)) + for _, id := range order { + acc := accumulators[id] + input := strings.TrimSpace(acc.Input.String()) + if input == "" || !json.Valid([]byte(input)) { + input = "{}" + } + toolUses = append(toolUses, kiroToolUse{ + ToolUseID: acc.ToolUseID, + Name: acc.Name, + Input: json.RawMessage(input), + }) + } + result.ToolUses = toolUses + return result, nil +} + +type kiroToolUseAccumulator struct { + ToolUseID string + Name string + Input strings.Builder +} + +func (c *kiroAPIClient) applyHeaders(req *http.Request) error { + invocationID, err := kiroNewUUID() + if err != nil { + return fmt.Errorf("generate invocation id: %w", err) + } + req.Header.Set("Content-Type", "application/x-amz-json-1.0") + req.Header.Set("X-Amz-Target", kiroGenerateTarget) + req.Header.Set("Authorization", "Bearer "+c.token) + req.Header.Set("Accept", "*/*") + req.Header.Set("Accept-Encoding", "gzip") + req.Header.Set("User-Agent", kiroUserAgent) + req.Header.Set("X-Amz-User-Agent", kiroAmzUserAgent) + req.Header.Set("Amz-Sdk-Request", kiroAmzSDKRequest) + req.Header.Set("Amz-Sdk-Invocation-Id", invocationID) + req.Header.Set("X-Amzn-Codewhisperer-Optout", "false") + for k, v := range c.extraHeaders { + req.Header.Set(k, v) + } + return nil +} + +func kiroDecodeHTTPError(resp *http.Response) error { + body, _ := io.ReadAll(resp.Body) + var envelope map[string]any + _ = json.Unmarshal(body, &envelope) + + e := &kiroAPIError{StatusCode: resp.StatusCode} + if code := resp.Header.Get("X-Amzn-Errortype"); code != "" { + e.Code = kiroCleanErrorCode(code) + } + if e.Code == "" { + if v, ok := envelope["code"].(string); ok { + e.Code = kiroCleanErrorCode(v) + } + if e.Code == "" { + if v, ok := envelope["__type"].(string); ok { + e.Code = kiroCleanErrorCode(v) + } + } + } + for _, key := range []string{"message", "Message", "errorMessage"} { + if v, ok := envelope[key].(string); ok && v != "" { + e.Message = v + break + } + } + return e +} + +func kiroDecodeStreamException(frame *kiroEventFrame) error { + var envelope map[string]any + _ = json.Unmarshal(frame.Payload, &envelope) + e := &kiroAPIError{Code: kiroCleanErrorCode(frame.EventType)} + for _, key := range []string{"message", "Message", "errorMessage"} { + if v, ok := envelope[key].(string); ok && v != "" { + e.Message = v + break + } + } + return e +} + +func kiroCleanErrorCode(value string) string { + value = strings.TrimSpace(value) + if value == "" { + return "" + } + value = strings.Split(value, ":")[0] + if strings.Contains(value, "#") { + parts := strings.Split(value, "#") + value = parts[len(parts)-1] + } + return value +} + +func kiroDefaultEndpoint(profileARN string) (string, error) { + parts := strings.Split(profileARN, ":") + if len(parts) < 6 || parts[0] != "arn" { + return "", fmt.Errorf("invalid profile ARN %q", profileARN) + } + region := parts[3] + if region == "" { + return "", fmt.Errorf("profile ARN %q does not contain a region", profileARN) + } + return "https://q." + region + ".amazonaws.com/", nil +} + +func kiroEndpointForRegion(region string) string { + return "https://q." + strings.TrimSpace(region) + ".amazonaws.com/" +} + +func kiroNewUUID() (string, error) { + var b [16]byte + if _, err := rand.Read(b[:]); err != nil { + return "", fmt.Errorf("generate uuid: %w", err) + } + b[6] = (b[6] & 0x0f) | 0x40 + b[8] = (b[8] & 0x3f) | 0x80 + return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", + b[0:4], b[4:6], b[6:8], b[8:10], b[10:16]), nil +} + +func cloneKiroHeaders(m map[string]string) map[string]string { + if len(m) == 0 { + return nil + } + out := make(map[string]string, len(m)) + maps.Copy(out, m) + return out +} + +// ── OIDC token refresh client ───────────────────────────────────────────────── + +type kiroOIDCClient struct { + endpoint string + httpClient *http.Client +} + +func newKiroOIDCClient(endpoint string, httpClient *http.Client) *kiroOIDCClient { + if httpClient == nil { + httpClient = &http.Client{} + } + return &kiroOIDCClient{endpoint: endpoint, httpClient: httpClient} +} + +func (c *kiroOIDCClient) RefreshToken(ctx context.Context, req kiroRefreshTokenRequest) (*kiroRefreshTokenResponse, error) { + body, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal refresh request: %w", err) + } + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, c.endpoint, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("build refresh request: %w", err) + } + invocationID, err := kiroNewUUID() + if err != nil { + return nil, err + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Accept", "*/*") + httpReq.Header.Set("Accept-Encoding", "gzip") + httpReq.Header.Set("User-Agent", kiroOIDCUserAgent) + httpReq.Header.Set("X-Amz-User-Agent", kiroOIDCAmzUserAgent) + httpReq.Header.Set("Amz-Sdk-Request", kiroAmzSDKRequest) + httpReq.Header.Set("Amz-Sdk-Invocation-Id", invocationID) + + resp, err := c.httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("send refresh request: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, kiroDecodeHTTPError(resp) + } + var result kiroRefreshTokenResponse + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return nil, fmt.Errorf("decode refresh response: %w", err) + } + if strings.TrimSpace(result.AccessToken) == "" { + return nil, fmt.Errorf("refresh response did not contain accessToken") + } + return &result, nil +} diff --git a/backend/copilot.go b/backend/copilot.go new file mode 100644 index 0000000..5f33676 --- /dev/null +++ b/backend/copilot.go @@ -0,0 +1,20 @@ +package backend + +import "net/http" + +// NewCopilot returns an OpenAI-compatible backend configured for the GitHub +// Copilot chat completions API. token is a short-lived Copilot bearer token. +// +// The required Copilot-specific headers (Openai-Intent, User-Agent) are +// injected on every request via the OpenAIBackend's extraHeaders mechanism. +func NewCopilot(token string) *OpenAIBackend { + return &OpenAIBackend{ + baseURL: "https://api.githubcopilot.com", + apiKey: token, + client: &http.Client{}, + extraHeaders: map[string]string{ + "Openai-Intent": "conversation-edits", + "User-Agent": "opencode/local", + }, + } +} diff --git a/backend/new.go b/backend/new.go index c9d0746..5370a59 100644 --- a/backend/new.go +++ b/backend/new.go @@ -40,10 +40,14 @@ func loadEnvFile(path string) { // so the file acts as a default; variables already set in the environment take // precedence. // -// OLLIE_BACKEND ollama | openai | openrouter (default: ollama) -// OLLIE_OLLAMA_URL base URL for ollama (default: http://localhost:11434) -// OLLIE_OPENAI_URL base URL for openai-compatible backends -// OLLIE_OPENAI_KEY API key (required for openai) +// OLLIE_BACKEND ollama | openai | openrouter | anthropic | copilot | kiro (default: ollama) +// OLLIE_OLLAMA_URL base URL for ollama (default: http://localhost:11434) +// OLLIE_OPENAI_URL base URL for openai-compatible backends +// OLLIE_OPENAI_KEY API key (required for openai/openrouter) +// OLLIE_ANTHROPIC_KEY API key (required for anthropic) +// OLLIE_COPILOT_TOKEN bearer token (required for copilot) +// OLLIE_KIRO_TOKEN bearer token or sqlite:// URL for kiro/codewhisperer +// (default: sqlite path auto-detected from Kiro CLI data dir) func New() (Backend, error) { home, _ := os.UserHomeDir() loadEnvFile(home + "/.config/ollie/env") @@ -56,9 +60,23 @@ func New() (Backend, error) { switch which { case "ollama": return NewOllama(os.Getenv("OLLIE_OLLAMA_URL")), nil - case "openai": + case "openai", "openrouter": return NewOpenAI(os.Getenv("OLLIE_OPENAI_URL"), os.Getenv("OLLIE_OPENAI_KEY")), nil + case "anthropic": + key := os.Getenv("OLLIE_ANTHROPIC_KEY") + if key == "" { + return nil, fmt.Errorf("OLLIE_ANTHROPIC_KEY is required for anthropic backend") + } + return NewAnthropic(key), nil + case "copilot": + token := os.Getenv("OLLIE_COPILOT_TOKEN") + if token == "" { + return nil, fmt.Errorf("OLLIE_COPILOT_TOKEN is required for copilot backend") + } + return NewCopilot(token), nil + case "kiro", "codewhisperer": + return NewCodeWhisperer(os.Getenv("OLLIE_KIRO_TOKEN")) default: - return nil, fmt.Errorf("unknown OLLIE_BACKEND %q (supported: ollama, openai)", which) + return nil, fmt.Errorf("unknown OLLIE_BACKEND %q (supported: ollama, openai, openrouter, anthropic, copilot)", which) } } diff --git a/backend/openai.go b/backend/openai.go index 5f36ab9..7a775c3 100644 --- a/backend/openai.go +++ b/backend/openai.go @@ -16,9 +16,10 @@ import ( // OpenAIBackend speaks the OpenAI /v1/chat/completions wire format. // Compatible with OpenRouter, OpenAI, and any other OpenAI-compatible API. type OpenAIBackend struct { - baseURL string - apiKey string - client *http.Client + baseURL string + apiKey string + client *http.Client + extraHeaders map[string]string // optional; applied after Authorization } func NewOpenAI(baseURL, apiKey string) *OpenAIBackend { @@ -161,6 +162,9 @@ func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages [ if b.apiKey != "" { httpReq.Header.Set("Authorization", "Bearer "+b.apiKey) } + for k, v := range b.extraHeaders { + httpReq.Header.Set(k, v) + } resp, err := b.client.Do(httpReq) if err != nil { diff --git a/interrupt_commands.go b/interrupt_commands.go new file mode 100644 index 0000000..59864f5 --- /dev/null +++ b/interrupt_commands.go @@ -0,0 +1,50 @@ +package main + +import ( + "fmt" + "strconv" + "strings" +) + +const interruptCommandUsage = "/interrupt " + +func parseInterruptCommandLine(input string) (prompt string, ok bool, err error) { + input = strings.TrimSpace(input) + if !strings.HasPrefix(input, "/") { + return "", false, nil + } + name, args, _ := strings.Cut(input[1:], " ") + if name != "interrupt" { + return "", false, nil + } + prompt, err = parseInterruptCommandArgs(args) + return prompt, true, err +} + +func parseInterruptCommandArgs(args string) (string, error) { + prompt := strings.TrimSpace(args) + if prompt == "" { + return "", fmt.Errorf("usage: %s", interruptCommandUsage) + } + if unquoted, err := strconv.Unquote(prompt); err == nil { + prompt = strings.TrimSpace(unquoted) + } else if len(prompt) >= 2 && prompt[0] == '\'' && prompt[len(prompt)-1] == '\'' { + prompt = strings.TrimSpace(prompt[1 : len(prompt)-1]) + } + if prompt == "" { + return "", fmt.Errorf("usage: %s", interruptCommandUsage) + } + return prompt, nil +} + +func interruptQueuedMessage(prompt string) string { + return "◇ interrupt queued: " + prompt +} + +func interruptExpiredMessage(count int) string { + return fmt.Sprintf("◇ %d interrupt(s) expired (turn ended)", count) +} + +func interruptQueueFullMessage() string { + return "Interrupt queue is full." +} diff --git a/main.go b/main.go index a6cea68..4b6ecd7 100644 --- a/main.go +++ b/main.go @@ -420,6 +420,20 @@ func newSessionID() string { return time.Now().Format("20060102-150405") + "-" + fmt.Sprintf("%06x", b) } +// defaultModelForBackend returns a sensible default model for the given backend label. +func defaultModelForBackend(name string) string { + switch name { + case "anthropic": + return "claude-sonnet-4-5" + case "openrouter": + return "deepseek/deepseek-v3.2" + case "kiro", "codewhisperer": + return "auto" + default: // ollama, local, groq, mistral, together, etc. + return "qwen3.5:9b" + } +} + // resolveBackendName returns a short human-readable backend label. func resolveBackendName() string { which := os.Getenv("OLLIE_BACKEND") @@ -767,7 +781,9 @@ func (s *appState) handleCommand(ctx context.Context, input string, out io.Write } s.loopcfg.Backend = be s.backendName = resolveBackendName() - fmt.Fprintf(out, "switched backend to: %s\n", s.backendName) + s.loopcfg.Model = defaultModelForBackend(s.backendName) + s.modelName = s.loopcfg.Model + fmt.Fprintf(out, "switched backend to: %s (model: %s)\n", s.backendName, s.modelName) return true case "/model": @@ -956,6 +972,7 @@ func (s *appState) handleCommand(ctx context.Context, input string, out io.Write func main() { sessionFlag := flag.String("session", "", "resume a session by ID") + promptFlag := flag.String("prompt", "", "run a single prompt non-interactively and exit") flag.Parse() extraArgs := flag.Args() @@ -973,12 +990,12 @@ func main() { os.Exit(1) } + backendName := resolveBackendName() + modelName := os.Getenv("OLLIE_MODEL") if modelName == "" { - modelName = "qwen3.5:9b" + modelName = defaultModelForBackend(backendName) } - - backendName := resolveBackendName() builtinExec := execpkg.New( home+"/.local/state/ollie", home+"/.cache/ollie/exec", @@ -1068,6 +1085,11 @@ func main() { exec.Command("sh", "-c", hook).Run() //nolint:errcheck } + if *promptFlag != "" { + s.processInput(context.Background(), *promptFlag, os.Stdout) + return + } + s.runInteractiveTTY(context.Background()) }