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:
parent
4d4c1708e6
commit
868ae09364
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
@ -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()
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
30
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())
|
||||
}
|
||||
|
||||
|
|
|
|||
Reference in New Issue