Add CodeWhisperer/Kiro backend + backend-aware model defaults

- backend/codewhisperer{,_internal}.go: full Amazon CodeWhisperer
  backend implementation — binary AWS event stream decoding, SQLite
  auth for both enterprise OIDC and personal social/GitHub flows, OIDC
  token refresh, and message encoding to the Kiro wire format
- backend/anthropic.go, copilot.go: new backends wired into New()
- backend/new.go: register anthropic, copilot, kiro/codewhisperer cases
- backend/openai.go: add extraHeaders hook for future use
- agent/loop.go: surface non-standard stop reasons as errors instead of
  silently dropping them
- main.go: defaultModelForBackend() sets a sensible default per backend
  (ollama→qwen3.5:9b, openrouter→deepseek/deepseek-v3.2,
   anthropic→claude-sonnet-4-5, kiro→auto); /backend switch now also
  resets the model to avoid stale foreign model IDs causing
  ValidationException; add -prompt flag for non-interactive batch mode

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Levi Neely 2026-04-09 20:02:23 +02:00
parent 4d4c1708e6
commit 868ae09364
9 changed files with 1889 additions and 13 deletions

View File

@ -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.

337
backend/anthropic.go Normal file
View File

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

807
backend/codewhisperer.go Normal file
View File

@ -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()
}

View File

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

20
backend/copilot.go Normal file
View File

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

View File

@ -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)
}
}

View File

@ -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 {

50
interrupt_commands.go Normal file
View File

@ -0,0 +1,50 @@
package main
import (
"fmt"
"strconv"
"strings"
)
const interruptCommandUsage = "/interrupt <prompt>"
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."
}

30
main.go
View File

@ -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())
}